Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
112 changes: 107 additions & 5 deletions crates/rmcp/src/transport/streamable_http_server/session/local.rs
Original file line number Diff line number Diff line change
@@ -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},
Expand All @@ -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,
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -426,6 +434,53 @@ pub struct StreamableHttpMessageReceiver {
pub inner: Receiver<ServerSseMessage>,
}

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<ServerSseMessage>,
handle: LocalSessionHandle,
http_request_id: Option<HttpRequestId>,
}

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<ServerSseMessage>,
handle: LocalSessionHandle,
http_request_id: Option<HttpRequestId>,
) -> 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<Option<Self::Item>> {
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 {
Expand Down Expand Up @@ -486,6 +541,22 @@ impl LocalSessionWorker {
self.unregister_resource(&resource);
}
}
fn remove_request_wise_channel(&mut self, id: HttpRequestId) -> Vec<RequestId> {
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
Expand Down Expand Up @@ -832,6 +903,9 @@ pub enum SessionEvent {
EstablishCommonChannel {
responder: oneshot::Sender<Result<StreamableHttpMessageReceiver, SessionError>>,
},
CancelRequestWiseChannel {
id: HttpRequestId,
},
}

#[derive(Debug, Clone)]
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand Down
118 changes: 113 additions & 5 deletions crates/rmcp/tests/test_streamable_http_disconnect_cancel.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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};

Expand Down Expand Up @@ -114,6 +110,47 @@ async fn spawn_stateless_server(json_response: bool) -> anyhow::Result<TestServe
})
}

async fn spawn_session_server() -> anyhow::Result<TestServer> {
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<CancelProbe, LocalSessionManager> =
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]
Expand Down Expand Up @@ -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.
Expand Down
Loading