Skip to content

Commit 74c65c4

Browse files
committed
fix(transport): cancel in-flight request on stateless streamable-HTTP client disconnect (#857)
A stateless streamable-HTTP request is one-shot (no session, no resumption), so if the client drops the response before the handler finishes, the request is terminal and should be cancelled. Previously the handler kept running with its `RequestContext::ct` never firing, so long-running or destructive tools could not observe a client disconnect. Give each stateless request its own cancellation token via `serve_directly_with_ct` and cancel it when the client disconnects. This covers both stateless paths: `serve_negotiated_request_directly` (per-request version negotiation) and the non-negotiated path. In each: - A disconnect while the handler is still producing its first message cancels it. The guard is disarmed once the handler emits anything, so a normal response is never cancelled. - The SSE response stream is wrapped in a guard that cancels the handler if the stream is dropped before it ends naturally. Stateful (resumable) mode is intentionally left unchanged: there a disconnect may be recovered via `Last-Event-ID`, so cancelling on disconnect would break resumption. Adds a regression test covering both stateless sub-modes (SSE and JSON).
1 parent 9df629e commit 74c65c4

3 files changed

Lines changed: 310 additions & 10 deletions

File tree

crates/rmcp/Cargo.toml

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -426,3 +426,12 @@ required-features = [
426426
"transport-streamable-http-client-reqwest",
427427
]
428428
path = "tests/test_streamable_http_connection_reuse.rs"
429+
430+
[[test]]
431+
name = "test_streamable_http_disconnect_cancel"
432+
required-features = [
433+
"server",
434+
"transport-streamable-http-server",
435+
"reqwest",
436+
]
437+
path = "tests/test_streamable_http_disconnect_cancel.rs"

crates/rmcp/src/transport/streamable_http_server/tower.rs

Lines changed: 105 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -1,12 +1,20 @@
11
use std::{
2-
borrow::Cow, collections::HashMap, convert::Infallible, fmt::Display, sync::Arc, time::Duration,
2+
borrow::Cow,
3+
collections::HashMap,
4+
convert::Infallible,
5+
fmt::Display,
6+
pin::Pin,
7+
sync::Arc,
8+
task::{Context, Poll},
9+
time::Duration,
310
};
411

512
use bytes::Bytes;
6-
use futures::{StreamExt, future::BoxFuture};
13+
use futures::{Stream, StreamExt, future::BoxFuture};
714
use http::{HeaderMap, Method, Request, Response, header::ALLOW};
815
use http_body::Body;
916
use http_body_util::{BodyExt, Full, combinators::BoxBody};
17+
use pin_project_lite::pin_project;
1018
use tokio_stream::wrappers::ReceiverStream;
1119
use tokio_util::sync::CancellationToken;
1220

@@ -22,7 +30,7 @@ use crate::{
2230
ProtocolVersion, RequestId, ServerJsonRpcMessage,
2331
},
2432
serve_server,
25-
service::serve_directly,
33+
service::serve_directly_with_ct,
2634
transport::{
2735
OneshotTransport, TransportAdapterIdentity,
2836
common::{
@@ -852,14 +860,28 @@ where
852860
request.request.extensions_mut().insert(parts);
853861
let (transport, mut receiver) =
854862
OneshotTransport::<RoleServer>::new(ClientJsonRpcMessage::Request(request));
855-
let service = serve_directly(service, transport, peer_info);
863+
// Give this stateless request its own cancellation token so a client
864+
// disconnect can cancel the in-flight handler (#857), as in the
865+
// non-negotiated stateless path below.
866+
let request_ct = CancellationToken::new();
867+
let service = serve_directly_with_ct(service, transport, peer_info, request_ct.clone());
856868
tokio::spawn(async move {
857869
let _ = service.waiting().await;
858870
});
859871

860872
let cancel = self.config.cancellation_token.child_token();
873+
// Cancel the handler if the client disconnects while it is still
874+
// producing its first message (this future is dropped before
875+
// `receiver.recv()` completes). Disarmed once the handler emits
876+
// anything, so a normal response is never cancelled.
877+
let mut disconnect_guard = Some(request_ct.clone().drop_guard());
861878
let first = tokio::select! {
862-
message = receiver.recv() => message,
879+
message = receiver.recv() => {
880+
if let Some(guard) = disconnect_guard.take() {
881+
guard.disarm();
882+
}
883+
message
884+
}
863885
_ = cancel.cancelled() => None,
864886
}
865887
.ok_or_else(|| {
@@ -873,14 +895,16 @@ where
873895
return jsonrpc_message_response(first, true);
874896
}
875897

898+
// The handler may still be streaming, so guard the response: dropping it
899+
// (client disconnect) must cancel the handler.
876900
let stream = futures::stream::once(async move { first })
877901
.chain(ReceiverStream::new(receiver))
878902
.map(|message| {
879903
tracing::trace!(?message);
880904
ServerSseMessage::from_message(message)
881905
});
882906
Ok(sse_stream_response(
883-
stream,
907+
CancelOnDisconnect::new(stream, request_ct),
884908
self.config.sse_keep_alive,
885909
self.config.cancellation_token.child_token(),
886910
))
@@ -1544,7 +1568,13 @@ where
15441568
request.request.extensions_mut().insert(part);
15451569
let (transport, mut receiver) =
15461570
OneshotTransport::<RoleServer>::new(ClientJsonRpcMessage::Request(request));
1547-
let service = serve_directly(service, transport, peer_info);
1571+
// Give this stateless request its own cancellation token so a
1572+
// client disconnect can cancel the in-flight handler (#857). A
1573+
// stateless request is one-shot (no session, no resumption), so a
1574+
// dropped response is terminal and safe to cancel.
1575+
let request_ct = CancellationToken::new();
1576+
let service =
1577+
serve_directly_with_ct(service, transport, peer_info, request_ct.clone());
15481578
tokio::spawn(async move {
15491579
// on service created
15501580
let _ = service.waiting().await;
@@ -1554,8 +1584,19 @@ where
15541584
// emits an intermediate notification or request, preserve
15551585
// the complete message sequence by falling back to SSE.
15561586
let cancel = self.config.cancellation_token.child_token();
1587+
// Cancel the handler if the client disconnects while it is
1588+
// still producing its first message (this future is dropped
1589+
// before `receiver.recv()` completes). Disarmed once the
1590+
// handler emits anything, so a normal response is never
1591+
// cancelled.
1592+
let mut disconnect_guard = Some(request_ct.clone().drop_guard());
15571593
let Some(message) = (tokio::select! {
1558-
res = receiver.recv() => res,
1594+
res = receiver.recv() => {
1595+
if let Some(guard) = disconnect_guard.take() {
1596+
guard.disarm();
1597+
}
1598+
res
1599+
}
15591600
_ = cancel.cancelled() => None,
15601601
}) else {
15611602
return Err(internal_error_response("empty response")(
@@ -1579,6 +1620,9 @@ where
15791620
.body(Full::new(Bytes::from(body)).boxed())
15801621
.expect("valid response"))
15811622
} else {
1623+
// The handler emitted an intermediate message and is still
1624+
// running, so guard the streamed sequence too: dropping it
1625+
// (client disconnect) must cancel the handler.
15821626
let first = futures::stream::once(async move {
15831627
ServerSseMessage::from_message(message)
15841628
});
@@ -1587,17 +1631,19 @@ where
15871631
ServerSseMessage::from_message(message)
15881632
});
15891633
Ok(sse_stream_response(
1590-
first.chain(remaining),
1634+
CancelOnDisconnect::new(first.chain(remaining), request_ct),
15911635
self.config.sse_keep_alive,
15921636
self.config.cancellation_token.child_token(),
15931637
))
15941638
}
15951639
} else {
1596-
// SSE mode (default): original behaviour preserved unchanged
1640+
// SSE mode (default): cancel the handler if the client
1641+
// disconnects (drops the response stream) before it completes.
15971642
let stream = ReceiverStream::new(receiver).map(|message| {
15981643
tracing::trace!(?message);
15991644
ServerSseMessage::from_message(message)
16001645
});
1646+
let stream = CancelOnDisconnect::new(stream, request_ct);
16011647
Ok(sse_stream_response(
16021648
stream,
16031649
self.config.sse_keep_alive,
@@ -1680,3 +1726,52 @@ where
16801726
})
16811727
}
16821728
}
1729+
1730+
pin_project! {
1731+
/// Wraps a stateless SSE response stream so a client disconnect cancels the
1732+
/// in-flight request.
1733+
///
1734+
/// A stateless streamable-HTTP request is one-shot: it has no session and no
1735+
/// resumption, so a dropped response stream means the client is gone for
1736+
/// good. When the stream is dropped *before* it ends naturally, the request's
1737+
/// cancellation token is fired, which stops the dedicated `serve_directly`
1738+
/// loop and cancels the handler's `RequestContext::ct` (see #857). If the
1739+
/// stream ends naturally (the request completed), the guard is disarmed so
1740+
/// normal completion cancels nothing.
1741+
struct CancelOnDisconnect<S> {
1742+
#[pin]
1743+
inner: S,
1744+
ct: Option<CancellationToken>,
1745+
}
1746+
impl<S> PinnedDrop for CancelOnDisconnect<S> {
1747+
fn drop(this: Pin<&mut Self>) {
1748+
let this = this.project();
1749+
if let Some(ct) = this.ct.take() {
1750+
ct.cancel();
1751+
}
1752+
}
1753+
}
1754+
}
1755+
1756+
impl<S> CancelOnDisconnect<S> {
1757+
fn new(inner: S, ct: CancellationToken) -> Self {
1758+
Self {
1759+
inner,
1760+
ct: Some(ct),
1761+
}
1762+
}
1763+
}
1764+
1765+
impl<S: Stream> Stream for CancelOnDisconnect<S> {
1766+
type Item = S::Item;
1767+
1768+
fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
1769+
let this = self.project();
1770+
let polled = this.inner.poll_next(cx);
1771+
if let Poll::Ready(None) = &polled {
1772+
// Ended naturally: the request completed, so don't cancel on drop.
1773+
*this.ct = None;
1774+
}
1775+
polled
1776+
}
1777+
}

0 commit comments

Comments
 (0)