From ad88b6f3ca8f88355b510d5899160284c5b28e57 Mon Sep 17 00:00:00 2001 From: Nick Cooper Date: Sun, 23 Aug 2026 10:59:05 -0400 Subject: [PATCH] fix: reject replies from retired streamable http sessions Keep the inbound transport generation with automatic handler responses and errors. Reject replies after session recovery, including during the shutdown drain, without adding wire metadata. Match the completing handler's token before removing its request id so an old handler cannot cancel a replacement that reuses the id. Cover late replies, id reuse, scoped send context, and shutdown races. --- crates/rmcp/src/service.rs | 376 +++++++++++++++++- .../src/transport/streamable_http_client.rs | 2 +- crates/rmcp/src/transport/worker.rs | 210 +++++++++- ...test_streamable_http_client_concurrency.rs | 154 ++++++- 4 files changed, 705 insertions(+), 37 deletions(-) diff --git a/crates/rmcp/src/service.rs b/crates/rmcp/src/service.rs index b6cc5e538..d8b432fe7 100644 --- a/crates/rmcp/src/service.rs +++ b/crates/rmcp/src/service.rs @@ -1323,6 +1323,105 @@ where tokio::task::spawn_local(future) } +/// Local routing metadata retained separately from a handler's wire response. +#[derive(Debug, Clone, Default)] +struct ReplyContext { + #[cfg(feature = "transport-worker")] + origin: Option, +} + +impl ReplyContext { + fn from_request(request: &mut impl GetExtensions) -> Self { + #[cfg(feature = "transport-worker")] + { + Self { + origin: request + .extensions_mut() + .remove::(), + } + } + #[cfg(not(feature = "transport-worker"))] + { + let _ = request; + Self::default() + } + } + + fn is_current(&self) -> bool { + #[cfg(feature = "transport-worker")] + { + self.origin + .as_ref() + .is_none_or(|origin| origin.is_current()) + } + #[cfg(not(feature = "transport-worker"))] + { + true + } + } + + fn send>( + self, + transport: &mut T, + message: TxJsonRpcMessage, + ) -> impl Future> + Send + 'static { + let send = self.is_current().then(|| { + #[cfg(feature = "transport-worker")] + { + crate::transport::worker::with_response_origin(self.origin.clone(), || { + transport.send(message) + }) + } + #[cfg(not(feature = "transport-worker"))] + { + transport.send(message) + } + }); + async move { + let Some(send) = send else { + return Ok(()); + }; + let result = send.await; + // Recovery may retire the origin between the initial check and the + // worker's dispatch check. Dropping that reply must not stop a drain + // from sending later, current-session replies. + if self.is_current() { result } else { Ok(()) } + } + } +} + +#[derive(Debug)] +struct HandlerResponse { + message: TxJsonRpcMessage, + owner: Arc, + context: ReplyContext, +} + +impl HandlerResponse { + fn finish(&self, handlers: &mut HashMap>) -> bool { + // Complete this handler, never the handler that may have reused its id. + self.owner.cancel(); + let id = match &self.message { + JsonRpcMessage::Response(response) => &response.id, + JsonRpcMessage::Error(error) => match &error.id { + Some(id) => id, + None => return false, + }, + _ => return false, + }; + if handlers + .get(id) + .is_some_and(|owner| Arc::ptr_eq(owner, &self.owner)) + { + handlers.remove(id); + true + } else { + tracing::debug!(%id, "dropping response: request was cancelled or its id was reused"); + false + } + } +} + #[instrument(skip_all)] fn serve_inner( service: S, @@ -1339,7 +1438,7 @@ where { const SINK_PROXY_BUFFER_SIZE: usize = 64; let (sink_proxy_tx, mut sink_proxy_rx) = - tokio::sync::mpsc::channel::>(SINK_PROXY_BUFFER_SIZE); + tokio::sync::mpsc::channel::>(SINK_PROXY_BUFFER_SIZE); let peer_info = peer.peer_info(); if R::IS_CLIENT { tracing::info!(?peer_info, "Service initialized as client"); @@ -1349,7 +1448,7 @@ where let mut local_responder_pool = HashMap::>>::new(); - let mut local_ct_pool = HashMap::::new(); + let mut local_ct_pool = HashMap::>::new(); let shared_service = Arc::new(service); // for return let service = shared_service.clone(); @@ -1380,7 +1479,7 @@ where enum Event { ProxyMessage(PeerSinkMessage), PeerMessage(RxJsonRpcMessage), - ToSink(TxJsonRpcMessage), + ToSink(HandlerResponse), SendTaskResult(SendTaskResult), ResponseSendTaskResult(Result<(), tokio::task::JoinError>), } @@ -1474,18 +1573,9 @@ where } } // response and error - Event::ToSink(m) => { - if let Some(id) = match &m { - JsonRpcMessage::Response(response) => Some(&response.id), - JsonRpcMessage::Error(error) => error.id.as_ref(), - _ => None, - } { - let Some(ct) = local_ct_pool.remove(id) else { - tracing::debug!(%id, "dropping response for cancelled request"); - continue; - }; - ct.cancel(); - let send = transport.send(m); + Event::ToSink(response) => { + if response.finish(&mut local_ct_pool) { + let send = response.context.send(&mut transport, response.message); let current_span = tracing::Span::current(); response_send_tasks.spawn(async move { let send_result = send.await; @@ -1538,6 +1628,7 @@ where .. })) => { tracing::debug!(%id, ?request, "received request"); + let reply_context = ReplyContext::from_request(&mut request); if let Err(error) = R::enforce_peer_request_association( &request, peer.peer_info().as_deref(), @@ -1547,7 +1638,9 @@ where // send directly: the sink proxy path would drop the // error since the request was never registered in // local_ct_pool - let send = transport.send(JsonRpcMessage::error(error, Some(id))); + let send = reply_context.send( + &mut transport, JsonRpcMessage::error(error, Some(id)), + ); let current_span = tracing::Span::current(); response_send_tasks.spawn(async move { if let Err(error) = send.await { @@ -1559,9 +1652,9 @@ where { let service = shared_service.clone(); let sink = sink_proxy_tx.clone(); - let request_ct = serve_loop_ct.child_token(); + let request_ct = Arc::new(serve_loop_ct.child_token()); let context_ct = request_ct.child_token(); - local_ct_pool.insert(id.clone(), request_ct); + local_ct_pool.insert(id.clone(), request_ct.clone()); let mut extensions = Extensions::new(); let mut meta = RequestMetaObject::new(); // avoid clone @@ -1591,7 +1684,11 @@ where JsonRpcMessage::error(error, Some(id)) } }; - let _send_result = sink.send(response).await; + let _send_result = sink.send(HandlerResponse { + message: response, + owner: request_ct, + context: reply_context, + }).await; }.instrument(current_span)); } } @@ -1748,8 +1845,11 @@ where } // Then drain any handler responses still in the channel // (handlers that finished after the loop broke). - while let Some(m) = sink_proxy_rx.recv().await { - if let Err(error) = transport.send(m).await { + while let Some(response) = sink_proxy_rx.recv().await { + if !response.finish(&mut local_ct_pool) { + continue; + } + if let Err(error) = response.context.send(&mut transport, response.message).await { tracing::error!(%error, "failed to send pending response during drain"); break; } @@ -1777,6 +1877,240 @@ where } } +#[cfg(all(test, feature = "client"))] +mod reply_context_tests { + use super::*; + use crate::model::{ClientJsonRpcMessage, ClientResult}; + + #[test] + fn old_handler_completion_preserves_reused_id_owner() { + for error in [false, true] { + let id = RequestId::Number(7); + let old = Arc::new(CancellationToken::new()); + let current = Arc::new(CancellationToken::new()); + let mut handlers = HashMap::from([(id.clone(), current.clone())]); + let message = if error { + ClientJsonRpcMessage::error( + McpError::internal_error("old handler", None), + Some(id.clone()), + ) + } else { + ClientJsonRpcMessage::response(ClientResult::empty(()), id.clone()) + }; + let response = HandlerResponse:: { + message, + owner: old.clone(), + context: ReplyContext::default(), + }; + assert!(!response.finish(&mut handlers)); + assert!(old.is_cancelled()); + assert!(!current.is_cancelled()); + assert!(Arc::ptr_eq(handlers.get(&id).unwrap(), ¤t)); + let response = HandlerResponse:: { + message: ClientJsonRpcMessage::response(ClientResult::empty(()), id), + owner: current.clone(), + context: ReplyContext::default(), + }; + assert!(response.finish(&mut handlers)); + assert!(current.is_cancelled()); + assert!(handlers.is_empty()); + } + } + + #[cfg(all(feature = "transport-worker", not(feature = "local")))] + mod drain { + use std::{ + io, + sync::atomic::{AtomicU64, Ordering}, + }; + + use tokio::{ + sync::{mpsc, oneshot}, + time::timeout, + }; + + use super::*; + use crate::{ + ClientHandler, + model::{CustomRequest, CustomResult, ServerJsonRpcMessage}, + transport::worker::ResponseOrigin, + }; + + const TIMEOUT: Duration = Duration::from_secs(5); + + struct DrainTransport { + incoming: mpsc::UnboundedReceiver, + outgoing: mpsc::UnboundedSender, + eof: Option>, + failed_send: Option>, + } + + impl Transport for DrainTransport { + type Error = io::Error; + + fn send( + &mut self, + message: ClientJsonRpcMessage, + ) -> impl Future> + Send + 'static { + let outgoing = self.outgoing.clone(); + let failed_send = self.failed_send.take(); + async move { + if let Some(failed_send) = failed_send { + failed_send.await.unwrap(); + return Err(io::Error::other("retired reply")); + } + outgoing.send(message).map_err(io::Error::other) + } + } + + async fn receive(&mut self) -> Option { + let message = self.incoming.recv().await; + if message.is_none() + && let Some(eof) = self.eof.take() + { + let _ = eof.send(()); + } + message + } + + async fn close(&mut self) -> Result<(), Self::Error> { + Ok(()) + } + } + + struct HeldRequest { + cancellation: CancellationToken, + reply: oneshot::Sender>, + } + + struct HeldClient(mpsc::UnboundedSender); + + impl ClientHandler for HeldClient { + async fn on_custom_request( + &self, + _request: CustomRequest, + context: RequestContext, + ) -> Result { + let (reply, result) = oneshot::channel(); + self.0 + .send(HeldRequest { + cancellation: context.ct, + reply, + }) + .unwrap(); + result.await.unwrap() + } + } + + fn inbound(id: i64, generation: &Arc) -> ServerJsonRpcMessage { + let mut message: ServerJsonRpcMessage = serde_json::from_value(serde_json::json!({ + "jsonrpc": "2.0", "id": id, "method": "test/held", + })) + .unwrap(); + let JsonRpcMessage::Request(request) = &mut message else { + unreachable!() + }; + request + .request + .extensions_mut() + .insert(ResponseOrigin::capture(generation)); + message + } + + #[tokio::test] + async fn shutdown_drain_preserves_current_handler_after_old_response_or_error() { + for error in [false, true] { + for replacement_id in [8, 7] { + let (incoming, incoming_rx) = mpsc::unbounded_channel(); + let (outgoing, mut sent) = mpsc::unbounded_channel(); + let (eof, ended) = oneshot::channel(); + let (started, mut handlers) = mpsc::unbounded_channel(); + let running = serve_directly::( + HeldClient(started), + DrainTransport { + incoming: incoming_rx, + outgoing, + eof: Some(eof), + failed_send: None, + }, + None, + ); + let generation = Arc::new(AtomicU64::new(0)); + incoming.send(inbound(7, &generation)).unwrap(); + let old = timeout(TIMEOUT, handlers.recv()).await.unwrap().unwrap(); + generation.fetch_add(1, Ordering::SeqCst); + incoming.send(inbound(replacement_id, &generation)).unwrap(); + let current = timeout(TIMEOUT, handlers.recv()).await.unwrap().unwrap(); + // Both handlers remain held until receive returns eof and + // the service leaves its main loop for the response drain. + drop(incoming); + timeout(TIMEOUT, ended).await.unwrap().unwrap(); + old.reply + .send(if error { + Err(McpError::internal_error("old handler", None)) + } else { + Ok(CustomResult::new(serde_json::json!({"reply": "old"}))) + }) + .unwrap(); + timeout(TIMEOUT, old.cancellation.cancelled()) + .await + .unwrap(); + assert!(!current.cancellation.is_cancelled()); + current + .reply + .send(Ok(CustomResult::new( + serde_json::json!({"reply": "current"}), + ))) + .unwrap(); + let message = timeout(TIMEOUT, sent.recv()).await.unwrap().unwrap(); + let value = serde_json::to_value(message).unwrap(); + assert_eq!(value["id"], replacement_id); + assert_eq!(value["result"], serde_json::json!({"reply": "current"})); + assert!(matches!( + timeout(TIMEOUT, running.waiting()).await.unwrap().unwrap(), + QuitReason::Closed + )); + assert!(sent.try_recv().is_err()); + } + } + } + + #[tokio::test] + async fn recovery_during_reply_send_does_not_stop_the_response_drain() { + let generation = Arc::new(AtomicU64::new(0)); + let context = ReplyContext { + origin: Some(ResponseOrigin::capture(&generation)), + }; + let (release, failed_send) = oneshot::channel(); + let (outgoing, mut sent) = mpsc::unbounded_channel(); + let mut transport = DrainTransport { + incoming: mpsc::unbounded_channel().1, + outgoing, + eof: None, + failed_send: Some(failed_send), + }; + let message = + || ClientJsonRpcMessage::response(ClientResult::empty(()), RequestId::Number(7)); + let send = context.send(&mut transport, message()); + generation.fetch_add(1, Ordering::SeqCst); + release.send(()).unwrap(); + assert!( + send.await.is_ok(), + "retired replies must not abort the drain" + ); + let context = ReplyContext { + origin: Some(ResponseOrigin::capture(&generation)), + }; + context.send(&mut transport, message()).await.unwrap(); + assert_eq!( + serde_json::to_value(sent.recv().await.unwrap()).unwrap(), + serde_json::to_value(message()).unwrap() + ); + assert!(sent.try_recv().is_err()); + } + } +} + #[cfg(all(test, feature = "server"))] mod sep2260_marker_tests { use std::sync::Arc; diff --git a/crates/rmcp/src/transport/streamable_http_client.rs b/crates/rmcp/src/transport/streamable_http_client.rs index d702bc1ca..82629df48 100644 --- a/crates/rmcp/src/transport/streamable_http_client.rs +++ b/crates/rmcp/src/transport/streamable_http_client.rs @@ -1386,7 +1386,7 @@ impl Worker for StreamableHttpClientWorker { } let cancellation_request_id = Self::cancellation_request_id(&send_request.message); - let stale = send_request.control_generation() != context.control_generation(); + let stale = !context.is_current_control(&send_request); if stale { // Do not send old controls to a replacement session. let result = match cancellation_request_id { diff --git a/crates/rmcp/src/transport/worker.rs b/crates/rmcp/src/transport/worker.rs index e32da7d7d..6ee930592 100644 --- a/crates/rmcp/src/transport/worker.rs +++ b/crates/rmcp/src/transport/worker.rs @@ -13,7 +13,7 @@ use tracing::{Instrument, Level}; use super::{IntoTransport, Transport}; use crate::{ - model::{CancelledNotification, JsonRpcMessage, RequestId}, + model::{CancelledNotification, GetExtensions, JsonRpcMessage, RequestId}, service::{RxJsonRpcMessage, ServiceRole, TxJsonRpcMessage}, }; @@ -81,6 +81,48 @@ pub trait Worker: Sized + Send + 'static { type RequestCancellations = Arc>>>; +/// The transport and generation that delivered an inbound request. Never serialized. +#[derive(Debug, Clone)] +pub(crate) struct ResponseOrigin { + transport: Arc, + generation: u64, +} + +impl ResponseOrigin { + /// Capture a transport's current generation before handing off a request. + pub(crate) fn capture(transport: &Arc) -> Self { + Self { + transport: transport.clone(), + generation: transport.load(Ordering::SeqCst), + } + } + + /// Return whether the originating transport still uses this generation. + pub(crate) fn is_current(&self) -> bool { + self.generation == self.transport.load(Ordering::SeqCst) + } +} + +tokio::task_local! { + static RESPONSE_ORIGIN: Option; +} + +/// Preserve reply origin during both send construction and future polling. +/// +/// Lazy wrappers polled in this future retain the context. Wrappers that spawn +/// tasks must construct the inner send before spawning it. Deferring that call +/// to a detached task, or serializing away request extensions, loses the origin. +pub(crate) fn with_response_origin( + origin: Option, + create_send: impl FnOnce() -> F, +) -> impl Future + Send + 'static +where + F: Future + Send + 'static, +{ + let send = RESPONSE_ORIGIN.sync_scope(origin.clone(), create_send); + RESPONSE_ORIGIN.scope(origin, send) +} + /// Keeps a request's cancellation token registered for a chosen lifetime. pub(crate) struct RequestCancellationRegistration { id: RequestId, @@ -135,6 +177,8 @@ pub struct WorkerSendRequest { pub responder: tokio::sync::oneshot::Sender>, cancellation: Option>, control_generation: u64, + #[cfg(feature = "transport-streamable-http-client")] + response_origin: Option, } impl WorkerSendRequest { @@ -156,11 +200,12 @@ impl WorkerSendRequest { self.cancellation.clone() } - /// Return the local generation captured when [`Transport::send`] created its future. + /// Return the local generation associated with this send. /// - /// This happens before polling or queue admission. The value is not sent over - /// the wire; the worker decides whether a message from an older generation is valid. - /// It identifies the outbound send, not the session that started an inbound handler. + /// Ordinary sends capture it when [`Transport::send`] creates its future, + /// before polling or queue admission. Automatic replies retain the generation + /// that delivered their inbound request. This value is not sent over the wire; + /// the worker decides whether a message from an older generation is valid. pub fn control_generation(&self) -> u64 { self.control_generation } @@ -301,17 +346,19 @@ pub struct WorkerContext { } impl WorkerContext { - /// Return the local generation that newly created sends will capture. + /// Return the local generation that new sends without a reply origin capture. /// /// Workers may use this value to check messages from an earlier connection - /// or session. The generic transport does not check it automatically. + /// or session. Automatic replies retain their inbound request's generation. + /// The generic transport does not check generations automatically. pub fn control_generation(&self) -> u64 { self.control_generation.load(Ordering::SeqCst) } /// Advance the local generation, wrapping at [`u64::MAX`], and return its new value. /// - /// Only subsequent calls to [`Transport::send`] capture the new value. + /// Subsequent sends without a reply origin capture the new value. Automatic + /// replies keep their inbound request's generation even when sent later. /// Advancing does not drain queues, cancel work, or reject older messages; /// the worker is responsible for those actions. pub fn advance_control_generation(&self) -> u64 { @@ -320,10 +367,27 @@ impl WorkerContext { .wrapping_add(1) } + /// Check both the captured generation and, for replies, the originating transport. + #[cfg(feature = "transport-streamable-http-client")] + pub(crate) fn is_current_control(&self, request: &WorkerSendRequest) -> bool { + request.control_generation == self.control_generation() + && request + .response_origin + .as_ref() + .is_none_or(|origin| Arc::ptr_eq(&origin.transport, &self.control_generation)) + } + + /// Queue an inbound message, recording request origin before waiting for capacity. pub async fn send_to_handler( &mut self, - item: RxJsonRpcMessage, + mut item: RxJsonRpcMessage, ) -> Result<(), WorkerQuitReason> { + if let JsonRpcMessage::Request(request) = &mut item { + request + .request + .extensions_mut() + .insert(ResponseOrigin::capture(&self.control_generation)); + } self.to_handler_tx .send(item) .await @@ -347,7 +411,18 @@ impl Transport for WorkerTransport { &mut self, item: TxJsonRpcMessage, ) -> impl Future> + Send + 'static { - let control_generation = self.control_generation.load(Ordering::SeqCst); + let response_origin = if matches!( + &item, + JsonRpcMessage::Response(_) | JsonRpcMessage::Error(_) + ) { + RESPONSE_ORIGIN.try_with(Clone::clone).ok().flatten() + } else { + None + }; + let control_generation = response_origin.as_ref().map_or_else( + || self.control_generation.load(Ordering::SeqCst), + |origin| origin.generation, + ); let mut cancellation_target = None; let registration = if W::supports_request_cancellation() { match &item { @@ -383,6 +458,8 @@ impl Transport for WorkerTransport { responder, cancellation: registration, control_generation, + #[cfg(feature = "transport-streamable-http-client")] + response_origin, }; async move { // Keep the stream alive until its cancellation is handled or abandoned. @@ -415,7 +492,10 @@ mod tests { use std::io; use super::*; - use crate::{model::ClientJsonRpcMessage, service::RoleClient}; + use crate::{ + model::{ClientJsonRpcMessage, ClientResult, ServerJsonRpcMessage}, + service::RoleClient, + }; struct TestWorker(tokio::sync::oneshot::Sender>); @@ -461,6 +541,114 @@ mod tests { .unwrap() } + #[tokio::test] + async fn inbound_origin_is_captured_before_queue_admission() { + let (context_tx, context_rx) = tokio::sync::oneshot::channel(); + let mut transport = WorkerTransport::spawn(TestWorker(context_tx)); + let mut context = context_rx.await.unwrap(); + let generation = context.control_generation.clone(); + let request = || { + serde_json::from_value::(serde_json::json!({ + "jsonrpc": "2.0", "id": 7, "method": "ping", + })) + .unwrap() + }; + for _ in 0..WorkerConfig::default().channel_buffer_capacity { + context.send_to_handler(request()).await.unwrap(); + } + let mut pending = Box::pin(context.send_to_handler(request())); + assert!(futures::poll!(pending.as_mut()).is_pending()); + generation.fetch_add(1, Ordering::SeqCst); + transport.receive().await.unwrap(); + pending.await.unwrap(); + for _ in 1..WorkerConfig::default().channel_buffer_capacity { + transport.receive().await.unwrap(); + } + let message = transport.receive().await.unwrap(); + assert_eq!( + serde_json::to_value(&message).unwrap(), + serde_json::to_value(request()).unwrap() + ); + let JsonRpcMessage::Request(request) = message else { + panic!("expected request") + }; + let origin = request + .request + .extensions() + .get::() + .unwrap(); + assert!(Arc::ptr_eq(&origin.transport, &generation)); + assert_eq!(origin.generation, 0); + assert!(!origin.is_current()); + transport.close().await.unwrap(); + } + + #[tokio::test] + async fn reply_origin_survives_send_creation_and_lazy_polling() { + for response in [ + ClientJsonRpcMessage::response(ClientResult::empty(()), RequestId::Number(7)), + ClientJsonRpcMessage::error( + crate::model::ErrorData::internal_error("handler error", None), + Some(RequestId::Number(7)), + ), + ] { + let (context_tx, context_rx) = tokio::sync::oneshot::channel(); + let mut transport = WorkerTransport::spawn(TestWorker(context_tx)); + let mut context = context_rx.await.unwrap(); + let origin = ResponseOrigin::capture(&context.control_generation); + context.advance_control_generation(); + let eager = + with_response_origin(Some(origin.clone()), || transport.send(response.clone())); + assert!(RESPONSE_ORIGIN.try_with(Clone::clone).is_err()); + let lazy = with_response_origin(Some(origin), || async move { + tokio::task::yield_now().await; + let result = transport.send(response).await; + (transport, result) + }); + let eager = tokio::spawn(eager); + let lazy = tokio::spawn(lazy); + for _ in 0..2 { + let request = context.from_handler_rx.recv().await.unwrap(); + assert_eq!(request.control_generation(), 0); + request.responder.send(Ok(())).unwrap(); + } + eager.await.unwrap().unwrap(); + let (mut transport, result) = lazy.await.unwrap(); + result.unwrap(); + assert!(RESPONSE_ORIGIN.try_with(Clone::clone).is_err()); + transport.close().await.unwrap(); + } + } + + #[cfg(feature = "transport-streamable-http-client")] + #[tokio::test] + async fn reply_origin_checks_transport_identity_even_when_generations_match() { + let (context_tx, context_rx) = tokio::sync::oneshot::channel(); + let mut transport = WorkerTransport::spawn(TestWorker(context_tx)); + let mut context = context_rx.await.unwrap(); + for foreign in [false, true] { + let generation = if foreign { + Arc::new(AtomicU64::new(context.control_generation())) + } else { + context.control_generation.clone() + }; + let origin = ResponseOrigin::capture(&generation); + let send = with_response_origin(Some(origin), || { + transport.send(ClientJsonRpcMessage::response( + ClientResult::empty(()), + RequestId::Number(7), + )) + }); + let send = tokio::spawn(send); + let request = context.from_handler_rx.recv().await.unwrap(); + assert_eq!(request.control_generation(), context.control_generation()); + assert_eq!(context.is_current_control(&request), !foreign); + request.responder.send(Ok(())).unwrap(); + send.await.unwrap().unwrap(); + } + transport.close().await.unwrap(); + } + #[tokio::test] async fn cancellation_matches_request_id_exactly() { let (context_tx, context_rx) = tokio::sync::oneshot::channel(); diff --git a/crates/rmcp/tests/test_streamable_http_client_concurrency.rs b/crates/rmcp/tests/test_streamable_http_client_concurrency.rs index 20b725831..d7a5739e5 100644 --- a/crates/rmcp/tests/test_streamable_http_client_concurrency.rs +++ b/crates/rmcp/tests/test_streamable_http_client_concurrency.rs @@ -15,14 +15,15 @@ use std::{ use futures::{StreamExt, stream::BoxStream}; use http::{HeaderName, HeaderValue}; use rmcp::{ + ClientHandler, model::{ CallToolRequestParams, CancelledNotificationParam, ClientInfo, ClientJsonRpcMessage, - ClientRequest, DiscoverResult, ProtocolVersion, Request, RequestId, RequestMetaObject, - ServerJsonRpcMessage, + ClientRequest, CustomRequest, CustomResult, DiscoverResult, ErrorData, ProtocolVersion, + Request, RequestId, RequestMetaObject, ServerJsonRpcMessage, }, service::{ - ClientLifecycleMode, PeerRequestOptions, RequestHandle, RoleClient, RunningService, - serve_client_with_lifecycle, + ClientLifecycleMode, PeerRequestOptions, RequestContext, RequestHandle, RoleClient, + RunningService, serve_client_with_lifecycle, }, transport::streamable_http_client::{ StreamableHttpClient, StreamableHttpClientTransport, StreamableHttpClientTransportConfig, @@ -744,6 +745,151 @@ async fn server_replies_still_run_while_recovery_waits_for_old_posts() -> anyhow harness.finish(vec![waiting, expired], 3, 2).await } +struct HeldServerRequest { + id: RequestId, + cancellation: CancellationToken, + reply: oneshot::Sender>, +} + +struct HeldRequestClient { + started: mpsc::UnboundedSender, +} + +impl ClientHandler for HeldRequestClient { + async fn on_custom_request( + &self, + _request: CustomRequest, + context: RequestContext, + ) -> Result { + let (reply, result) = oneshot::channel(); + self.started + .send(HeldServerRequest { + id: context.id, + cancellation: context.ct, + reply, + }) + .expect("test receives the server request"); + result.await.expect("test releases the handler") + } +} + +async fn check_late_session_handler_reply( + old_result: Result, + replacement_id: i64, +) -> anyhow::Result<()> { + let (started, mut requests) = mpsc::unbounded_channel(); + let (control_tx, mut controls) = mpsc::unbounded_channel(); + let (incoming, incoming_rx) = mpsc::unbounded_channel(); + let (handler_tx, mut handlers) = mpsc::unbounded_channel(); + let counts = Arc::new(Counts::default()); + counts.manual_controls.store(true, SeqCst); + let transport = StreamableHttpClientTransport::with_client( + ScriptedClient { + started, + controls: control_tx, + incoming: Arc::new(Mutex::new(Some(incoming_rx))), + reinitializing: mpsc::unbounded_channel().0, + counts: counts.clone(), + }, + config(), + ); + let client = serve_client_with_lifecycle( + HeldRequestClient { + started: handler_tx, + }, + transport, + ClientLifecycleMode::Initialize, + ) + .await?; + + incoming.send(sse(json!({ + "jsonrpc": "2.0", "id": 7, "method": "test/held", + })))?; + let old = next_event(&mut handlers).await; + assert_eq!(old.id, RequestId::Number(7)); + + let peer = client.peer().clone(); + let call = + tokio::spawn(async move { peer.call_tool(CallToolRequestParams::new("recover")).await }); + next_event(&mut requests).await.expire(); + let retry = next_event(&mut requests).await; + assert_eq!(retry.session.as_deref(), Some("session-2")); + let final_response = serde_json::to_value(retry.result())?; + let (replacement_stream, replacement_rx) = mpsc::unbounded_channel(); + retry + .finish_and_wait(Ok(StreamableHttpPostResponse::Sse( + UnboundedReceiverStream::new(replacement_rx).boxed(), + None, + ))) + .await?; + replacement_stream.send(sse(json!({ + "jsonrpc": "2.0", "id": replacement_id, "method": "test/held", + })))?; + let replacement = next_event(&mut handlers).await; + assert_eq!(replacement.id, RequestId::Number(replacement_id)); + replacement_stream.send(sse(final_response))?; + timeout(TEST_TIMEOUT, call).await???; + + old.reply.send(old_result).expect("old handler is waiting"); + // Completion must cancel only the completing handler's token. Waiting for + // that acknowledgement orders the old completion before the new one. + timeout(TEST_TIMEOUT, async { + tokio::select! { + _ = old.cancellation.cancelled() => {} + _ = replacement.cancellation.cancelled() => { + panic!("old completion cancelled the replacement handler"); + } + } + }) + .await?; + assert!(!replacement.cancellation.is_cancelled()); + replacement + .reply + .send(Ok(CustomResult::new(json!({ "reply": "replacement" })))) + .expect("replacement handler is waiting"); + + let control = next_event(&mut controls).await; + assert_eq!(control.session.as_deref(), Some("session-2")); + assert_eq!(control.message["id"], replacement_id); + assert_eq!(control.message["result"], json!({ "reply": "replacement" })); + control + .reply + .send(Ok(StreamableHttpPostResponse::Accepted)) + .unwrap(); + timeout(TEST_TIMEOUT, replacement.cancellation.cancelled()).await?; + // Closing drains automatic send tasks, so a delayed old reply cannot evade + // the assertion by starting after the replacement reply was acknowledged. + timeout(TEST_TIMEOUT, client.cancel()).await??; + assert!( + controls.try_recv().is_err(), + "no old-session reply may be posted" + ); + assert_eq!(counts.initialized.load(SeqCst), 2); + Ok(()) +} + +#[tokio::test] +async fn old_session_handler_response_does_not_reach_replacement_session() -> anyhow::Result<()> { + check_late_session_handler_reply(Ok(CustomResult::new(json!({ "reply": "old" }))), 8).await +} + +#[tokio::test] +async fn old_session_handler_error_does_not_reach_replacement_session() -> anyhow::Result<()> { + check_late_session_handler_reply(Err(ErrorData::internal_error("old handler error", None)), 8) + .await +} + +#[tokio::test] +async fn old_session_handler_response_preserves_reused_request_id() -> anyhow::Result<()> { + check_late_session_handler_reply(Ok(CustomResult::new(json!({ "reply": "old" }))), 7).await +} + +#[tokio::test] +async fn old_session_handler_error_preserves_reused_request_id() -> anyhow::Result<()> { + check_late_session_handler_reply(Err(ErrorData::internal_error("old handler error", None)), 7) + .await +} + #[tokio::test] async fn a_version_barrier_allows_the_server_reply_it_needs() -> anyhow::Result<()> { let mut harness = Harness::start(config().max_concurrent_requests(2)).await?;