use nautilus_core::{
UUID4, UnixNanos,
python::{IntoPyObjectNautilusExt, to_pyvalue_err},
};
use nautilus_model::identifiers::{ClientId, TraderId, Venue};
use pyo3::{basic::CompareOp, prelude::*};
use ustr::Ustr;
use crate::{
messages::system::{
QueueCondition, QueueState, QueueStateChanged, ReconnectSocket, SocketState,
SocketStateChanged, socket_endpoint,
},
runner::SystemChannel,
};
#[pymethods]
#[pyo3_stub_gen::derive::gen_stub_pymethods]
impl SystemChannel {
const fn __hash__(&self) -> isize {
*self as isize
}
}
#[pymethods]
#[pyo3_stub_gen::derive::gen_stub_pymethods]
impl QueueCondition {
const fn __hash__(&self) -> isize {
*self as isize
}
}
#[pymethods]
#[pyo3_stub_gen::derive::gen_stub_pymethods]
impl QueueState {
const fn __hash__(&self) -> isize {
*self as isize
}
}
#[pymethods]
#[pyo3_stub_gen::derive::gen_stub_pymethods]
impl SocketState {
const fn __hash__(&self) -> isize {
*self as isize
}
}
#[pymethods]
#[pyo3_stub_gen::derive::gen_stub_pymethods]
impl ReconnectSocket {
#[new]
fn py_new(
trader_id: TraderId,
client_id: ClientId,
endpoint: &str,
ts_init: u64,
) -> PyResult<Self> {
Ok(Self::new(
trader_id,
client_id,
socket_endpoint(endpoint).map_err(to_pyvalue_err)?,
UnixNanos::from(ts_init),
))
}
fn __richcmp__(&self, other: &Self, op: CompareOp, py: Python<'_>) -> Py<PyAny> {
match op {
CompareOp::Eq => self.eq(other).into_py_any_unwrap(py),
CompareOp::Ne => self.ne(other).into_py_any_unwrap(py),
_ => py.NotImplemented(),
}
}
fn __repr__(&self) -> String {
self.to_string()
}
#[getter]
const fn trader_id(&self) -> TraderId {
self.trader_id
}
#[getter]
const fn client_id(&self) -> ClientId {
self.client_id
}
#[getter]
fn endpoint(&self) -> &str {
self.endpoint.as_str()
}
#[getter]
const fn ts_init(&self) -> u64 {
self.ts_init.as_u64()
}
}
#[pymethods]
#[pyo3_stub_gen::derive::gen_stub_pymethods]
impl QueueStateChanged {
#[new]
#[expect(clippy::too_many_arguments)]
fn py_new(
trader_id: TraderId,
channel: SystemChannel,
condition: QueueCondition,
state: QueueState,
queue_depth: usize,
mean_dispatch_ns: u64,
event_id: UUID4,
ts_event: u64,
ts_init: u64,
) -> Self {
Self::new(
trader_id,
channel,
condition,
state,
queue_depth,
mean_dispatch_ns,
event_id,
UnixNanos::from(ts_event),
UnixNanos::from(ts_init),
)
}
fn __richcmp__(&self, other: &Self, op: CompareOp, py: Python<'_>) -> Py<PyAny> {
match op {
CompareOp::Eq => self.eq(other).into_py_any_unwrap(py),
CompareOp::Ne => self.ne(other).into_py_any_unwrap(py),
_ => py.NotImplemented(),
}
}
fn __repr__(&self) -> String {
self.to_string()
}
#[getter]
#[pyo3(name = "trader_id")]
const fn py_trader_id(&self) -> TraderId {
self.trader_id
}
#[getter]
#[pyo3(name = "channel")]
const fn py_channel(&self) -> SystemChannel {
self.channel
}
#[getter]
#[pyo3(name = "condition")]
const fn py_condition(&self) -> QueueCondition {
self.condition
}
#[getter]
#[pyo3(name = "state")]
const fn py_state(&self) -> QueueState {
self.state
}
#[getter]
#[pyo3(name = "queue_depth")]
const fn py_queue_depth(&self) -> usize {
self.queue_depth
}
#[getter]
#[pyo3(name = "mean_dispatch_ns")]
const fn py_mean_dispatch_ns(&self) -> u64 {
self.mean_dispatch_ns
}
#[getter]
#[pyo3(name = "event_id")]
const fn py_event_id(&self) -> UUID4 {
self.event_id
}
#[getter]
#[pyo3(name = "ts_event")]
const fn py_ts_event(&self) -> u64 {
self.ts_event.as_u64()
}
#[getter]
#[pyo3(name = "ts_init")]
const fn py_ts_init(&self) -> u64 {
self.ts_init.as_u64()
}
}
#[pymethods]
#[pyo3_stub_gen::derive::gen_stub_pymethods]
impl SocketStateChanged {
#[new]
#[expect(clippy::too_many_arguments)]
fn py_new(
trader_id: TraderId,
client_id: ClientId,
venue: Option<Venue>,
endpoint: &str,
state: SocketState,
event_id: UUID4,
ts_event: u64,
ts_init: u64,
) -> Self {
Self::new(
trader_id,
client_id,
venue,
Ustr::from(endpoint),
state,
event_id,
UnixNanos::from(ts_event),
UnixNanos::from(ts_init),
)
}
fn __richcmp__(&self, other: &Self, op: CompareOp, py: Python<'_>) -> Py<PyAny> {
match op {
CompareOp::Eq => self.eq(other).into_py_any_unwrap(py),
CompareOp::Ne => self.ne(other).into_py_any_unwrap(py),
_ => py.NotImplemented(),
}
}
fn __repr__(&self) -> String {
self.to_string()
}
#[getter]
#[pyo3(name = "trader_id")]
const fn py_trader_id(&self) -> TraderId {
self.trader_id
}
#[getter]
#[pyo3(name = "client_id")]
const fn py_client_id(&self) -> ClientId {
self.client_id
}
#[getter]
#[pyo3(name = "venue")]
const fn py_venue(&self) -> Option<Venue> {
self.venue
}
#[getter]
#[pyo3(name = "endpoint")]
fn py_endpoint(&self) -> &str {
self.endpoint.as_str()
}
#[getter]
#[pyo3(name = "state")]
const fn py_state(&self) -> SocketState {
self.state
}
#[getter]
#[pyo3(name = "event_id")]
const fn py_event_id(&self) -> UUID4 {
self.event_id
}
#[getter]
#[pyo3(name = "ts_event")]
const fn py_ts_event(&self) -> u64 {
self.ts_event.as_u64()
}
#[getter]
#[pyo3(name = "ts_init")]
const fn py_ts_init(&self) -> u64 {
self.ts_init.as_u64()
}
}
#[cfg(test)]
mod tests {
use pyo3::exceptions::PyValueError;
use rstest::rstest;
use super::*;
#[rstest]
fn reconnect_socket_python_constructor_assigns_all_fields() {
let trader_id = TraderId::from("TRADER-001");
let client_id = ClientId::from("POLYMARKET");
let command =
ReconnectSocket::py_new(trader_id, client_id, "polymarket-market-streams", 11).unwrap();
assert_eq!(command.trader_id, trader_id);
assert_eq!(command.client_id, client_id);
assert_eq!(command.endpoint.as_str(), "polymarket-market-streams");
assert_eq!(command.ts_init, UnixNanos::from(11));
}
#[rstest]
fn reconnect_socket_python_constructor_rejects_raw_urls() {
pyo3::Python::initialize();
let trader_id = TraderId::from("TRADER-001");
let client_id = ClientId::from("POLYMARKET");
Python::attach(|py| {
let command_error = ReconnectSocket::py_new(
trader_id,
client_id,
"wss://user:secret@example.com/feed",
11,
)
.unwrap_err();
assert!(command_error.is_instance_of::<PyValueError>(py));
assert_eq!(
command_error.value(py).to_string(),
"Socket endpoint must contain only ASCII letters, digits, '.', '-', or '_'",
);
});
}
}