nautilus_infrastructure/python/redis/
msgbus.rs1use bytes::Bytes;
17use futures::{pin_mut, stream::StreamExt};
18use nautilus_common::{
19 enums::SerializationEncoding,
20 msgbus::{
21 BusMessage, BusPayloadType, MessageBusBacking, MessageBusBackingFactory, MessageBusConfig,
22 },
23 python::{config_error_to_pyvalue_err, msgbus::get_global_msgbus_factory_registry},
24};
25use nautilus_core::{
26 UUID4,
27 python::{call_python, to_pyruntime_err, to_pyvalue_err},
28};
29use nautilus_model::identifiers::TraderId;
30use pyo3::{IntoPyObjectExt, prelude::*, pybacked::PyBackedBytes};
31use serde_json::Value;
32use ustr::Ustr;
33
34use crate::redis::msgbus::{RedisMessageBusBacking, RedisMessageBusConfig, RedisMessageBusFactory};
35
36#[pymethods]
37#[pyo3_stub_gen::derive::gen_stub_pymethods]
38impl RedisMessageBusConfig {
39 #[new]
43 #[expect(clippy::too_many_arguments)]
44 #[pyo3(signature = (host=None, port=None, username=None, password=None, ssl=None, connection_timeout=None, response_timeout=None, number_of_retries=None, exponent_base=None, max_delay=None, factor=None))]
45 fn py_new(
46 host: Option<String>,
47 port: Option<u16>,
48 username: Option<String>,
49 password: Option<String>,
50 ssl: Option<bool>,
51 connection_timeout: Option<u16>,
52 response_timeout: Option<u16>,
53 number_of_retries: Option<usize>,
54 exponent_base: Option<u64>,
55 max_delay: Option<u64>,
56 factor: Option<u64>,
57 ) -> Self {
58 let default = Self::default();
59 Self {
60 host,
61 port,
62 username,
63 password,
64 ssl: ssl.unwrap_or(default.ssl),
65 connection_timeout: connection_timeout.unwrap_or(default.connection_timeout),
66 response_timeout: response_timeout.unwrap_or(default.response_timeout),
67 number_of_retries: number_of_retries.unwrap_or(default.number_of_retries),
68 exponent_base: exponent_base.unwrap_or(default.exponent_base),
69 max_delay: max_delay.unwrap_or(default.max_delay),
70 factor: factor.unwrap_or(default.factor),
71 }
72 }
73
74 #[getter]
75 fn host(&self) -> Option<&str> {
76 self.host.as_deref()
77 }
78
79 #[getter]
80 const fn port(&self) -> Option<u16> {
81 self.port
82 }
83
84 #[getter]
85 fn username(&self) -> Option<&str> {
86 self.username.as_deref()
87 }
88
89 #[getter]
90 fn password(&self) -> Option<&str> {
91 self.password.as_deref()
92 }
93
94 #[getter]
95 const fn ssl(&self) -> bool {
96 self.ssl
97 }
98
99 #[getter]
100 const fn connection_timeout(&self) -> u16 {
101 self.connection_timeout
102 }
103
104 #[getter]
105 const fn response_timeout(&self) -> u16 {
106 self.response_timeout
107 }
108
109 #[getter]
110 const fn number_of_retries(&self) -> usize {
111 self.number_of_retries
112 }
113
114 #[getter]
115 const fn exponent_base(&self) -> u64 {
116 self.exponent_base
117 }
118
119 #[getter]
120 const fn max_delay(&self) -> u64 {
121 self.max_delay
122 }
123
124 #[getter]
125 const fn factor(&self) -> u64 {
126 self.factor
127 }
128}
129
130#[derive(Debug, Clone)]
131#[pyclass(
132 name = "RedisMessageBusFactory",
133 module = "nautilus_trader.infrastructure",
134 from_py_object
135)]
136#[pyo3_stub_gen::derive::gen_stub_pyclass(module = "nautilus_trader.infrastructure")]
137pub struct PyRedisMessageBusFactory {
138 inner: RedisMessageBusFactory,
139}
140
141#[pymethods]
142#[pyo3_stub_gen::derive::gen_stub_pymethods]
143impl PyRedisMessageBusFactory {
144 #[new]
146 #[pyo3(signature = (config=None))]
147 fn py_new(config: Option<RedisMessageBusConfig>) -> Self {
148 Self {
149 inner: RedisMessageBusFactory::new(config.unwrap_or_default()),
150 }
151 }
152}
153
154#[expect(clippy::needless_pass_by_value)]
155fn extract_redis_msgbus_config(
156 py: Python<'_>,
157 factory: Py<PyAny>,
158) -> PyResult<Box<dyn MessageBusBackingFactory>> {
159 Ok(Box::new(factory.extract::<RedisMessageBusConfig>(py)?))
160}
161
162#[expect(clippy::needless_pass_by_value)]
163fn extract_redis_msgbus_factory(
164 py: Python<'_>,
165 factory: Py<PyAny>,
166) -> PyResult<Box<dyn MessageBusBackingFactory>> {
167 let factory = factory.extract::<PyRedisMessageBusFactory>(py)?;
168 Ok(Box::new(factory.inner))
169}
170
171pub(in crate::python) fn register_redis_msgbus_factory() -> PyResult<()> {
172 let registry = get_global_msgbus_factory_registry();
173 registry
174 .register(
175 stringify!(RedisMessageBusConfig).to_string(),
176 extract_redis_msgbus_config,
177 )
178 .map_err(to_pyruntime_err)?;
179 registry
180 .register(
181 stringify!(RedisMessageBusFactory).to_string(),
182 extract_redis_msgbus_factory,
183 )
184 .map_err(to_pyruntime_err)
185}
186
187#[derive(Debug)]
188#[pyclass(
189 name = "RedisMessageBusBacking",
190 module = "nautilus_trader.infrastructure"
191)]
192#[pyo3_stub_gen::derive::gen_stub_pyclass(module = "nautilus_trader.infrastructure")]
193pub struct PyRedisMessageBusBacking {
194 inner: RedisMessageBusBacking,
195}
196
197#[pymethods]
198#[pyo3_stub_gen::derive::gen_stub_pymethods]
199impl PyRedisMessageBusBacking {
200 #[new]
201 #[expect(
202 clippy::needless_pass_by_value,
203 reason = "PyBackedBytes is required for generated Python bytes stubs"
204 )]
205 fn py_new(
206 trader_id: TraderId,
207 instance_id: UUID4,
208 config_json: PyBackedBytes,
209 ) -> PyResult<Self> {
210 let (config, backing) = parse_config(config_json.as_ref())?;
211 let inner = RedisMessageBusBacking::new(trader_id, instance_id, config, backing)
212 .map_err(to_pyvalue_err)?;
213 Ok(Self { inner })
214 }
215
216 #[pyo3(name = "is_closed")]
217 fn py_is_closed(&self) -> bool {
218 MessageBusBacking::is_closed(&self.inner)
219 }
220
221 #[pyo3(name = "publish")]
222 #[expect(
223 clippy::needless_pass_by_value,
224 reason = "PyBackedBytes is required for generated Python bytes stubs"
225 )]
226 fn py_publish(&self, topic: &str, payload: PyBackedBytes) {
227 let message = BusMessage::new(
228 Ustr::from(topic),
229 BusPayloadType::Custom(Ustr::default()),
230 Bytes::copy_from_slice(payload.as_ref()),
231 SerializationEncoding::default(),
232 );
233 MessageBusBacking::publish(&self.inner, message);
234 }
235
236 #[pyo3(name = "stream")]
237 fn py_stream<'py>(
238 &mut self,
239 callback: Py<PyAny>,
240 py: Python<'py>,
241 ) -> PyResult<Bound<'py, PyAny>> {
242 let stream_rx = self.inner.get_stream_receiver().map_err(to_pyruntime_err)?;
243 let stream = RedisMessageBusBacking::stream(stream_rx);
244 pyo3_async_runtimes::tokio::future_into_py(py, async move {
245 pin_mut!(stream);
246 while let Some(msg) = stream.next().await {
247 Python::attach(|py| -> PyResult<()> {
248 call_python(py, &callback, msg.into_py_any(py)?);
249 Ok(())
250 })?;
251 }
252 Ok(())
253 })
254 }
255
256 #[pyo3(name = "close")]
257 fn py_close(&mut self) {
258 MessageBusBacking::close(&mut self.inner);
259 }
260}
261
262fn parse_config(config_json: &[u8]) -> PyResult<(MessageBusConfig, RedisMessageBusConfig)> {
263 let mut value: Value = serde_json::from_slice(config_json).map_err(to_pyvalue_err)?;
264 let backing = parse_backing_config(&mut value)?;
265 let config = serde_json::from_value::<MessageBusConfig>(value).map_err(to_pyvalue_err)?;
266 config.validate().map_err(config_error_to_pyvalue_err)?;
267
268 Ok((config, backing))
269}
270
271fn parse_backing_config(value: &mut Value) -> PyResult<RedisMessageBusConfig> {
272 let Value::Object(config) = value else {
273 return Err(to_pyvalue_err("MessageBusConfig must be a JSON object"));
274 };
275
276 let Some(database) = config.remove("database") else {
277 return Ok(RedisMessageBusConfig::default());
278 };
279
280 let mut database = match database {
281 Value::Null => return Ok(RedisMessageBusConfig::default()),
282 Value::Object(database) => database,
283 _ => {
284 return Err(to_pyvalue_err(
285 "MessageBusConfig.database must be a JSON object",
286 ));
287 }
288 };
289
290 if let Some(database_type) = database.remove("type") {
291 match database_type {
292 Value::String(database_type) if database_type == "redis" => {}
293 Value::String(database_type) => {
294 return Err(to_pyvalue_err(format!(
295 "MessageBusConfig.database.type must be 'redis', was '{database_type}'"
296 )));
297 }
298 other => {
299 return Err(to_pyvalue_err(format!(
300 "MessageBusConfig.database.type must be a string, was {other}"
301 )));
302 }
303 }
304 }
305
306 serde_json::from_value(Value::Object(database)).map_err(to_pyvalue_err)
307}
308
309#[cfg(test)]
310mod tests {
311 use rstest::rstest;
312 use serde_json::json;
313
314 use super::*;
315
316 #[rstest]
317 fn test_parse_config_splits_legacy_database_config() {
318 let config_json = json!({
319 "database": {
320 "type": "redis",
321 "host": "localhost",
322 "port": 6380,
323 "ssl": true,
324 },
325 "buffer_interval_ms": 100,
326 "streams_prefix": "signals",
327 "stream_per_topic": false,
328 "external_streams": ["signals"],
329 });
330
331 let (config, backing) = parse_config(config_json.to_string().as_bytes()).unwrap();
332
333 assert_eq!(config.buffer_interval_ms, Some(100));
334 assert_eq!(config.streams_prefix, "signals");
335 assert!(!config.stream_per_topic);
336 assert_eq!(config.external_streams, Some(vec!["signals".to_string()]));
337 assert_eq!(backing.host, Some("localhost".to_string()));
338 assert_eq!(backing.port, Some(6380));
339 assert!(backing.ssl);
340 }
341
342 #[rstest]
343 fn test_parse_config_accepts_python_message_bus_config_json() {
344 let config_json = json!({
345 "database": {
346 "type": "redis",
347 "host": "redis.example.com",
348 "port": 6380,
349 "username": "user",
350 "password": "secret",
351 "ssl": true,
352 "connection_timeout": 30,
353 "response_timeout": 10,
354 "number_of_retries": 3,
355 "exponent_base": 3,
356 "max_delay": 15,
357 "factor": 4,
358 },
359 "encoding": "msgpack",
360 "timestamps_as_iso8601": true,
361 "buffer_interval_ms": null,
362 "autotrim_mins": null,
363 "use_trader_prefix": true,
364 "use_trader_id": false,
365 "use_instance_id": true,
366 "streams_prefix": "stream",
367 "stream_per_topic": false,
368 "external_streams": ["signals"],
369 "types_filter": ["nautilus_trader.model.data:QuoteTick"],
370 "heartbeat_interval_secs": null,
371 });
372
373 let (config, backing) = parse_config(config_json.to_string().as_bytes()).unwrap();
374
375 assert_eq!(config.encoding, SerializationEncoding::MsgPack);
376 assert!(config.timestamps_as_iso8601);
377 assert_eq!(config.buffer_interval_ms, None);
378 assert_eq!(config.autotrim_mins, None);
379 assert!(config.use_trader_prefix);
380 assert!(!config.use_trader_id);
381 assert!(config.use_instance_id);
382 assert_eq!(config.streams_prefix, "stream");
383 assert!(!config.stream_per_topic);
384 assert_eq!(config.external_streams, Some(vec!["signals".to_string()]));
385 assert_eq!(
386 config.types_filter,
387 Some(vec!["nautilus_trader.model.data:QuoteTick".to_string()])
388 );
389 assert_eq!(config.heartbeat_interval_secs, None);
390 assert_eq!(backing.host, Some("redis.example.com".to_string()));
391 assert_eq!(backing.port, Some(6380));
392 assert_eq!(backing.username, Some("user".to_string()));
393 assert_eq!(backing.password, Some("secret".to_string()));
394 assert!(backing.ssl);
395 assert_eq!(backing.connection_timeout, 30);
396 assert_eq!(backing.response_timeout, 10);
397 assert_eq!(backing.number_of_retries, 3);
398 assert_eq!(backing.exponent_base, 3);
399 assert_eq!(backing.max_delay, 15);
400 assert_eq!(backing.factor, 4);
401 }
402
403 #[rstest]
404 fn test_parse_config_rejects_non_redis_database_type() {
405 Python::initialize();
406 let config_json = json!({
407 "database": {
408 "type": "postgres",
409 },
410 });
411
412 let result = parse_config(config_json.to_string().as_bytes());
413
414 assert_eq!(
415 result.unwrap_err().to_string(),
416 "ValueError: MessageBusConfig.database.type must be 'redis', was 'postgres'"
417 );
418 }
419}