diff --git a/crates/rmcp/src/service.rs b/crates/rmcp/src/service.rs index b6cc5e538..7cf14e03a 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: &impl GetExtensions) -> Self { + #[cfg(feature = "transport-worker")] + { + Self { + origin: request + .extensions() + .get::() + .cloned(), + } + } + #[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 { + 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(&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?;