From 033df847a3863a4edfa310a093b152075df51fc8 Mon Sep 17 00:00:00 2001 From: Charlie Marsh Date: Sun, 13 Sep 2026 19:56:24 -0400 Subject: [PATCH] fix(http): cancel session work when the transport closes --- .../src/transport/streamable_http_client.rs | 97 ++++--- ...est_streamable_http_client_cancellation.rs | 262 ++++++++++++++++++ 2 files changed, 316 insertions(+), 43 deletions(-) create mode 100644 crates/rmcp/tests/test_streamable_http_client_cancellation.rs diff --git a/crates/rmcp/src/transport/streamable_http_client.rs b/crates/rmcp/src/transport/streamable_http_client.rs index 8af045fcf..58b0cd352 100644 --- a/crates/rmcp/src/transport/streamable_http_client.rs +++ b/crates/rmcp/src/transport/streamable_http_client.rs @@ -1042,46 +1042,18 @@ impl StreamableHttpClientWorker { Ok((new_session_id, negotiated_version, new_protocol_headers)) } -} -impl Worker for StreamableHttpClientWorker { - type Role = RoleClient; - type Error = StreamableHttpError; - fn is_control_message(message: &ClientJsonRpcMessage) -> bool { - match message { - ClientJsonRpcMessage::Response(_) | ClientJsonRpcMessage::Error(_) => true, - ClientJsonRpcMessage::Notification(notification) => matches!( - notification.notification, - ClientNotification::CancelledNotification(_) - ), - ClientJsonRpcMessage::Request(_) => false, - } - } - fn supports_request_cancellation() -> bool { - true - } - fn err_closed() -> Self::Error { - StreamableHttpError::TransportChannelClosed - } - fn err_join(e: tokio::task::JoinError) -> Self::Error { - StreamableHttpError::TokioJoinError(e) - } - fn config(&self) -> super::worker::WorkerConfig { - super::worker::WorkerConfig { - name: Some("StreamableHttpClientWorker".into()), - channel_buffer_capacity: self.config.channel_buffer_capacity, - } - } - async fn run( + /// Runs initialization and message processing, leaving session cleanup to the owner. + async fn run_session( self, mut context: super::worker::WorkerContext, - ) -> Result<(), WorkerQuitReason> { + session_cleanup_info: &mut Option>, + ) -> Result<(), WorkerQuitReason>> { let channel_buffer_capacity = self.config.channel_buffer_capacity; let (sse_worker_tx, mut sse_worker_rx) = tokio::sync::mpsc::channel::(channel_buffer_capacity); let config = self.config.clone(); let transport_task_ct = context.cancellation_token.clone(); - let _drop_guard = transport_task_ct.clone().drop_guard(); let WorkerSendRequest { responder, message: startup_request, @@ -1155,7 +1127,7 @@ impl Worker for StreamableHttpClientWorker { let mut tool_header_cache: HashMap> = HashMap::new(); // Store session info for cleanup when run() exits (not spawned, so cleanup completes before close() returns) - let mut session_cleanup_info = session_id.as_ref().map(|sid| SessionCleanupInfo { + *session_cleanup_info = session_id.as_ref().map(|sid| SessionCleanupInfo { client: self.client.clone(), uri: config.uri.clone(), session_id: sid.clone(), @@ -1234,7 +1206,7 @@ impl Worker for StreamableHttpClientWorker { } // Each POST uses the session and headers chosen when it starts. // Only this loop updates the current session and protocol version. - let loop_result: Result<(), WorkerQuitReason> = 'main_loop: loop { + 'main_loop: loop { if retrying_recovery && recovery_posts.is_empty() && posts.is_empty() { retrying_recovery = false; } @@ -1283,7 +1255,7 @@ impl Worker for StreamableHttpClientWorker { session_id = new_session_id; negotiated_version = new_version; protocol_headers = new_headers; - session_cleanup_info = session_id.as_ref().map(|sid| SessionCleanupInfo { + *session_cleanup_info = session_id.as_ref().map(|sid| SessionCleanupInfo { client: self.client.clone(), uri: config.uri.clone(), session_id: sid.clone(), @@ -1538,7 +1510,7 @@ impl Worker for StreamableHttpClientWorker { &initialize_response, config.custom_headers.clone(), ); - session_cleanup_info = + *session_cleanup_info = session_id.as_ref().map(|session_id| SessionCleanupInfo { client: self.client.clone(), uri: config.uri.clone(), @@ -1571,7 +1543,7 @@ impl Worker for StreamableHttpClientWorker { protocol_headers .insert(HeaderName::from_static("mcp-protocol-version"), value); } - if let Some(cleanup) = &mut session_cleanup_info { + if let Some(cleanup) = session_cleanup_info.as_mut() { cleanup.protocol_headers = protocol_headers.clone(); } } @@ -1787,15 +1759,54 @@ impl Worker for StreamableHttpClientWorker { } } } + } + } +} + +impl Worker for StreamableHttpClientWorker { + type Role = RoleClient; + type Error = StreamableHttpError; + fn is_control_message(message: &ClientJsonRpcMessage) -> bool { + match message { + ClientJsonRpcMessage::Response(_) | ClientJsonRpcMessage::Error(_) => true, + ClientJsonRpcMessage::Notification(notification) => matches!( + notification.notification, + ClientNotification::CancelledNotification(_) + ), + ClientJsonRpcMessage::Request(_) => false, + } + } + fn supports_request_cancellation() -> bool { + true + } + fn err_closed() -> Self::Error { + StreamableHttpError::TransportChannelClosed + } + fn err_join(e: tokio::task::JoinError) -> Self::Error { + StreamableHttpError::TokioJoinError(e) + } + fn config(&self) -> super::worker::WorkerConfig { + super::worker::WorkerConfig { + name: Some("StreamableHttpClientWorker".into()), + channel_buffer_capacity: self.config.channel_buffer_capacity, + } + } + async fn run( + self, + context: super::worker::WorkerContext, + ) -> Result<(), WorkerQuitReason> { + let transport_task_ct = context.cancellation_token.clone(); + let _drop_guard = transport_task_ct.clone().drop_guard(); + let mut session_cleanup_info = None; + let loop_result = tokio::select! { + biased; + _ = transport_task_ct.cancelled() => Err(WorkerQuitReason::Cancelled), + result = self.run_session(context, &mut session_cleanup_info) => result, }; - // Stop outstanding http requests before deleting their session. + // Dropping the session future releases pending POSTs and aborts its stream tasks. + // DELETE must remain outside cancellation so close() can await bounded cleanup. transport_task_ct.cancel(); - drop(posts); - drop(control_posts); - drop(pending_message); - drop(recovery_posts); - streams.abort_all(); // Cleanup session before returning (ensures close() waits for session deletion) // Use a timeout to prevent indefinite hangs if the server is unresponsive diff --git a/crates/rmcp/tests/test_streamable_http_client_cancellation.rs b/crates/rmcp/tests/test_streamable_http_client_cancellation.rs new file mode 100644 index 000000000..21116cd23 --- /dev/null +++ b/crates/rmcp/tests/test_streamable_http_client_cancellation.rs @@ -0,0 +1,262 @@ +//! Transport cancellation drops active HTTP work before bounded session cleanup. +#![cfg(all(feature = "client", feature = "transport-streamable-http-client"))] + +use std::{collections::HashMap, future::pending, io, sync::Arc, time::Duration}; + +use futures::{StreamExt, stream::BoxStream}; +use http::{HeaderName, HeaderValue}; +use rmcp::{ + model::{ + ClientJsonRpcMessage, InitializeResult, ProtocolVersion, RequestId, ServerCapabilities, + ServerJsonRpcMessage, ServerResult, + }, + transport::{ + Transport, + streamable_http_client::{ + StreamableHttpClient, StreamableHttpClientTransport, + StreamableHttpClientTransportConfig, StreamableHttpError, StreamableHttpPostResponse, + }, + }, +}; +use serde_json::json; +use sse_stream::{Error as SseError, Sse}; +use tokio::{ + sync::{mpsc, oneshot}, + time::timeout, +}; +use tokio_util::sync::CancellationToken; + +type PostResult = Result>; +type Post = (ClientJsonRpcMessage, oneshot::Sender); + +#[derive(Debug)] +struct Delete { + session: Arc, + auth: Option, + complete: oneshot::Sender<()>, +} + +#[derive(Clone)] +struct Client { + posts: mpsc::UnboundedSender, + deletes: mpsc::UnboundedSender, + get_started: CancellationToken, + get_dropped: CancellationToken, +} + +impl StreamableHttpClient for Client { + type Error = io::Error; + + async fn post_message( + &self, + _uri: Arc, + message: ClientJsonRpcMessage, + _session: Option>, + _auth: Option, + _headers: HashMap, + ) -> PostResult { + let (reply, response) = oneshot::channel(); + self.posts.send((message, reply)).unwrap(); + response.await.unwrap() + } + + async fn get_stream( + &self, + _uri: Arc, + _session: Option>, + _last_event_id: Option, + _auth: Option, + _headers: HashMap, + ) -> Result>, StreamableHttpError> { + let _dropped = self.get_dropped.clone().drop_guard(); + self.get_started.cancel(); + pending().await + } + + async fn delete_session( + &self, + _uri: Arc, + session: Arc, + auth: Option, + _headers: HashMap, + ) -> Result<(), StreamableHttpError> { + let (complete, response) = oneshot::channel(); + self.deletes + .send(Delete { + session, + auth, + complete, + }) + .unwrap(); + response.await.unwrap(); + Ok(()) + } +} + +struct Harness { + transport: StreamableHttpClientTransport, + posts: mpsc::UnboundedReceiver, + deletes: mpsc::UnboundedReceiver, + get_started: CancellationToken, + get_dropped: CancellationToken, +} + +impl Harness { + fn new() -> Self { + let (posts, incoming) = mpsc::unbounded_channel(); + let (deletes, cleanup) = mpsc::unbounded_channel(); + let client = Client { + posts, + deletes, + get_started: CancellationToken::new(), + get_dropped: CancellationToken::new(), + }; + Self { + transport: StreamableHttpClientTransport::with_client( + client.clone(), + StreamableHttpClientTransportConfig::with_uri("http://scripted/mcp") + .auth_header("test-token"), + ), + posts: incoming, + deletes: cleanup, + get_started: client.get_started, + get_dropped: client.get_dropped, + } + } + + async fn next_post(&mut self) -> Post { + timeout(Duration::from_secs(1), self.posts.recv()) + .await + .unwrap() + .unwrap() + } +} + +fn initialize() -> ClientJsonRpcMessage { + serde_json::from_value(json!({ + "jsonrpc": "2.0", "id": 1, "method": "initialize", + "params": {"protocolVersion": "2025-11-25", "capabilities": {}, + "clientInfo": {"name": "test", "version": "1"}} + })) + .unwrap() +} + +#[rstest::rstest] +#[case::initialize_post(false, false)] +#[case::initialize_sse(false, true)] +#[case::fallback_post(true, false)] +#[case::fallback_sse(true, true)] +#[tokio::test] +async fn dropping_transport_cancels_initialization( + #[case] fallback: bool, + #[case] sse: bool, +) -> anyhow::Result<()> { + let mut harness = Harness::new(); + if fallback { + let discover = tokio::spawn(harness.transport.send(serde_json::from_value(json!({ + "jsonrpc": "2.0", "id": 0, "method": "server/discover" + }))?)); + let (_, reply) = harness.next_post().await; + reply + .send(Ok(StreamableHttpPostResponse::Json( + serde_json::from_value(json!({ + "jsonrpc": "2.0", "id": 0, + "error": {"code": -32601, "message": "Method not found"} + }))?, + None, + ))) + .unwrap(); + discover.await??; + assert!(harness.transport.receive().await.is_some()); + } + + let send = tokio::spawn(harness.transport.send(initialize())); + let (_, mut reply) = harness.next_post().await; + if sse { + let dropped = CancellationToken::new(); + let guard = dropped.clone().drop_guard(); + let stream = futures::stream::once(async move { + let _guard = guard; + pending().await + }) + .boxed(); + reply + .send(Ok(StreamableHttpPostResponse::Sse(stream, None))) + .unwrap(); + send.await??; + drop(harness.transport); + timeout(Duration::from_secs(1), dropped.cancelled()).await?; + } else { + drop(harness.transport); + timeout(Duration::from_secs(1), reply.closed()).await?; + assert!(send.await?.is_err()); + } + Ok(()) +} + +#[rstest::rstest] +#[case::initialized_post(false, false)] +#[case::get_headers(true, false)] +#[case::cleanup_timeout(false, true)] +#[tokio::test(start_paused = true)] +async fn close_preserves_session_cleanup( + #[case] finish_initialized: bool, + #[case] stall_delete: bool, +) -> anyhow::Result<()> { + let mut harness = Harness::new(); + let send = tokio::spawn(harness.transport.send(initialize())); + let (_, reply) = harness.next_post().await; + reply + .send(Ok(StreamableHttpPostResponse::Json( + ServerJsonRpcMessage::response( + ServerResult::InitializeResult( + InitializeResult::new(ServerCapabilities::default()) + .with_protocol_version(ProtocolVersion::V_2025_11_25), + ), + RequestId::Number(1), + ), + Some("test-session".into()), + ))) + .unwrap(); + send.await??; + assert!(harness.transport.receive().await.is_some()); + + let initialized = tokio::spawn(harness.transport.send(serde_json::from_value(json!({ + "jsonrpc": "2.0", "method": "notifications/initialized" + }))?)); + let (_, reply) = harness.next_post().await; + let pending_initialized = if finish_initialized { + reply + .send(Ok(StreamableHttpPostResponse::Accepted)) + .unwrap(); + timeout(Duration::from_secs(1), harness.get_started.cancelled()).await?; + None + } else { + Some(reply) + }; + + let closed = harness.transport.cancel_token(); + let close = tokio::spawn(async move { harness.transport.close().await }); + let mut delete = timeout(Duration::from_secs(1), harness.deletes.recv()) + .await? + .unwrap(); + assert!(closed.is_cancelled()); + assert_eq!(delete.session.as_ref(), "test-session"); + assert_eq!(delete.auth.as_deref(), Some("test-token")); + assert!(!close.is_finished(), "close must wait for DELETE"); + if let Some(mut reply) = pending_initialized { + assert!(initialized.await?.is_err()); + timeout(Duration::from_secs(1), reply.closed()).await?; + } else { + timeout(Duration::from_secs(1), harness.get_dropped.cancelled()).await?; + initialized.await??; + } + if stall_delete { + tokio::time::advance(Duration::from_secs(5)).await; + timeout(Duration::from_secs(1), delete.complete.closed()).await?; + } else { + delete.complete.send(()).unwrap(); + } + timeout(Duration::from_secs(1), close).await???; + Ok(()) +}