use std::{
collections::HashMap,
sync::{Arc, Mutex},
};
use rho_sdk::tool::{ToolProgress, ToolProgressSender};
use rmcp::model::{ProgressNotificationParam, ProgressToken};
#[derive(Clone, Debug, Default)]
pub(crate) struct McpProgressRouter {
subscribers: Arc<Mutex<HashMap<ProgressToken, ToolProgressSender>>>,
}
impl McpProgressRouter {
pub(crate) fn new() -> Self {
Self::default()
}
pub(crate) fn subscribe(
&self,
token: ProgressToken,
sender: ToolProgressSender,
) -> ProgressSubscription {
self.lock().insert(token.clone(), sender);
ProgressSubscription {
token,
router: self.clone(),
}
}
pub(crate) async fn dispatch(&self, params: ProgressNotificationParam) {
let Some(sender) = self.lock().get(¶ms.progress_token).cloned() else {
return;
};
sender.send(tool_progress(params)).await;
}
fn unsubscribe(&self, token: &ProgressToken) {
self.lock().remove(token);
}
fn lock(&self) -> std::sync::MutexGuard<'_, HashMap<ProgressToken, ToolProgressSender>> {
self.subscribers
.lock()
.unwrap_or_else(|error| error.into_inner())
}
}
pub(crate) struct ProgressSubscription {
token: ProgressToken,
router: McpProgressRouter,
}
impl Drop for ProgressSubscription {
fn drop(&mut self) {
self.router.unsubscribe(&self.token);
}
}
fn tool_progress(params: ProgressNotificationParam) -> ToolProgress {
let message = params
.message
.unwrap_or_else(|| "MCP server reported progress".into());
let progress = ToolProgress::message(message);
match params.total {
Some(total) if total.is_finite() && total > 0.0 => {
progress.units(whole_units(params.progress), whole_units(total))
}
_ => progress,
}
}
fn whole_units(value: f64) -> u64 {
if !value.is_finite() || value <= 0.0 {
return 0;
}
value.round() as u64
}
#[cfg(test)]
#[path = "progress_tests.rs"]
mod tests;