nautilus_infrastructure/python/sql/
cache.rs1use bytes::Bytes;
17use nautilus_common::{
18 cache::database::{CacheDatabaseAdapter, CacheDatabaseFactory},
19 live::get_runtime,
20 python::cache::get_global_cache_database_factory_registry,
21 signal::Signal,
22};
23use nautilus_core::python::to_pyruntime_err;
24use nautilus_model::{
25 data::{Bar, CustomData, DataType, QuoteTick, TradeTick},
26 events::{OrderSnapshot, PositionSnapshot},
27 identifiers::{AccountId, ClientId, ClientOrderId, InstrumentId, PositionId},
28 python::{
29 account::{account_any_to_pyobject, pyobject_to_account_any},
30 events::order::pyobject_to_order_event,
31 instruments::{instrument_any_to_pyobject, pyobject_to_instrument_any},
32 orders::{order_any_to_pyobject, pyobject_to_order_any},
33 },
34 types::Currency,
35};
36use pyo3::{IntoPyObjectExt, prelude::*};
37
38use crate::sql::{
39 cache::{PostgresCacheConfig, PostgresCacheDatabase},
40 queries::DatabaseQueries,
41};
42
43#[pymethods]
44impl PostgresCacheDatabase {
45 #[staticmethod]
51 #[pyo3(name = "connect")]
52 #[pyo3(signature = (host=None, port=None, username=None, password=None, database=None))]
53 fn py_connect(
54 host: Option<String>,
55 port: Option<u16>,
56 username: Option<String>,
57 password: Option<String>,
58 database: Option<String>,
59 ) -> PyResult<Self> {
60 let result = get_runtime()
61 .block_on(async { Self::connect(host, port, username, password, database).await });
62 result.map_err(to_pyruntime_err)
63 }
64
65 #[pyo3(name = "close")]
66 fn py_close(&mut self) -> PyResult<()> {
67 self.close().map_err(to_pyruntime_err)
68 }
69
70 #[pyo3(name = "flush_db")]
71 fn py_flush_db(&mut self) -> PyResult<()> {
72 self.flush().map_err(to_pyruntime_err)
73 }
74
75 #[pyo3(name = "load")]
76 fn py_load(&self) -> PyResult<std::collections::HashMap<String, Vec<u8>>> {
77 get_runtime()
78 .block_on(async { DatabaseQueries::load(&self.pool).await })
79 .map(|m| m.into_iter().collect())
80 .map_err(to_pyruntime_err)
81 }
82
83 #[pyo3(name = "load_currency")]
84 fn py_load_currency(&self, code: &str) -> PyResult<Option<Currency>> {
85 let result = get_runtime()
86 .block_on(async { DatabaseQueries::load_currency(&self.pool, code).await });
87 result.map_err(to_pyruntime_err)
88 }
89
90 #[pyo3(name = "load_currencies")]
91 fn py_load_currencies(&self) -> PyResult<Vec<Currency>> {
92 let result =
93 get_runtime().block_on(async { DatabaseQueries::load_currencies(&self.pool).await });
94 result.map_err(to_pyruntime_err)
95 }
96
97 #[pyo3(name = "load_instrument")]
98 fn py_load_instrument(
99 &self,
100 py: Python,
101 instrument_id: InstrumentId,
102 ) -> PyResult<Option<Py<PyAny>>> {
103 get_runtime().block_on(async {
104 let result = DatabaseQueries::load_instrument(&self.pool, &instrument_id)
105 .await
106 .map_err(to_pyruntime_err)?;
107
108 match result {
109 Some(instrument) => {
110 let py_object = instrument_any_to_pyobject(py, instrument)?;
111 Ok(Some(py_object))
112 }
113 None => Ok(None),
114 }
115 })
116 }
117
118 #[pyo3(name = "load_instruments")]
119 fn py_load_instruments(&self, py: Python) -> PyResult<Vec<Py<PyAny>>> {
120 get_runtime().block_on(async {
121 let result = DatabaseQueries::load_instruments(&self.pool)
122 .await
123 .map_err(to_pyruntime_err)?;
124 let mut instruments = Vec::new();
125
126 for instrument in result {
127 let py_object = instrument_any_to_pyobject(py, instrument)?;
128 instruments.push(py_object);
129 }
130 Ok(instruments)
131 })
132 }
133
134 #[pyo3(name = "load_order")]
135 fn py_load_order(
136 &self,
137 py: Python,
138 client_order_id: ClientOrderId,
139 ) -> PyResult<Option<Py<PyAny>>> {
140 get_runtime().block_on(async {
141 let result = DatabaseQueries::load_order(&self.pool, &client_order_id)
142 .await
143 .map_err(to_pyruntime_err)?;
144
145 match result {
146 Some(order) => {
147 let py_object = order_any_to_pyobject(py, order)?;
148 Ok(Some(py_object))
149 }
150 None => Ok(None),
151 }
152 })
153 }
154
155 #[pyo3(name = "load_account")]
156 fn py_load_account(&self, py: Python, account_id: AccountId) -> PyResult<Option<Py<PyAny>>> {
157 get_runtime().block_on(async {
158 let result = DatabaseQueries::load_account(&self.pool, &account_id)
159 .await
160 .map_err(to_pyruntime_err)?;
161
162 match result {
163 Some(account) => {
164 let py_object = account_any_to_pyobject(py, account)?;
165 Ok(Some(py_object))
166 }
167 None => Ok(None),
168 }
169 })
170 }
171
172 #[pyo3(name = "load_quotes")]
173 fn py_load_quotes(&self, py: Python, instrument_id: InstrumentId) -> PyResult<Vec<Py<PyAny>>> {
174 get_runtime().block_on(async {
175 let result = DatabaseQueries::load_quotes(&self.pool, &instrument_id)
176 .await
177 .map_err(to_pyruntime_err)?;
178 let mut quotes = Vec::new();
179
180 for quote in result {
181 let py_object = quote.into_py_any(py)?;
182 quotes.push(py_object);
183 }
184 Ok(quotes)
185 })
186 }
187
188 #[pyo3(name = "load_trades")]
189 fn py_load_trades(&self, py: Python, instrument_id: InstrumentId) -> PyResult<Vec<Py<PyAny>>> {
190 get_runtime().block_on(async {
191 let result = DatabaseQueries::load_trades(&self.pool, &instrument_id)
192 .await
193 .map_err(to_pyruntime_err)?;
194 let mut trades = Vec::new();
195
196 for trade in result {
197 let py_object = trade.into_py_any(py)?;
198 trades.push(py_object);
199 }
200 Ok(trades)
201 })
202 }
203
204 #[pyo3(name = "load_bars")]
205 fn py_load_bars(&self, py: Python, instrument_id: InstrumentId) -> PyResult<Vec<Py<PyAny>>> {
206 get_runtime().block_on(async {
207 let result = DatabaseQueries::load_bars(&self.pool, &instrument_id)
208 .await
209 .map_err(to_pyruntime_err)?;
210 let mut bars = Vec::new();
211
212 for bar in result {
213 let py_object = bar.into_py_any(py)?;
214 bars.push(py_object);
215 }
216 Ok(bars)
217 })
218 }
219
220 #[pyo3(name = "load_signals")]
221 fn py_load_signals(&self, name: &str) -> PyResult<Vec<Signal>> {
222 get_runtime().block_on(async {
223 DatabaseQueries::load_signals(&self.pool, name)
224 .await
225 .map_err(to_pyruntime_err)
226 })
227 }
228
229 #[pyo3(name = "load_custom_data")]
230 #[expect(clippy::needless_pass_by_value)]
231 fn py_load_custom_data(&self, data_type: DataType) -> PyResult<Vec<CustomData>> {
232 get_runtime()
233 .block_on(async { DatabaseQueries::load_custom_data(&self.pool, &data_type).await })
234 .map_err(to_pyruntime_err)
235 }
236
237 #[pyo3(name = "load_order_snapshot")]
238 fn py_load_order_snapshot(
239 &self,
240 client_order_id: ClientOrderId,
241 ) -> PyResult<Option<OrderSnapshot>> {
242 get_runtime().block_on(async {
243 DatabaseQueries::load_order_snapshot(&self.pool, &client_order_id)
244 .await
245 .map_err(to_pyruntime_err)
246 })
247 }
248
249 #[pyo3(name = "load_position_snapshot")]
250 fn py_load_position_snapshot(
251 &self,
252 position_id: PositionId,
253 ) -> PyResult<Option<PositionSnapshot>> {
254 get_runtime().block_on(async {
255 DatabaseQueries::load_position_snapshot(&self.pool, &position_id)
256 .await
257 .map_err(to_pyruntime_err)
258 })
259 }
260
261 #[pyo3(name = "add")]
262 fn py_add(&self, key: String, value: Vec<u8>) -> PyResult<()> {
263 self.add(key, Bytes::from(value)).map_err(to_pyruntime_err)
264 }
265
266 #[pyo3(name = "add_currency")]
267 fn py_add_currency(&self, currency: Currency) -> PyResult<()> {
268 self.add_currency(¤cy).map_err(to_pyruntime_err)
269 }
270
271 #[pyo3(name = "add_instrument")]
272 fn py_add_instrument(&self, py: Python, instrument: Py<PyAny>) -> PyResult<()> {
273 let instrument_any = pyobject_to_instrument_any(py, instrument)?;
274 self.add_instrument(&instrument_any)
275 .map_err(to_pyruntime_err)
276 }
277
278 #[pyo3(name = "add_order")]
279 #[pyo3(signature = (order, client_id=None))]
280 fn py_add_order(
281 &self,
282 py: Python,
283 order: Py<PyAny>,
284 client_id: Option<ClientId>,
285 ) -> PyResult<()> {
286 let order_any = pyobject_to_order_any(py, order)?;
287 self.add_order(&order_any, client_id)
288 .map_err(to_pyruntime_err)
289 }
290
291 #[pyo3(name = "add_order_snapshot")]
292 #[expect(clippy::needless_pass_by_value)]
293 fn py_add_order_snapshot(&self, snapshot: OrderSnapshot) -> PyResult<()> {
294 self.add_order_snapshot(&snapshot).map_err(to_pyruntime_err)
295 }
296
297 #[pyo3(name = "add_position_snapshot")]
298 #[expect(clippy::needless_pass_by_value)]
299 fn py_add_position_snapshot(&self, snapshot: PositionSnapshot) -> PyResult<()> {
300 self.add_position_snapshot(&snapshot)
301 .map_err(to_pyruntime_err)
302 }
303
304 #[pyo3(name = "add_account")]
305 fn py_add_account(&self, py: Python, account: Py<PyAny>) -> PyResult<()> {
306 let account_any = pyobject_to_account_any(py, account)?;
307 self.add_account(&account_any).map_err(to_pyruntime_err)
308 }
309
310 #[pyo3(name = "add_quote")]
311 fn py_add_quote(&self, quote: QuoteTick) -> PyResult<()> {
312 self.add_quote("e).map_err(to_pyruntime_err)
313 }
314
315 #[pyo3(name = "add_trade")]
316 fn py_add_trade(&self, trade: TradeTick) -> PyResult<()> {
317 self.add_trade(&trade).map_err(to_pyruntime_err)
318 }
319
320 #[pyo3(name = "add_bar")]
321 fn py_add_bar(&self, bar: Bar) -> PyResult<()> {
322 self.add_bar(&bar).map_err(to_pyruntime_err)
323 }
324
325 #[pyo3(name = "add_signal")]
326 #[expect(clippy::needless_pass_by_value)]
327 fn py_add_signal(&self, signal: Signal) -> PyResult<()> {
328 self.add_signal(&signal).map_err(to_pyruntime_err)
329 }
330
331 #[pyo3(name = "add_custom_data")]
332 #[expect(clippy::needless_pass_by_value)]
333 fn py_add_custom_data(&self, data: CustomData) -> PyResult<()> {
334 self.add_custom_data(&data).map_err(to_pyruntime_err)
335 }
336
337 #[pyo3(name = "update_order")]
338 fn py_update_order(&self, py: Python, order_event: Py<PyAny>) -> PyResult<()> {
339 let event = pyobject_to_order_event(py, order_event)?;
340 self.update_order(&event).map_err(to_pyruntime_err)
341 }
342
343 #[pyo3(name = "update_account")]
344 fn py_update_account(&self, py: Python, order: Py<PyAny>) -> PyResult<()> {
345 let order_any = pyobject_to_account_any(py, order)?;
346 self.update_account(&order_any).map_err(to_pyruntime_err)
347 }
348}
349
350#[pymethods]
351#[pyo3_stub_gen::derive::gen_stub_pymethods]
352impl PostgresCacheConfig {
353 #[new]
357 #[pyo3(signature = (host=None, port=None, username=None, password=None, database=None))]
358 fn py_new(
359 host: Option<String>,
360 port: Option<u16>,
361 username: Option<String>,
362 password: Option<String>,
363 database: Option<String>,
364 ) -> Self {
365 Self {
366 host,
367 port,
368 username,
369 password,
370 database,
371 }
372 }
373
374 #[getter]
375 fn host(&self) -> Option<&str> {
376 self.host.as_deref()
377 }
378
379 #[getter]
380 const fn port(&self) -> Option<u16> {
381 self.port
382 }
383
384 #[getter]
385 fn username(&self) -> Option<&str> {
386 self.username.as_deref()
387 }
388
389 #[getter]
390 fn password(&self) -> Option<&str> {
391 self.password.as_deref()
392 }
393
394 #[getter]
395 fn database(&self) -> Option<&str> {
396 self.database.as_deref()
397 }
398}
399
400#[expect(clippy::needless_pass_by_value)]
401fn extract_postgres_cache_database_factory(
402 py: Python<'_>,
403 factory: Py<PyAny>,
404) -> PyResult<Box<dyn CacheDatabaseFactory>> {
405 Ok(Box::new(factory.extract::<PostgresCacheConfig>(py)?))
406}
407
408pub(in crate::python) fn register_postgres_cache_database_factory() -> PyResult<()> {
409 get_global_cache_database_factory_registry()
410 .register(
411 stringify!(PostgresCacheConfig).to_string(),
412 extract_postgres_cache_database_factory,
413 )
414 .map_err(to_pyruntime_err)
415}