From 7ce0e69108c74cbdd4e2a795e9071dac8e774878 Mon Sep 17 00:00:00 2001 From: Huangshuo Kuang <141250392+kkkhs@users.noreply.github.com> Date: Wed, 7 Oct 2026 11:25:49 +0000 Subject: [PATCH 1/2] fix(http): cancel local session requests on disconnect --- .../streamable_http_server/session/local.rs | 112 ++++++++++++++++- .../test_streamable_http_disconnect_cancel.rs | 118 +++++++++++++++++- 2 files changed, 220 insertions(+), 10 deletions(-) diff --git a/crates/rmcp/src/transport/streamable_http_server/session/local.rs b/crates/rmcp/src/transport/streamable_http_server/session/local.rs index e03e5b736..a49f37c9a 100644 --- a/crates/rmcp/src/transport/streamable_http_server/session/local.rs +++ b/crates/rmcp/src/transport/streamable_http_server/session/local.rs @@ -1,11 +1,14 @@ use std::{ collections::{HashMap, HashSet, VecDeque}, num::ParseIntError, + pin::Pin, sync::Arc, + task::{Context, Poll}, time::{Duration, Instant}, }; use futures::{Stream, StreamExt}; +use pin_project_lite::pin_project; use thiserror::Error; use tokio::sync::{ mpsc::{Receiver, Sender}, @@ -17,9 +20,10 @@ use tracing::instrument; use crate::{ RoleServer, model::{ - CancelledNotificationParam, ClientJsonRpcMessage, ClientNotification, ClientRequest, - JsonRpcNotification, JsonRpcRequest, Notification, ProgressNotificationParam, - ProgressToken, RequestId, ServerJsonRpcMessage, ServerNotification, + CancelledNotification, CancelledNotificationParam, ClientJsonRpcMessage, + ClientNotification, ClientRequest, JsonRpcNotification, JsonRpcRequest, Notification, + ProgressNotificationParam, ProgressToken, RequestId, ServerJsonRpcMessage, + ServerNotification, }, transport::{ WorkerTransport, @@ -109,7 +113,11 @@ impl SessionManager for LocalSessionManager { let receiver = handle.establish_request_wise_channel().await?; let http_request_id = receiver.http_request_id; handle.push_message(message, http_request_id).await?; - Ok(ReceiverStream::new(receiver.inner)) + Ok(RequestWiseResponseStream::new( + ReceiverStream::new(receiver.inner), + handle.clone(), + http_request_id, + )) } async fn create_standalone_stream( @@ -426,6 +434,53 @@ pub struct StreamableHttpMessageReceiver { pub inner: Receiver, } +pin_project! { + /// Cancels the in-flight session request when the request-wise response + /// stream is dropped before it reaches natural completion. + struct RequestWiseResponseStream { + #[pin] + inner: ReceiverStream, + handle: LocalSessionHandle, + http_request_id: Option, + } + + impl PinnedDrop for RequestWiseResponseStream { + fn drop(this: Pin<&mut Self>) { + let this = this.project(); + if let Some(id) = this.http_request_id.take() { + this.handle.cancel_request_wise_channel_on_disconnect(id); + } + } + } +} + +impl RequestWiseResponseStream { + fn new( + inner: ReceiverStream, + handle: LocalSessionHandle, + http_request_id: Option, + ) -> Self { + Self { + inner, + handle, + http_request_id, + } + } +} + +impl Stream for RequestWiseResponseStream { + type Item = ServerSseMessage; + + fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + let this = self.project(); + let polled = this.inner.poll_next(cx); + if let Poll::Ready(None) = &polled { + *this.http_request_id = None; + } + polled + } +} + impl LocalSessionWorker { fn unregister_resource(&mut self, resource: &ResourceKey) { let Some(http_request_id) = self.resource_router.remove(resource) else { @@ -486,6 +541,22 @@ impl LocalSessionWorker { self.unregister_resource(&resource); } } + fn remove_request_wise_channel(&mut self, id: HttpRequestId) -> Vec { + let Some(channel) = self.tx_router.remove(&id) else { + return Vec::new(); + }; + channel + .resources + .into_iter() + .filter_map(|resource| { + self.resource_router.remove(&resource); + match resource { + ResourceKey::McpRequestId(request_id) => Some(request_id), + ResourceKey::ProgressToken(_) => None, + } + }) + .collect() + } fn evict_expired_channels(&mut self) { let ttl = self.session_config.completed_cache_ttl; self.tx_router @@ -813,6 +884,9 @@ pub enum SessionEvent { id: HttpRequestId, responder: oneshot::Sender>, }, + CancelRequestWiseChannel { + id: HttpRequestId, + }, Resume { last_event_id: EventId, responder: oneshot::Sender>, @@ -914,6 +988,19 @@ impl LocalSessionHandle { .map_err(|_| SessionError::SessionServiceTerminated)? } + fn cancel_request_wise_channel_on_disconnect(&self, request_id: HttpRequestId) { + let event = SessionEvent::CancelRequestWiseChannel { id: request_id }; + match self.event_tx.try_send(event) { + Ok(()) | Err(tokio::sync::mpsc::error::TrySendError::Closed(_)) => {} + Err(tokio::sync::mpsc::error::TrySendError::Full(event)) => { + let event_tx = self.event_tx.clone(); + tokio::spawn(async move { + let _ = event_tx.send(event).await; + }); + } + } + } + /// Establish a common channel for general purpose messages. pub async fn establish_common_channel( &self, @@ -1186,9 +1273,24 @@ impl Worker for LocalSessionWorker { id, responder, }) => { - let _handle_result = self.tx_router.remove(&id); + self.remove_request_wise_channel(id); let _ = responder.send(Ok(())); } + InnerEvent::FromHttpService(SessionEvent::CancelRequestWiseChannel { id }) => { + let request_ids = self.remove_request_wise_channel(id); + for request_id in request_ids { + context + .send_to_handler(ClientJsonRpcMessage::notification( + ClientNotification::CancelledNotification( + CancelledNotification::new(CancelledNotificationParam::new( + Some(request_id), + Some("client disconnected".to_owned()), + )), + ), + )) + .await?; + } + } InnerEvent::FromHttpService(SessionEvent::Resume { last_event_id, responder, diff --git a/crates/rmcp/tests/test_streamable_http_disconnect_cancel.rs b/crates/rmcp/tests/test_streamable_http_disconnect_cancel.rs index 05efe2110..e2f5edacb 100644 --- a/crates/rmcp/tests/test_streamable_http_disconnect_cancel.rs +++ b/crates/rmcp/tests/test_streamable_http_disconnect_cancel.rs @@ -5,13 +5,9 @@ not(feature = "local") ))] -//! Regression test for #857: when a stateless streamable-HTTP client disconnects +//! Regression tests for #857 and #1325: when a streamable-HTTP client disconnects //! (drops the response) while a tool handler is still awaiting, the per-request //! `RequestContext::ct` should fire so the handler can cancel cooperatively. -//! -//! Stateless requests are one-shot (no session, no resumption), so a dropped -//! response is terminal and safe to cancel — unlike the stateful/resumable path, -//! where a disconnect may be recovered via `Last-Event-ID`. use std::{sync::Arc, time::Duration}; @@ -114,6 +110,47 @@ async fn spawn_stateless_server(json_response: bool) -> anyhow::Result anyhow::Result { + let started = Arc::new(Notify::new()); + let cancelled = Arc::new(Notify::new()); + let probe = CancelProbe { + started: started.clone(), + cancelled: cancelled.clone(), + }; + + let server_ct = CancellationToken::new(); + let config = StreamableHttpServerConfig::default() + // A short keep-alive lets the SSE server notice a dropped connection + // quickly (hyper only observes the disconnect on its next write). + .with_sse_keep_alive(Some(Duration::from_millis(100))) + .with_cancellation_token(server_ct.child_token()); + + let service: StreamableHttpService = + StreamableHttpService::new( + move || Ok(probe.clone()), + Arc::new(LocalSessionManager::default()), + config, + ); + let router = axum::Router::new().nest_service("/mcp", service); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await?; + let addr = listener.local_addr()?; + tokio::spawn({ + let ct = server_ct.clone(); + async move { + let _ = axum::serve(listener, router) + .with_graceful_shutdown(async move { ct.cancelled_owned().await }) + .await; + } + }); + + Ok(TestServer { + url: format!("http://{addr}/mcp"), + server_ct, + started, + cancelled, + }) +} + /// SSE mode: the response is a stream; dropping it (client disconnect) must fire /// the handler's cancellation token. #[tokio::test] @@ -154,6 +191,77 @@ async fn stateless_sse_client_disconnect_cancels_request() -> anyhow::Result<()> Ok(()) } +/// Session mode: dropping the request-wise response stream must cancel the +/// in-flight request owned by the local session worker. +#[tokio::test] +async fn stateful_sse_client_disconnect_cancels_request() -> anyhow::Result<()> { + let server = spawn_session_server().await?; + let client = reqwest::Client::builder() + .pool_max_idle_per_host(0) + .build()?; + + let init = client + .post(&server.url) + .header("Content-Type", "application/json") + .header("Accept", "application/json, text/event-stream") + .header("MCP-Protocol-Version", "2025-06-18") + .body(r#"{"jsonrpc":"2.0","id":0,"method":"initialize","params":{"protocolVersion":"2025-06-18","capabilities":{},"clientInfo":{"name":"test","version":"0.1.0"}}}"#) + .send() + .await?; + assert!( + init.status().is_success(), + "initialize failed: {:?}", + init.status() + ); + let session_id = init + .headers() + .get("mcp-session-id") + .expect("initialize response should include session id") + .to_str()? + .to_owned(); + let _ = init.text().await?; + + let initialized = client + .post(&server.url) + .header("Content-Type", "application/json") + .header("Accept", "application/json, text/event-stream") + .header("mcp-session-id", &session_id) + .header("MCP-Protocol-Version", "2025-06-18") + .body(r#"{"jsonrpc":"2.0","method":"notifications/initialized"}"#) + .send() + .await?; + assert_eq!(initialized.status(), reqwest::StatusCode::ACCEPTED); + + let call = client + .post(&server.url) + .header("Content-Type", "application/json") + .header("Accept", "application/json, text/event-stream") + .header("mcp-session-id", &session_id) + .header("MCP-Protocol-Version", "2025-06-18") + .body(CALL_BODY) + .send() + .await?; + assert!( + call.status().is_success(), + "tools/call failed: {:?}", + call.status() + ); + + tokio::time::timeout(Duration::from_secs(5), server.started.notified()) + .await + .expect("tool handler should start"); + + drop(call); + drop(client); + + tokio::time::timeout(Duration::from_secs(10), server.cancelled.notified()) + .await + .expect("RequestContext::ct should fire after client disconnect (stateful SSE)"); + + server.server_ct.cancel(); + Ok(()) +} + /// JSON-direct mode: the server holds the connection open awaiting the single /// response. A client that disconnects while the handler is running must still /// fire the handler's cancellation token. From e0dc4bae47db0ee40715d2442c0401e240c4d998 Mon Sep 17 00:00:00 2001 From: Huangshuo Kuang <141250392+kkkhs@users.noreply.github.com> Date: Sat, 10 Oct 2026 07:17:10 +0000 Subject: [PATCH 2/2] fix(http): preserve SessionEvent discriminants --- .../src/transport/streamable_http_server/session/local.rs | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/crates/rmcp/src/transport/streamable_http_server/session/local.rs b/crates/rmcp/src/transport/streamable_http_server/session/local.rs index a49f37c9a..5e6f19006 100644 --- a/crates/rmcp/src/transport/streamable_http_server/session/local.rs +++ b/crates/rmcp/src/transport/streamable_http_server/session/local.rs @@ -884,9 +884,6 @@ pub enum SessionEvent { id: HttpRequestId, responder: oneshot::Sender>, }, - CancelRequestWiseChannel { - id: HttpRequestId, - }, Resume { last_event_id: EventId, responder: oneshot::Sender>, @@ -906,6 +903,9 @@ pub enum SessionEvent { EstablishCommonChannel { responder: oneshot::Sender>, }, + CancelRequestWiseChannel { + id: HttpRequestId, + }, } #[derive(Debug, Clone)]