use std::collections::HashMap;
use nautilus_core::python::{to_pyruntime_err, to_pyvalue_err};
use nautilus_model::{
data::{
Bar, InstrumentStatus, MarkPriceUpdate, OptionGreeks, OrderBookDelta, OrderBookDepth10,
QuoteTick, TradeTick,
},
python::data::data_to_pyobject,
};
use nautilus_serialization::arrow::{ArrowSchemaProvider, custom::CustomDataDecoder};
use pyo3::{IntoPyObjectExt, prelude::*};
use crate::backend::session::{DataBackendSession, DataQueryResult, QueryError};
struct SendPtr<T>(*mut T);
unsafe impl<T> Send for SendPtr<T> {}
#[repr(C)]
#[pyclass(frozen, eq, eq_int, from_py_object)]
#[pyo3_stub_gen::derive::gen_stub_pyclass_enum(module = "nautilus_trader.persistence")]
#[derive(Clone, Copy, Debug, Hash, PartialEq, Eq)]
pub enum NautilusDataType {
OrderBookDelta = 1,
OrderBookDepth10 = 2,
QuoteTick = 3,
TradeTick = 4,
Bar = 5,
MarkPriceUpdate = 6,
OptionGreeks = 7,
InstrumentStatus = 8,
}
#[pymethods]
#[pyo3_stub_gen::derive::gen_stub_pymethods]
impl NautilusDataType {
#[expect(
clippy::trivially_copy_pass_by_ref,
reason = "PyO3 special methods use a borrowed receiver"
)]
const fn __hash__(&self) -> isize {
*self as isize
}
}
#[pymethods]
#[pyo3_stub_gen::derive::gen_stub_pymethods]
impl DataBackendSession {
#[new]
#[pyo3(signature=(chunk_size=10_000))]
fn py_new(chunk_size: usize) -> PyResult<Self> {
if chunk_size == 0 {
return Err(to_pyvalue_err("chunk_size must be positive"));
}
Ok(Self::new(chunk_size))
}
#[pyo3(name = "add_file")]
#[pyo3(signature = (data_type, table_name, file_path, sql_query=None))]
fn py_add_file(
mut slf: PyRefMut<'_, Self>,
data_type: NautilusDataType,
table_name: &str,
file_path: &str,
sql_query: Option<&str>,
) -> PyResult<()> {
let _guard = slf.runtime.enter();
match data_type {
NautilusDataType::OrderBookDelta => slf
.add_file::<OrderBookDelta>(table_name, file_path, sql_query, None)
.map_err(to_pyruntime_err),
NautilusDataType::OrderBookDepth10 => slf
.add_file::<OrderBookDepth10>(table_name, file_path, sql_query, None)
.map_err(to_pyruntime_err),
NautilusDataType::QuoteTick => slf
.add_file::<QuoteTick>(table_name, file_path, sql_query, None)
.map_err(to_pyruntime_err),
NautilusDataType::TradeTick => slf
.add_file::<TradeTick>(table_name, file_path, sql_query, None)
.map_err(to_pyruntime_err),
NautilusDataType::Bar => slf
.add_file::<Bar>(table_name, file_path, sql_query, None)
.map_err(to_pyruntime_err),
NautilusDataType::MarkPriceUpdate => slf
.add_file::<MarkPriceUpdate>(table_name, file_path, sql_query, None)
.map_err(to_pyruntime_err),
NautilusDataType::OptionGreeks => slf
.add_file::<OptionGreeks>(table_name, file_path, sql_query, None)
.map_err(to_pyruntime_err),
NautilusDataType::InstrumentStatus => slf
.add_file::<InstrumentStatus>(table_name, file_path, sql_query, None)
.map_err(to_pyruntime_err),
}
}
#[pyo3(name = "add_custom_file")]
#[pyo3(signature = (type_name, table_name, file_path, sql_query=None))]
fn py_add_custom_file(
mut slf: PyRefMut<'_, Self>,
type_name: &str,
table_name: &str,
file_path: &str,
sql_query: Option<&str>,
) -> PyResult<()> {
let _guard = slf.runtime.enter();
let mut metadata = HashMap::new();
metadata.insert("type_name".to_string(), type_name.to_string());
let base_schema = CustomDataDecoder::get_schema(Some(metadata));
base_schema.field_with_name("ts_init").map_err(|_| {
to_pyruntime_err(format!(
"custom data type '{type_name}' is not registered with an Arrow schema containing ts_init"
))
})?;
slf.add_file::<CustomDataDecoder>(table_name, file_path, sql_query, Some(type_name))
.map_err(to_pyruntime_err)
}
fn to_query_result(mut slf: PyRefMut<'_, Self>) -> DataQueryResult {
let py = slf.py();
let chunk_size = slf.chunk_size;
let ptr = SendPtr(&raw mut *slf);
let query_result = unsafe {
py.detach(move || {
let p = ptr;
(*p.0).get_query_result()
})
};
DataQueryResult::new(query_result, chunk_size)
}
#[pyo3(name = "register_object_store_from_uri")]
#[pyo3(signature = (uri, storage_options=None))]
fn py_register_object_store_from_uri(
mut slf: PyRefMut<'_, Self>,
uri: &str,
storage_options: Option<HashMap<String, String>>,
) -> PyResult<()> {
let storage_options = storage_options.map(|m| m.into_iter().collect());
slf.register_object_store_from_uri(uri, storage_options)
.map_err(to_pyruntime_err)
}
}
#[pymethods]
#[pyo3_stub_gen::derive::gen_stub_pymethods]
impl DataQueryResult {
#[pyo3(name = "to_list")]
fn py_to_list(mut slf: PyRefMut<'_, Self>) -> PyResult<Vec<Py<PyAny>>> {
let py = slf.py();
let ptr = SendPtr(&raw mut *slf);
let data = unsafe {
py.detach(move || -> Result<Vec<_>, QueryError> {
let p = ptr;
let result = &mut *p.0;
let mut data = Vec::new();
for chunk in result.by_ref() {
let chunk = chunk?;
if chunk.is_empty() {
break;
}
data.extend(chunk);
}
Ok(data)
})
}
.map_err(to_pyruntime_err)?;
data.into_iter()
.map(|item| data_to_pyobject(py, item))
.collect()
}
const fn __iter__(slf: PyRef<'_, Self>) -> PyRef<'_, Self> {
slf
}
fn __next__(mut slf: PyRefMut<'_, Self>) -> PyResult<Option<Py<PyAny>>> {
let py = slf.py();
let ptr = SendPtr(&raw mut *slf);
let acc = unsafe {
py.detach(move || {
let p = ptr;
(*p.0).next()
})
};
match acc {
Some(Ok(acc)) if !acc.is_empty() => {
let objects: Vec<Py<PyAny>> = acc
.into_iter()
.map(|item| data_to_pyobject(py, item))
.collect::<PyResult<_>>()?;
Ok(Some(objects.into_py_any(py)?))
}
Some(Err(e)) => Err(to_pyruntime_err(e)),
_ => Ok(None),
}
}
}