From 48637fd77664774230571e131215fe957678d506 Mon Sep 17 00:00:00 2001 From: monody0007 <52037177+monody0007@users.noreply.github.com> Date: Mon, 28 Sep 2026 10:08:12 -0700 Subject: [PATCH 1/2] fix(client): keep ProgressDispatcher lock off the await path ProgressDispatcher held its tokio RwLock read guard while awaiting the bounded per-subscriber channel, so one subscriber that fell 16 notifications behind stalled subscribe() and delivery for every other token sharing the dispatcher. ProgressSubscriber::drop also unregistered through tokio::spawn, which panics outside a runtime and could run after the same token had been subscribed again, removing the new subscription. Use a std RwLock that is only held to look up or edit the map, and unregister synchronously on drop, only when the entry still belongs to the dropped subscriber. --- crates/rmcp/src/handler/client/progress.rs | 70 +++++++++--- crates/rmcp/tests/test_progress_subscriber.rs | 101 +++++++++++++++++- 2 files changed, 156 insertions(+), 15 deletions(-) diff --git a/crates/rmcp/src/handler/client/progress.rs b/crates/rmcp/src/handler/client/progress.rs index 04f31610f..4bb5b5f1b 100644 --- a/crates/rmcp/src/handler/client/progress.rs +++ b/crates/rmcp/src/handler/client/progress.rs @@ -1,10 +1,15 @@ -use std::{collections::HashMap, sync::Arc}; +use std::{ + collections::HashMap, + sync::{Arc, PoisonError, RwLock}, +}; use futures::{Stream, StreamExt}; -use tokio::sync::RwLock; use tokio_stream::wrappers::ReceiverStream; use crate::model::{ProgressNotificationParam, ProgressToken}; +// A synchronous lock: it is never held across an `.await`, so a subscriber that is not +// keeping up only back-pressures its own token, and `ProgressSubscriber::drop` can +// unregister without spawning onto a runtime. type Dispatcher = Arc>>>; @@ -22,8 +27,13 @@ impl ProgressDispatcher { /// Handle a progress notification by sending it to the appropriate subscriber pub async fn handle_notification(&self, notification: ProgressNotificationParam) { - let token = ¬ification.progress_token; - if let Some(sender) = self.dispatcher.read().await.get(token).cloned() { + let sender = self + .dispatcher + .read() + .unwrap_or_else(PoisonError::into_inner) + .get(¬ification.progress_token) + .cloned(); + if let Some(sender) = sender { let send_result = sender.send(notification).await; if let Err(e) = send_result { tracing::warn!("Failed to send progress notification: {e}"); @@ -38,7 +48,7 @@ impl ProgressDispatcher { let (sender, receiver) = tokio::sync::mpsc::channel(Self::CHANNEL_SIZE); self.dispatcher .write() - .await + .unwrap_or_else(PoisonError::into_inner) .insert(progress_token.clone(), sender); let receiver = ReceiverStream::new(receiver); ProgressSubscriber { @@ -50,13 +60,18 @@ impl ProgressDispatcher { /// Unsubscribe from progress notifications for a specific token. pub async fn unsubscribe(&self, token: &ProgressToken) { - self.dispatcher.write().await.remove(token); + self.dispatcher + .write() + .unwrap_or_else(PoisonError::into_inner) + .remove(token); } /// Clear all dispatcher. pub async fn clear(&self) { - let mut dispatcher = self.dispatcher.write().await; - dispatcher.clear(); + self.dispatcher + .write() + .unwrap_or_else(PoisonError::into_inner) + .clear(); } } @@ -89,12 +104,39 @@ impl Stream for ProgressSubscriber { impl Drop for ProgressSubscriber { fn drop(&mut self) { - let token = self.progress_token.clone(); self.receiver.close(); - let dispatcher = self.dispatcher.clone(); - tokio::spawn(async move { - let mut dispatcher = dispatcher.write_owned().await; - dispatcher.remove(&token); - }); + // Only remove the entry if it still belongs to this subscriber: the token may + // have been subscribed again since, and that subscription must stay registered. + let mut dispatcher = self + .dispatcher + .write() + .unwrap_or_else(PoisonError::into_inner); + if dispatcher + .get(&self.progress_token) + .is_some_and(|sender| sender.is_closed()) + { + dispatcher.remove(&self.progress_token); + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::model::NumberOrString; + + #[test] + fn dropping_a_subscriber_unregisters_it_without_a_runtime() { + let runtime = tokio::runtime::Builder::new_current_thread() + .build() + .expect("build runtime"); + let dispatcher = ProgressDispatcher::new(); + let token = ProgressToken(NumberOrString::Number(1)); + let subscriber = runtime.block_on(dispatcher.subscribe(token.clone())); + assert!(dispatcher.dispatcher.read().unwrap().contains_key(&token)); + + // Dropped outside of any Tokio runtime context. + drop(subscriber); + assert!(dispatcher.dispatcher.read().unwrap().is_empty()); } } diff --git a/crates/rmcp/tests/test_progress_subscriber.rs b/crates/rmcp/tests/test_progress_subscriber.rs index 8df1b7cad..ac12d33ae 100644 --- a/crates/rmcp/tests/test_progress_subscriber.rs +++ b/crates/rmcp/tests/test_progress_subscriber.rs @@ -1,10 +1,13 @@ #![cfg(not(feature = "local"))] +use std::time::Duration; + use futures::StreamExt; use rmcp::{ ClientHandler, Peer, RoleServer, ServerHandler, ServiceExt, handler::{client::progress::ProgressDispatcher, server::tool::ToolRouter}, model::{ - CallToolRequestParams, ClientRequest, ProgressNotificationParam, Request, RequestMetaObject, + CallToolRequestParams, ClientRequest, NumberOrString, ProgressNotificationParam, + ProgressToken, Request, RequestMetaObject, }, service::PeerRequestOptions, tool, tool_handler, tool_router, @@ -132,3 +135,99 @@ async fn test_progress_subscriber() -> anyhow::Result<()> { tokio::time::sleep(tokio::time::Duration::from_secs(1)).await; Ok(()) } + +fn progress(token: &ProgressToken, step: u32) -> ProgressNotificationParam { + ProgressNotificationParam::new(token.clone(), step as f64) +} + +// A subscriber that is not keeping up must only back-pressure its own token, +// not every other subscription sharing the dispatcher. +#[tokio::test] +async fn test_slow_subscriber_does_not_block_other_tokens() -> anyhow::Result<()> { + const SLOW_NOTIFICATIONS: u32 = 64; + let dispatcher = ProgressDispatcher::new(); + let slow_token = ProgressToken(NumberOrString::Number(1)); + let mut slow = dispatcher.subscribe(slow_token.clone()).await; + + // Deliver more notifications than the slow subscriber can buffer while it is not polled. + let producer = tokio::spawn({ + let dispatcher = dispatcher.clone(); + let slow_token = slow_token.clone(); + async move { + for step in 0..SLOW_NOTIFICATIONS { + dispatcher + .handle_notification(progress(&slow_token, step)) + .await; + } + } + }); + tokio::time::sleep(Duration::from_millis(100)).await; + assert!( + !producer.is_finished(), + "slow subscriber should be applying back-pressure" + ); + + let fast_token = ProgressToken(NumberOrString::Number(2)); + let mut fast = tokio::time::timeout( + Duration::from_secs(1), + dispatcher.subscribe(fast_token.clone()), + ) + .await + .expect("subscribe() stalled behind an unrelated slow subscriber"); + tokio::time::timeout( + Duration::from_secs(1), + dispatcher.handle_notification(progress(&fast_token, 7)), + ) + .await + .expect("delivery to an unrelated token stalled behind a slow subscriber"); + let received = tokio::time::timeout(Duration::from_secs(1), fast.next()) + .await + .expect("fast subscriber did not receive its notification") + .expect("fast subscriber stream ended"); + assert_eq!(received.progress, 7.0); + + // Once the slow subscriber catches up it still gets every notification, in order. + for step in 0..SLOW_NOTIFICATIONS { + let received = tokio::time::timeout(Duration::from_secs(1), slow.next()) + .await + .expect("slow subscriber did not receive a queued notification") + .expect("slow subscriber stream ended"); + assert_eq!(received.progress, step as f64); + } + tokio::time::timeout(Duration::from_secs(1), producer).await??; + Ok(()) +} + +#[tokio::test] +async fn test_resubscribing_a_dropped_token_keeps_the_new_subscriber() { + let dispatcher = ProgressDispatcher::new(); + let token = ProgressToken(NumberOrString::Number(1)); + drop(dispatcher.subscribe(token.clone()).await); + let mut subscriber = dispatcher.subscribe(token.clone()).await; + // Give any deferred cleanup of the dropped subscriber a chance to run. + tokio::time::sleep(Duration::from_millis(50)).await; + + dispatcher.handle_notification(progress(&token, 3)).await; + let received = tokio::time::timeout(Duration::from_secs(1), subscriber.next()) + .await + .expect("new subscriber did not receive its notification") + .expect("dropping the old subscriber unregistered the new one"); + assert_eq!(received.progress, 3.0); +} + +#[tokio::test] +async fn test_dropping_a_replaced_subscriber_keeps_its_replacement() { + let dispatcher = ProgressDispatcher::new(); + let token = ProgressToken(NumberOrString::Number(1)); + let replaced = dispatcher.subscribe(token.clone()).await; + let mut replacement = dispatcher.subscribe(token.clone()).await; + drop(replaced); + tokio::time::sleep(Duration::from_millis(50)).await; + + dispatcher.handle_notification(progress(&token, 5)).await; + let received = tokio::time::timeout(Duration::from_secs(1), replacement.next()) + .await + .expect("replacement subscriber did not receive its notification") + .expect("dropping the replaced subscriber unregistered its replacement"); + assert_eq!(received.progress, 5.0); +} From 57ee24c6724421b65767fa23a7b056f4ff341291 Mon Sep 17 00:00:00 2001 From: monody0007 <52037177+monody0007@users.noreply.github.com> Date: Mon, 5 Oct 2026 13:39:36 -0700 Subject: [PATCH 2/2] fix(client): run no caller code under the progress registry lock ProgressSubscriber::drop takes the registry's std RwLock, so any caller code that runs while the lock is held can re-enter it on the same thread and deadlock: - Dropping the last sender of a progress channel wakes the receiver's task synchronously, and a scheduler may drop a cancelled task (and the ProgressSubscriber it owns) right there. subscribe() (replaced sender), unsubscribe(), clear() and ProgressSubscriber::drop all destroyed senders while still holding the write lock. - The derived Debug of ProgressDispatcher formatted the RwLock, which holds a read guard while it writes the map into the caller's fmt::Write; a writer that drops a subscriber there blocks forever. Move removed senders out of the map under the lock and drop them after the guard is released. Replace the derived Debug with one that copies the subscribed tokens out first and formats them after release. Add unit tests that drop a subscriber from a waker for each sender path and from a writer for Debug. --- crates/rmcp/src/handler/client/progress.rs | 230 +++++++++++++++++++-- 1 file changed, 217 insertions(+), 13 deletions(-) diff --git a/crates/rmcp/src/handler/client/progress.rs b/crates/rmcp/src/handler/client/progress.rs index 4bb5b5f1b..b6678a188 100644 --- a/crates/rmcp/src/handler/client/progress.rs +++ b/crates/rmcp/src/handler/client/progress.rs @@ -9,16 +9,34 @@ use tokio_stream::wrappers::ReceiverStream; use crate::model::{ProgressNotificationParam, ProgressToken}; // A synchronous lock: it is never held across an `.await`, so a subscriber that is not // keeping up only back-pressures its own token, and `ProgressSubscriber::drop` can -// unregister without spawning onto a runtime. +// unregister without spawning onto a runtime. Because that `drop` takes the lock, no caller +// code may run while it is held: senders are dropped (which can wake a task that drops its +// subscriber) and `Debug` output is written (into a caller's writer) only after release. type Dispatcher = Arc>>>; /// A dispatcher for progress notifications. -#[derive(Debug, Clone, Default)] +#[derive(Clone, Default)] pub struct ProgressDispatcher { pub(crate) dispatcher: Dispatcher, } +impl std::fmt::Debug for ProgressDispatcher { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + // Copy the tokens out first: the formatter writes into caller code. + let subscriptions: Vec = self + .dispatcher + .read() + .unwrap_or_else(PoisonError::into_inner) + .keys() + .cloned() + .collect(); + f.debug_struct("ProgressDispatcher") + .field("subscriptions", &subscriptions) + .finish() + } +} + impl ProgressDispatcher { const CHANNEL_SIZE: usize = 16; pub fn new() -> Self { @@ -46,10 +64,12 @@ impl ProgressDispatcher { /// If you drop the returned `ProgressSubscriber`, it will automatically unsubscribe from notifications for that token. pub async fn subscribe(&self, progress_token: ProgressToken) -> ProgressSubscriber { let (sender, receiver) = tokio::sync::mpsc::channel(Self::CHANNEL_SIZE); - self.dispatcher + let replaced = self + .dispatcher .write() .unwrap_or_else(PoisonError::into_inner) .insert(progress_token.clone(), sender); + drop(replaced); let receiver = ReceiverStream::new(receiver); ProgressSubscriber { progress_token, @@ -60,18 +80,23 @@ impl ProgressDispatcher { /// Unsubscribe from progress notifications for a specific token. pub async fn unsubscribe(&self, token: &ProgressToken) { - self.dispatcher + let removed = self + .dispatcher .write() .unwrap_or_else(PoisonError::into_inner) .remove(token); + drop(removed); } /// Clear all dispatcher. pub async fn clear(&self) { - self.dispatcher - .write() - .unwrap_or_else(PoisonError::into_inner) - .clear(); + let removed = std::mem::take( + &mut *self + .dispatcher + .write() + .unwrap_or_else(PoisonError::into_inner), + ); + drop(removed); } } @@ -105,26 +130,102 @@ impl Stream for ProgressSubscriber { impl Drop for ProgressSubscriber { fn drop(&mut self) { self.receiver.close(); - // Only remove the entry if it still belongs to this subscriber: the token may - // have been subscribed again since, and that subscription must stay registered. + // Only remove the entry if its receiver is closed. The token may have been + // subscribed again since; that subscription's receiver is still open, so it + // stays registered. let mut dispatcher = self .dispatcher .write() .unwrap_or_else(PoisonError::into_inner); - if dispatcher + let removed = if dispatcher .get(&self.progress_token) .is_some_and(|sender| sender.is_closed()) { - dispatcher.remove(&self.progress_token); - } + dispatcher.remove(&self.progress_token) + } else { + None + }; + drop(dispatcher); + drop(removed); } } #[cfg(test)] mod tests { + use std::{ + sync::{Mutex, mpsc::RecvTimeoutError}, + task::{Context, Wake, Waker}, + time::Duration, + }; + + use futures::FutureExt; + use super::*; use crate::model::NumberOrString; + /// A task that owns a subscriber and drops it as soon as it is woken, as a + /// scheduler may synchronously do with a cancelled task. + struct DropOnWake(Mutex>); + + impl DropOnWake { + fn new(subscriber: ProgressSubscriber) -> Arc { + Arc::new(Self(Mutex::new(Some(subscriber)))) + } + + /// Polls `subscriber` with this task's waker, so closing its channel wakes us. + fn watch(self: &Arc, subscriber: &mut ProgressSubscriber) { + let waker = Waker::from(self.clone()); + let poll = subscriber.poll_next_unpin(&mut Context::from_waker(&waker)); + assert!(poll.is_pending()); + } + + fn watch_own(self: &Arc) { + let mut subscriber = self.0.lock().unwrap(); + self.watch(subscriber.as_mut().unwrap()); + } + + fn dropped(&self) -> bool { + self.0.lock().unwrap().is_none() + } + } + + impl Wake for DropOnWake { + fn wake(self: Arc) { + self.wake_by_ref(); + } + + fn wake_by_ref(self: &Arc) { + let subscriber = self.0.lock().unwrap().take(); + drop(subscriber); + } + } + + /// Runs `scenario` on its own thread: a deadlock blocks synchronously, so only + /// another thread can notice it. + fn assert_completes(scenario: impl FnOnce() + Send + 'static) { + let (done_tx, done_rx) = std::sync::mpsc::channel(); + let handle = std::thread::spawn(move || { + scenario(); + let _ = done_tx.send(()); + }); + match done_rx.recv_timeout(Duration::from_secs(5)) { + Ok(()) => {} + Err(RecvTimeoutError::Timeout) => { + panic!("deadlocked: a subscriber dropped under the lock could not take it") + } + Err(RecvTimeoutError::Disconnected) => { + std::panic::resume_unwind(handle.join().unwrap_err()) + } + } + } + + fn subscribe(dispatcher: &ProgressDispatcher, token: &ProgressToken) -> ProgressSubscriber { + dispatcher + .subscribe(token.clone()) + .now_or_never() + .expect("subscribe does not wait") + } + #[test] fn dropping_a_subscriber_unregisters_it_without_a_runtime() { let runtime = tokio::runtime::Builder::new_current_thread() @@ -139,4 +240,107 @@ mod tests { drop(subscriber); assert!(dispatcher.dispatcher.read().unwrap().is_empty()); } + + #[test] + fn clear_tolerates_a_subscriber_dropped_on_wake() { + assert_completes(|| { + let dispatcher = ProgressDispatcher::new(); + let token = ProgressToken(NumberOrString::Number(1)); + let task = DropOnWake::new(subscribe(&dispatcher, &token)); + task.watch_own(); + + dispatcher.clear().now_or_never().unwrap(); + assert!(task.dropped()); + assert!(dispatcher.dispatcher.read().unwrap().is_empty()); + }); + } + + #[test] + fn unsubscribe_tolerates_a_subscriber_dropped_on_wake() { + assert_completes(|| { + let dispatcher = ProgressDispatcher::new(); + let token = ProgressToken(NumberOrString::Number(1)); + let task = DropOnWake::new(subscribe(&dispatcher, &token)); + task.watch_own(); + + dispatcher.unsubscribe(&token).now_or_never().unwrap(); + assert!(task.dropped()); + assert!(dispatcher.dispatcher.read().unwrap().is_empty()); + }); + } + + #[test] + fn resubscribe_tolerates_the_replaced_subscriber_dropped_on_wake() { + assert_completes(|| { + let dispatcher = ProgressDispatcher::new(); + let token = ProgressToken(NumberOrString::Number(1)); + let task = DropOnWake::new(subscribe(&dispatcher, &token)); + task.watch_own(); + + let _replacement = subscribe(&dispatcher, &token); + assert!(task.dropped()); + // The replaced subscriber's drop must not unregister the replacement. + let registry = dispatcher.dispatcher.read().unwrap(); + assert!( + registry + .get(&token) + .is_some_and(|sender| !sender.is_closed()) + ); + }); + } + + #[test] + fn subscriber_drop_tolerates_another_subscriber_dropped_on_wake() { + assert_completes(|| { + let dispatcher = ProgressDispatcher::new(); + let first = ProgressToken(NumberOrString::Number(1)); + let second = ProgressToken(NumberOrString::Number(2)); + let mut subscriber = subscribe(&dispatcher, &first); + let task = DropOnWake::new(subscribe(&dispatcher, &second)); + task.watch(&mut subscriber); + + drop(subscriber); + assert!(task.dropped()); + assert!(dispatcher.dispatcher.read().unwrap().is_empty()); + }); + } + + /// A caller's writer that drops a subscriber it owns once the subscribed token is + /// written, which a `Debug` that formats the registry in place does under the lock. + struct DropOnWrite { + subscriber: Option, + output: String, + } + + impl std::fmt::Write for DropOnWrite { + fn write_str(&mut self, text: &str) -> std::fmt::Result { + if text.contains("drop-on-write") { + drop(self.subscriber.take()); + } + self.output.push_str(text); + Ok(()) + } + } + + #[test] + fn debug_tolerates_a_subscriber_dropped_by_the_writer() { + assert_completes(|| { + use std::fmt::Write; + + let dispatcher = ProgressDispatcher::new(); + let token = ProgressToken(NumberOrString::String("drop-on-write".into())); + let mut writer = DropOnWrite { + subscriber: Some(subscribe(&dispatcher, &token)), + output: String::new(), + }; + + write!(writer, "{dispatcher:?}").unwrap(); + assert!(writer.subscriber.is_none()); + assert_eq!( + writer.output, + r#"ProgressDispatcher { subscriptions: [ProgressToken(String("drop-on-write"))] }"# + ); + assert!(dispatcher.dispatcher.read().unwrap().is_empty()); + }); + } }