diff --git a/python/Cargo.toml b/python/Cargo.toml index 0bc3de1a..1a50135e 100644 --- a/python/Cargo.toml +++ b/python/Cargo.toml @@ -14,6 +14,7 @@ serde_yaml = "0.8.23" flume = "0.10.14" arrow = { version = "45.0.0", features = ["pyarrow"] } pythonize = "0.19.0" +futures = "0.3.28" [lib] name = "dora_ros2_bridge" diff --git a/python/src/lib.rs b/python/src/lib.rs index f18f4d77..ecd58d30 100644 --- a/python/src/lib.rs +++ b/python/src/lib.rs @@ -8,10 +8,11 @@ use std::{ use ::dora_ros2_bridge::{ros2_client, rustdds}; use dora_ros2_bridge_msg_gen::types::Message; use eyre::{eyre, Context, ContextCompat}; +use futures::{Stream, StreamExt}; use pyo3::{ prelude::{pyclass, pymethods, pymodule}, types::PyModule, - PyAny, PyObject, PyResult, Python, + wrap_pyfunction, PyAny, PyErr, PyObject, PyResult, Python, ToPyObject, }; use typed::{ deserialize::{Ros2Value, TypedDeserializer}, @@ -226,6 +227,29 @@ impl Ros2Subscription { } } +impl Ros2Subscription { + fn as_stream( + &self, + ) -> impl Stream> + '_ + { + self.subscription + .async_stream_seed(self.deserializer.clone()) + } +} + +impl Stream for Ros2Subscription { + type Item = Result<(Ros2Value, ros2_client::MessageInfo), rustdds::dds::ReadError>; + + fn poll_next( + self: std::pin::Pin<&mut Self>, + cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + let s = self.as_stream(); + futures::pin_mut!(s); + s.poll_next_unpin(cx) + } +} + #[pymodule] fn dora_ros2_bridge(_py: Python, m: &PyModule) -> PyResult<()> { m.add_class::()?;