11use 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
512use bytes:: Bytes ;
6- use futures:: { StreamExt , future:: BoxFuture } ;
13+ use futures:: { Stream , StreamExt , future:: BoxFuture } ;
714use http:: { HeaderMap , Method , Request , Response , header:: ALLOW } ;
815use http_body:: Body ;
916use http_body_util:: { BodyExt , Full , combinators:: BoxBody } ;
17+ use pin_project_lite:: pin_project;
1018use tokio_stream:: wrappers:: ReceiverStream ;
1119use 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