Skip to main content

nautilus_infrastructure/python/sql/
cache.rs

1// -------------------------------------------------------------------------------------------------
2//  Copyright (C) 2015-2026 Nautech Systems Pty Ltd. All rights reserved.
3//  https://nautechsystems.io
4//
5//  Licensed under the GNU Lesser General Public License Version 3.0 (the "License");
6//  You may not use this file except in compliance with the License.
7//  You may obtain a copy of the License at https://www.gnu.org/licenses/lgpl-3.0.en.html
8//
9//  Unless required by applicable law or agreed to in writing, software
10//  distributed under the License is distributed on an "AS IS" BASIS,
11//  WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12//  See the License for the specific language governing permissions and
13//  limitations under the License.
14// -------------------------------------------------------------------------------------------------
15
16use 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    /// Connects to the Postgres cache database using the provided connection parameters.
46    ///
47    /// # Errors
48    ///
49    /// Returns an error if establishing the database connection fails.
50    #[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(&currency).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(&quote).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    /// Configuration for a Postgres-backed cache database.
354    ///
355    /// Missing fields are resolved from Postgres environment variables and then built-in defaults.
356    #[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}