use std::{marker::PhantomData, sync::Arc};
use dashmap::DashMap;
use tokio::sync::{mpsc, watch};
use zenoh::Result;
use super::{GoalId, GoalInfo, GoalStatus, Time, ZAction, messages::*};
use crate::{
Builder, entity::TypeInfo, msg::ZMessage, qos::QosProfile, topic_name::qualify_topic_name,
};
pub mod goal_state {
pub struct Active;
pub struct Terminated;
}
pub struct ZActionClientBuilder<'a, A: ZAction> {
pub action_name: String,
pub node: &'a crate::node::ZNode,
pub goal_service_qos: Option<QosProfile>,
pub result_service_qos: Option<QosProfile>,
pub cancel_service_qos: Option<QosProfile>,
pub feedback_topic_qos: Option<QosProfile>,
pub status_topic_qos: Option<QosProfile>,
pub goal_type_info: Option<TypeInfo>,
pub result_type_info: Option<TypeInfo>,
pub feedback_type_info: Option<TypeInfo>,
pub _phantom: std::marker::PhantomData<A>,
}
impl<'a, A: ZAction> ZActionClientBuilder<'a, A> {
pub fn with_goal_service_qos(mut self, qos: QosProfile) -> Self {
self.goal_service_qos = Some(qos);
self
}
pub fn with_result_service_qos(mut self, qos: QosProfile) -> Self {
self.result_service_qos = Some(qos);
self
}
pub fn with_cancel_service_qos(mut self, qos: QosProfile) -> Self {
self.cancel_service_qos = Some(qos);
self
}
pub fn with_feedback_topic_qos(mut self, qos: QosProfile) -> Self {
self.feedback_topic_qos = Some(qos);
self
}
pub fn with_status_topic_qos(mut self, qos: QosProfile) -> Self {
self.status_topic_qos = Some(qos);
self
}
pub fn with_goal_type_info(mut self, info: TypeInfo) -> Self {
self.goal_type_info = Some(info);
self
}
pub fn with_result_type_info(mut self, info: TypeInfo) -> Self {
self.result_type_info = Some(info);
self
}
pub fn with_feedback_type_info(mut self, info: TypeInfo) -> Self {
self.feedback_type_info = Some(info);
self
}
}
impl<'a, A: ZAction> ZActionClientBuilder<'a, A> {
pub fn new(action_name: &str, node: &'a crate::node::ZNode) -> Self {
Self {
action_name: action_name.to_string(),
node,
goal_service_qos: None,
result_service_qos: None,
cancel_service_qos: None,
feedback_topic_qos: None,
status_topic_qos: None,
goal_type_info: None,
result_type_info: None,
feedback_type_info: None,
_phantom: std::marker::PhantomData,
}
}
}
impl<'a, A: ZAction> Builder for ZActionClientBuilder<'a, A> {
type Output = ZActionClient<A>;
fn build(self) -> Result<Self::Output> {
let action_name = self.node.remap_rules.apply(&self.action_name);
if action_name.is_empty() {
return Err(zenoh::Error::from("Action name cannot be empty"));
}
let qualified_action_name = qualify_topic_name(
&action_name,
&self.node.entity.namespace,
&self.node.entity.name,
)?;
tracing::debug!(
"Action name: '{}', namespace: '{}', qualified: '{}'",
action_name,
self.node.entity.namespace,
qualified_action_name
);
let goal_service_name = format!("{}/_action/send_goal", qualified_action_name);
let result_service_name = format!("{}/_action/get_result", qualified_action_name);
let cancel_service_name = format!("{}/_action/cancel_goal", qualified_action_name);
let feedback_topic_name = format!("{}/_action/feedback", qualified_action_name);
let status_topic_name = format!("{}/_action/status", qualified_action_name);
let goal_type_info = Some(self.goal_type_info.unwrap_or_else(A::send_goal_type_info));
let mut goal_client_builder = self
.node
.create_client_impl::<GoalService<A>>(&goal_service_name, goal_type_info);
if let Some(qos) = self.goal_service_qos {
goal_client_builder.entity.qos = qos.to_protocol_qos();
}
let goal_client = goal_client_builder.build()?;
let result_type_info = Some(
self.result_type_info
.unwrap_or_else(A::get_result_type_info),
);
let mut result_client_builder = self
.node
.create_client_impl::<ResultService<A>>(&result_service_name, result_type_info)
.with_querier_timeout(std::time::Duration::MAX);
if let Some(qos) = self.result_service_qos {
result_client_builder.entity.qos = qos.to_protocol_qos();
}
let result_client = result_client_builder.build()?;
tracing::debug!("Created result client for: {}", result_service_name);
let cancel_type_info = Some(A::cancel_goal_type_info());
let mut cancel_client_builder = self
.node
.create_client_impl::<CancelService<A>>(&cancel_service_name, cancel_type_info);
if let Some(qos) = self.cancel_service_qos {
cancel_client_builder.entity.qos = qos.to_protocol_qos();
}
let cancel_client = cancel_client_builder.build()?;
let goal_board = Arc::new(GoalBoard {
active_goals: DashMap::new(),
});
let feedback_type_info = Some(
self.feedback_type_info
.unwrap_or_else(A::feedback_type_info),
);
let mut feedback_sub_builder = self
.node
.create_sub_impl::<FeedbackMessage<A>>(&feedback_topic_name, feedback_type_info);
if let Some(qos) = self.feedback_topic_qos {
feedback_sub_builder.entity.qos = qos.to_protocol_qos();
}
tracing::debug!(
"Creating feedback subscriber with callback for {}",
feedback_topic_name
);
let goal_board_feedback = goal_board.clone();
let feedback_sub =
feedback_sub_builder.build_with_callback(move |msg: FeedbackMessage<A>| {
tracing::trace!("Feedback callback received for goal {:?}", msg.goal_id);
if let Some(channels) = goal_board_feedback.active_goals.get(&msg.goal_id) {
tracing::trace!("Routing feedback to goal {:?}", msg.goal_id);
let _ = channels.feedback_tx.send(msg.feedback);
} else {
tracing::warn!("No active goal found for feedback {:?}", msg.goal_id);
}
})?;
tracing::debug!("Feedback subscriber created successfully");
let status_type_info = Some(A::status_type_info());
let mut status_sub_builder = self
.node
.create_sub_impl::<StatusMessage>(&status_topic_name, status_type_info);
if let Some(qos) = self.status_topic_qos {
status_sub_builder.entity.qos = qos.to_protocol_qos();
}
let goal_board_status = goal_board.clone();
let status_sub = status_sub_builder.build_with_callback(move |msg: StatusMessage| {
tracing::trace!(
"Status callback received with {} statuses",
msg.status_list.len()
);
for status_info in msg.status_list {
if let Some(channels) = goal_board_status
.active_goals
.get(&status_info.goal_info.goal_id)
{
tracing::trace!(
"Routing status {:?} to goal {:?}",
status_info.status,
status_info.goal_info.goal_id
);
let _ = channels.status_tx.send(status_info.status);
} else {
tracing::trace!(
"No active goal found for status {:?}",
status_info.goal_info.goal_id
);
}
}
})?;
Ok(ZActionClient {
action_name: qualified_action_name,
graph: self.node.graph.clone(),
goal_client: Arc::new(goal_client),
result_client: Arc::new(result_client),
cancel_client: Arc::new(cancel_client),
feedback_sub: Arc::new(feedback_sub),
status_sub: Arc::new(status_sub),
goal_board,
})
}
}
pub struct ZActionClient<A: ZAction> {
action_name: String,
graph: Arc<crate::graph::Graph>,
goal_client: Arc<crate::service::ZClient<GoalService<A>>>,
result_client: Arc<crate::service::ZClient<ResultService<A>>>,
cancel_client: Arc<crate::service::ZClient<CancelService<A>>>,
feedback_sub:
Arc<crate::pubsub::ZSub<FeedbackMessage<A>, (), <FeedbackMessage<A> as ZMessage>::Serdes>>,
status_sub: Arc<crate::pubsub::ZSub<StatusMessage, (), <StatusMessage as ZMessage>::Serdes>>,
goal_board: Arc<GoalBoard<A>>,
}
impl<A: ZAction> std::fmt::Debug for ZActionClient<A> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ZActionClient")
.field("goal_client", &self.goal_client)
.finish_non_exhaustive()
}
}
impl<A: ZAction> Clone for ZActionClient<A> {
fn clone(&self) -> Self {
Self {
action_name: self.action_name.clone(),
graph: self.graph.clone(),
goal_client: self.goal_client.clone(),
result_client: self.result_client.clone(),
cancel_client: self.cancel_client.clone(),
feedback_sub: self.feedback_sub.clone(),
status_sub: self.status_sub.clone(),
goal_board: self.goal_board.clone(),
}
}
}
impl<A: ZAction> ZActionClient<A> {
pub async fn wait_for_server(&self, timeout: std::time::Duration) -> bool {
self.graph
.wait_for_action_server(self.action_name.as_str(), timeout)
.await
}
pub async fn send_goal(&self, goal: A::Goal) -> Result<GoalHandle<A, goal_state::Active>> {
let goal_id = GoalId::new();
let (feedback_tx, feedback_rx) = mpsc::unbounded_channel();
let (status_tx, status_rx) = watch::channel(GoalStatus::Unknown);
self.goal_board.active_goals.insert(
goal_id,
GoalChannels {
feedback_tx,
status_tx,
},
);
let request = SendGoalRequest { goal_id, goal };
tracing::debug!("Sending goal request for goal_id: {:?}", goal_id);
let response = match self.goal_client.call(&request).await {
Ok(response) => response,
Err(error) => {
self.goal_board.active_goals.remove(&goal_id);
return Err(error);
}
};
if !response.accepted {
self.goal_board.active_goals.remove(&goal_id);
return Err(zenoh::Error::from("Goal rejected".to_string()));
}
if let Some(channels) = self.goal_board.active_goals.get(&goal_id) {
channels.status_tx.send_if_modified(|s| {
if *s == GoalStatus::Unknown {
*s = GoalStatus::Accepted;
true
} else {
false
}
});
}
Ok(GoalHandle {
id: goal_id,
client: Arc::new(self.clone()),
feedback_rx: Some(feedback_rx),
status_rx: Some(status_rx),
_state: PhantomData,
})
}
pub async fn cancel_goal(&self, goal_id: GoalId) -> Result<CancelGoalServiceResponse> {
let goal_info = GoalInfo::new(goal_id);
let request = CancelGoalServiceRequest { goal_info };
self.cancel_client.call(&request).await
}
pub async fn cancel_all_goals(&self) -> Result<CancelGoalServiceResponse> {
let zero_goal_id = GoalId([0u8; 16]);
let goal_info = GoalInfo {
goal_id: zero_goal_id,
stamp: Time::zero(),
};
let request = CancelGoalServiceRequest { goal_info };
self.cancel_client.call(&request).await
}
pub fn feedback_stream(&self, goal_id: GoalId) -> Option<mpsc::UnboundedReceiver<A::Feedback>> {
self.goal_board
.active_goals
.get_mut(&goal_id)
.map(|mut channels| {
let (tx, rx) = mpsc::unbounded_channel();
channels.feedback_tx = tx;
rx
})
}
pub fn status_watch(&self, goal_id: GoalId) -> Option<watch::Receiver<GoalStatus>> {
self.goal_board
.active_goals
.get(&goal_id)
.map(|channels| channels.status_tx.subscribe())
}
pub async fn get_result(&self, goal_id: GoalId) -> Result<A::Result> {
let request = GetResultRequest { goal_id };
let response: GetResultResponse<A> = self.result_client.call(&request).await?;
Ok(response.result)
}
}
struct GoalBoard<A: ZAction> {
active_goals: DashMap<GoalId, GoalChannels<A>>,
}
struct GoalChannels<A: ZAction> {
feedback_tx: mpsc::UnboundedSender<A::Feedback>,
status_tx: watch::Sender<GoalStatus>,
}
pub struct GoalHandle<A: ZAction, State = goal_state::Active> {
id: GoalId,
client: Arc<ZActionClient<A>>,
feedback_rx: Option<mpsc::UnboundedReceiver<A::Feedback>>,
status_rx: Option<watch::Receiver<GoalStatus>>,
_state: PhantomData<State>,
}
impl<A: ZAction> GoalHandle<A, goal_state::Active> {
pub fn id(&self) -> GoalId {
self.id
}
pub fn feedback(&mut self) -> Option<mpsc::UnboundedReceiver<A::Feedback>> {
self.feedback_rx.take()
}
pub fn status_watch(&mut self) -> Option<watch::Receiver<GoalStatus>> {
self.status_rx.take()
}
pub async fn cancel(&self) -> Result<CancelGoalServiceResponse> {
self.client.cancel_goal(self.id).await
}
pub async fn result(self) -> Result<A::Result> {
let res = self.client.get_result(self.id).await;
self.client.goal_board.active_goals.remove(&self.id);
res
}
pub async fn result_with_timeout(self, timeout: std::time::Duration) -> Result<A::Result> {
let res = match tokio::time::timeout(timeout, self.client.get_result(self.id)).await {
Ok(res) => res,
Err(_) => Err(crate::error::Error::timeout(timeout)),
};
self.client.goal_board.active_goals.remove(&self.id);
res
}
}