Skip to main content

nautilus_infrastructure/python/redis/
msgbus.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 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    /// Configuration for a Redis-backed message bus backing.
40    ///
41    /// Redis 6.2 or higher is required for correct operation.
42    #[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    /// Creates a Redis message bus backing factory.
145    #[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}