use std::future::Future;
use std::pin::Pin;
use std::time::Duration;
use futures::stream::{Stream, StreamExt};
use pamoja_core::{Actuator, Error, Result, Sensor};
use r2r::{Context, Node, Publisher, QosProfile, WrappedTypesupport};
pub struct Ros2Node {
node: Node,
}
impl Ros2Node {
pub fn new(name: &str, namespace: &str) -> Result<Self> {
let context = Context::create().map_err(map_err)?;
let node = Node::create(context, name, namespace).map_err(map_err)?;
Ok(Self { node })
}
pub fn publisher<T: WrappedTypesupport>(&mut self, topic: &str) -> Result<RosPublisher<T>> {
let publisher = self
.node
.create_publisher::<T>(topic, QosProfile::default())
.map_err(map_err)?;
Ok(RosPublisher { publisher })
}
pub fn subscriber<T: WrappedTypesupport + Send + 'static>(
&mut self,
topic: &str,
) -> Result<RosSubscriber<T>> {
let stream = self
.node
.subscribe::<T>(topic, QosProfile::default())
.map_err(map_err)?;
Ok(RosSubscriber {
stream: Box::pin(stream),
})
}
pub fn service<S: r2r::WrappedServiceTypeSupport + Send + 'static>(
&mut self,
name: &str,
) -> Result<RosService<S>> {
let requests = self
.node
.create_service::<S>(name, QosProfile::default())
.map_err(map_err)?;
Ok(RosService {
requests: Box::pin(requests),
})
}
pub fn client<S: r2r::WrappedServiceTypeSupport + 'static>(
&mut self,
name: &str,
) -> Result<RosClient<S>> {
let client = self
.node
.create_client::<S>(name, QosProfile::default())
.map_err(map_err)?;
Ok(RosClient { client })
}
pub fn action_client<T: r2r::WrappedActionTypeSupport + 'static>(
&mut self,
name: &str,
) -> Result<RosActionClient<T>> {
let client = self.node.create_action_client::<T>(name).map_err(map_err)?;
Ok(RosActionClient { client })
}
pub fn action_server<T: r2r::WrappedActionTypeSupport + Send + 'static>(
&mut self,
name: &str,
) -> Result<RosActionServer<T>> {
let goals = self.node.create_action_server::<T>(name).map_err(map_err)?;
Ok(RosActionServer {
goals: Box::pin(goals),
})
}
pub fn spin_once(&mut self, timeout: Duration) {
self.node.spin_once(timeout);
}
}
pub struct RosPublisher<T: WrappedTypesupport> {
publisher: Publisher<T>,
}
impl<T: WrappedTypesupport + 'static> Actuator for RosPublisher<T> {
type Command = T;
async fn apply(&mut self, command: T) -> Result<()> {
self.publisher.publish(&command).map_err(map_err)
}
}
pub struct RosSubscriber<T> {
stream: Pin<Box<dyn Stream<Item = T> + Send>>,
}
impl<T> Sensor for RosSubscriber<T> {
type Reading = T;
async fn read(&mut self) -> Result<T> {
self.stream.next().await.ok_or(Error::Closed)
}
}
pub struct RosService<S>
where
S: r2r::WrappedServiceTypeSupport,
{
requests: Pin<Box<dyn Stream<Item = r2r::ServiceRequest<S>> + Send>>,
}
impl<S: r2r::WrappedServiceTypeSupport + 'static> RosService<S> {
pub async fn next_request(&mut self) -> Option<r2r::ServiceRequest<S>> {
self.requests.next().await
}
}
pub struct RosClient<S>
where
S: r2r::WrappedServiceTypeSupport,
{
client: r2r::Client<S>,
}
impl<S: r2r::WrappedServiceTypeSupport + 'static> RosClient<S> {
pub async fn ready(&self) -> Result<()> {
r2r::Node::is_available(&self.client)
.map_err(map_err)?
.await
.map_err(map_err)
}
pub async fn call(&self, request: &S::Request) -> Result<S::Response> {
self.client
.request(request)
.map_err(map_err)?
.await
.map_err(map_err)
}
}
pub struct RosActionClient<T>
where
T: r2r::WrappedActionTypeSupport,
{
client: r2r::ActionClient<T>,
}
impl<T: r2r::WrappedActionTypeSupport + 'static> RosActionClient<T> {
pub async fn ready(&self) -> Result<()> {
r2r::Node::is_available(&self.client)
.map_err(map_err)?
.await
.map_err(map_err)
}
pub async fn send_goal(&self, goal: T::Goal) -> Result<RosGoal<T>>
where
T::Result: Send + 'static,
T::Feedback: Send + 'static,
{
let (_handle, result, feedback) = self
.client
.send_goal_request(goal)
.map_err(map_err)?
.await
.map_err(map_err)?;
let result = Box::pin(async move {
let (_status, value) = result.await.map_err(map_err)?;
Ok(value)
});
Ok(RosGoal {
result,
feedback: Box::pin(feedback),
})
}
}
pub struct RosGoal<T>
where
T: r2r::WrappedActionTypeSupport,
{
result: Pin<Box<dyn Future<Output = Result<T::Result>> + Send>>,
feedback: Pin<Box<dyn Stream<Item = T::Feedback> + Send>>,
}
impl<T: r2r::WrappedActionTypeSupport> RosGoal<T> {
pub async fn next_feedback(&mut self) -> Option<T::Feedback> {
self.feedback.next().await
}
pub async fn result(self) -> Result<T::Result> {
self.result.await
}
}
pub struct RosActionServer<T>
where
T: r2r::WrappedActionTypeSupport,
{
goals: Pin<Box<dyn Stream<Item = r2r::ActionServerGoalRequest<T>> + Send>>,
}
impl<T: r2r::WrappedActionTypeSupport + 'static> RosActionServer<T> {
pub async fn next_goal(&mut self) -> Option<r2r::ActionServerGoalRequest<T>> {
self.goals.next().await
}
}
fn map_err<E: core::fmt::Display>(err: E) -> Error {
Error::Transport(err.to_string())
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn chatter_round_trips_through_ros2() {
let mut node = Ros2Node::new("pamoja_bridge_test", "").unwrap();
let mut publisher = node
.publisher::<r2r::std_msgs::msg::String>("/pamoja_chatter")
.unwrap();
let mut subscriber = node
.subscriber::<r2r::std_msgs::msg::String>("/pamoja_chatter")
.unwrap();
let spinner = std::thread::spawn(move || {
for _ in 0..400 {
node.spin_once(Duration::from_millis(50));
}
});
let received = tokio::time::timeout(Duration::from_secs(15), async {
loop {
publisher
.apply(r2r::std_msgs::msg::String {
data: "hello".to_string(),
})
.await
.unwrap();
tokio::select! {
msg = subscriber.read() => return msg.unwrap(),
_ = tokio::time::sleep(Duration::from_millis(200)) => {}
}
}
})
.await
.expect("a message should arrive within the timeout");
assert_eq!(received.data, "hello");
let _ = spinner.join();
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
#[ignore = "needs rmw_zenoh; run via `cargo xtask ros` or RMW_IMPLEMENTATION=rmw_zenoh_cpp"]
async fn ros2_twist_is_received_over_zenoh() {
use crate::msg::Twist;
use pamoja_core::Transport;
use pamoja_zenoh::{ZenohConfig, ZenohTransport};
use std::time::Duration;
let mut zenoh = ZenohTransport::new(ZenohConfig::new().multicast_scouting(true));
zenoh.connect().await.unwrap();
zenoh.subscribe("0/cmd_vel/**").await.unwrap();
let mut node = Ros2Node::new("pamoja_interop_test", "").unwrap();
let mut publisher = node
.publisher::<r2r::geometry_msgs::msg::Twist>("/cmd_vel")
.unwrap();
let spinner = std::thread::spawn(move || {
for _ in 0..400 {
node.spin_once(Duration::from_millis(50));
}
});
let sample = tokio::time::timeout(Duration::from_secs(20), async {
loop {
publisher.apply(twist(0.6, 0.4)).await.unwrap();
tokio::select! {
msg = zenoh.recv() => return msg.unwrap().unwrap(),
_ = tokio::time::sleep(Duration::from_millis(250)) => {}
}
}
})
.await
.expect("a ROS 2 publication should arrive over Zenoh");
assert!(
sample
.key
.starts_with("0/cmd_vel/geometry_msgs::msg::dds_::Twist_/RIHS01_"),
"unexpected key: {}",
sample.key,
);
let decoded = Twist::from_cdr(&sample.payload).expect("payload should be a CDR Twist");
assert!((decoded.linear.x - 0.6).abs() < 1e-9);
assert!((decoded.angular.z - 0.4).abs() < 1e-9);
let _ = spinner.join();
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn set_bool_service_round_trips() {
use r2r::std_srvs::srv::SetBool;
use std::time::Duration;
let mut node = Ros2Node::new("pamoja_service_test", "").unwrap();
let mut service = node
.service::<SetBool::Service>("/pamoja_set_bool")
.unwrap();
let client = node.client::<SetBool::Service>("/pamoja_set_bool").unwrap();
let spinner = std::thread::spawn(move || {
for _ in 0..400 {
node.spin_once(Duration::from_millis(50));
}
});
let server = tokio::spawn(async move {
if let Some(request) = service.next_request().await {
let response = SetBool::Response {
success: request.message.data,
message: "ok".to_string(),
};
let _ = request.respond(response);
}
});
tokio::time::timeout(Duration::from_secs(10), client.ready())
.await
.expect("the service should become available")
.unwrap();
let response = tokio::time::timeout(
Duration::from_secs(10),
client.call(&SetBool::Request { data: true }),
)
.await
.expect("the call should return")
.unwrap();
assert!(response.success);
assert_eq!(response.message, "ok");
let _ = server.await;
let _ = spinner.join();
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn fibonacci_action_round_trips() {
use r2r::example_interfaces::action::Fibonacci;
use std::time::Duration;
let mut node = Ros2Node::new("pamoja_action_test", "").unwrap();
let mut server = node
.action_server::<Fibonacci::Action>("/pamoja_fib")
.unwrap();
let client = node
.action_client::<Fibonacci::Action>("/pamoja_fib")
.unwrap();
let spinner = std::thread::spawn(move || {
for _ in 0..400 {
node.spin_once(Duration::from_millis(50));
}
});
let server_task = tokio::spawn(async move {
if let Some(request) = server.next_goal().await {
if let Ok((mut goal, _cancel)) = request.accept() {
let _ = goal.publish_feedback(Fibonacci::Feedback {
sequence: vec![0, 1],
});
let _ = goal.succeed(Fibonacci::Result {
sequence: vec![0, 1, 1, 2, 3, 5],
});
}
}
});
tokio::time::timeout(Duration::from_secs(10), client.ready())
.await
.expect("the action server should become available")
.unwrap();
let goal = tokio::time::timeout(
Duration::from_secs(10),
client.send_goal(Fibonacci::Goal { order: 5 }),
)
.await
.expect("the goal should be accepted")
.unwrap();
let result = tokio::time::timeout(Duration::from_secs(10), goal.result())
.await
.expect("the result should arrive")
.unwrap();
assert_eq!(result.sequence, vec![0, 1, 1, 2, 3, 5]);
let _ = server_task.await;
let _ = spinner.join();
}
fn twist(vx: f64, wz: f64) -> r2r::geometry_msgs::msg::Twist {
r2r::geometry_msgs::msg::Twist {
linear: r2r::geometry_msgs::msg::Vector3 {
x: vx,
y: 0.0,
z: 0.0,
},
angular: r2r::geometry_msgs::msg::Vector3 {
x: 0.0,
y: 0.0,
z: wz,
},
}
}
}