diff --git a/README.md b/README.md index 6ab22d67c..44892e833 100644 --- a/README.md +++ b/README.md @@ -1631,6 +1631,17 @@ let transport = StreamableHttpClientTransport::from_uri("http://localhost:8000/m let client = ClientInfo::default().serve(transport).await?; ``` +The client allows up to 16 ordinary http POSTs at once. Configure this with +`StreamableHttpClientTransportConfig::with_uri(url).max_concurrent_requests(n)`; +`1` keeps ordinary POSTs serial, and `0` is treated as `1`. An open sse response +stream does not count against this limit. Cancellation and replies use a +separate queue with one extra POST slot. Each control POST has a five-second +timeout after it starts. Session recovery waits up to five seconds for old +POSTs, then stops any that remain. Those POSTs are not retried because the +server may have processed them. Configure this wait and the separate +initialization timeout with `session_recovery_timeout`. Callers still decide +which tools may run at the same time and which need approval. + #### Server-Sent Events (SSE) Streamable HTTP responses arrive as either a single `application/json` body or a diff --git a/crates/rmcp/Cargo.toml b/crates/rmcp/Cargo.toml index 9dd27a251..5dc37e817 100644 --- a/crates/rmcp/Cargo.toml +++ b/crates/rmcp/Cargo.toml @@ -295,6 +295,11 @@ name = "test_streamable_http_json_response" required-features = ["server", "client", "transport-streamable-http-server", "reqwest"] path = "tests/test_streamable_http_json_response.rs" +[[test]] +name = "test_streamable_http_client_concurrency" +required-features = ["client", "transport-streamable-http-client"] +path = "tests/test_streamable_http_client_concurrency.rs" + [[test]] name = "test_streamable_http_protocol_version" required-features = ["server", "client", "transport-streamable-http-server", "reqwest"] diff --git a/crates/rmcp/src/transport/streamable_http_client.rs b/crates/rmcp/src/transport/streamable_http_client.rs index fe563ee3f..bba6ef8fd 100644 --- a/crates/rmcp/src/transport/streamable_http_client.rs +++ b/crates/rmcp/src/transport/streamable_http_client.rs @@ -1,11 +1,15 @@ use std::{ borrow::Cow, - collections::{HashMap, HashSet}, + collections::{HashMap, HashSet, VecDeque}, sync::Arc, time::Duration, }; -use futures::{Stream, StreamExt, future::BoxFuture, stream::BoxStream}; +use futures::{ + Stream, StreamExt, + future::BoxFuture, + stream::{BoxStream, FuturesUnordered}, +}; use http::{HeaderName, HeaderValue}; pub use sse_stream::Error as SseError; use sse_stream::Sse; @@ -26,13 +30,17 @@ use crate::{ service::InboundStreamOrigin, transport::{ common::{client_side_sse::SseAutoReconnectStream, mcp_headers}, - worker::{Worker, WorkerQuitReason, WorkerSendRequest, WorkerTransport}, + worker::{ + RequestCancellationRegistration, Worker, WorkerQuitReason, WorkerSendRequest, + WorkerTransport, + }, }, }; type BoxedSseStream = BoxStream<'static, Result>; type SseTaskResult = (Option, Result<(), StreamableHttpError>); const SESSION_CLEANUP_TIMEOUT: Duration = Duration::from_secs(5); +const CONTROL_POST_TIMEOUT: Duration = Duration::from_secs(5); fn build_request_headers( base: &HashMap, @@ -207,6 +215,12 @@ pub enum StreamableHttpError { ReservedHeaderConflict(String), #[error("Session expired (HTTP 404)")] SessionExpired, + /// Session recovery timed out. The server may have processed an interrupted POST. + #[error("Session recovery timed out; the server may have processed the POST")] + SessionRecoveryTimeout, + /// A cancellation or reply POST did not finish in time. + #[error("Control POST timed out")] + ControlRequestTimeout, } impl StreamableHttpError { @@ -475,6 +489,21 @@ pub struct StreamableHttpClientWorker { pub config: StreamableHttpClientTransportConfig, } +struct PostResult { + send_request: WorkerSendRequest>, + // None means the send future was dropped, or the request or transport was cancelled. + response: Option>>, + // The protocol version used to send this POST. + version: ProtocolVersion, +} + +struct PostSession { + id: Option>, + headers: HashMap, + version: ProtocolVersion, + cancellation: CancellationToken, +} + impl StreamableHttpClientWorker { pub fn new_simple(url: impl Into>) -> Self { Self { @@ -494,6 +523,84 @@ impl StreamableHttpClientWorker { } impl StreamableHttpClientWorker { + // Run initialization and protocol-version changes without other ordinary POSTs. + fn is_ordering_barrier( + message: &ClientJsonRpcMessage, + negotiated_version: &ProtocolVersion, + ) -> bool { + match message { + ClientJsonRpcMessage::Request(request) => { + matches!( + &request.request, + ClientRequest::InitializeRequest(_) | ClientRequest::DiscoverRequest(_) + ) || request + .request + .get_meta() + .protocol_version() + .is_some_and(|version| &version != negotiated_version) + } + ClientJsonRpcMessage::Notification(notification) => matches!( + ¬ification.notification, + ClientNotification::InitializedNotification(_) + ), + _ => false, + } + } + + fn post_request( + client: C, + config: &StreamableHttpClientTransportConfig, + mut send_request: WorkerSendRequest, + session: PostSession, + transport_cancellation: CancellationToken, + ) -> BoxFuture<'static, PostResult> { + let uri = config.uri.clone(); + let auth_header = config.auth_header.clone(); + let max_sse_event_size = config.max_sse_event_size; + let is_control = Self::is_control_message(&send_request.message); + let cancellation = send_request + .cancellation_token() + .unwrap_or_else(|| transport_cancellation.child_token()); + Box::pin(async move { + let response = tokio::select! { + biased; + _ = cancellation.cancelled() => None, + _ = send_request.responder.closed() => None, + _ = session.cancellation.cancelled() => { + Some(Err(StreamableHttpError::SessionRecoveryTimeout)) + }, + _ = tokio::time::sleep(CONTROL_POST_TIMEOUT), if is_control => { + Some(Err(StreamableHttpError::ControlRequestTimeout)) + }, + response = client.post_message_with_max_sse_event_size( + uri, + send_request.message.clone(), + session.id, + auth_header, + session.headers, + max_sse_event_size, + ) => Some(response), + }; + PostResult { + send_request, + response, + version: session.version, + } + }) + } + + fn cancellation_request_id(message: &ClientJsonRpcMessage) -> Option<&RequestId> { + match message { + ClientJsonRpcMessage::Notification(notification) => match ¬ification.notification { + ClientNotification::CancelledNotification(cancelled) => { + cancelled.params.request_id.as_ref() + } + _ => None, + }, + _ => None, + } + } + fn client_request_id(message: &ClientJsonRpcMessage) -> Option { match message { ClientJsonRpcMessage::Request(request) => Some(request.id.clone()), @@ -819,6 +926,22 @@ impl StreamableHttpClientWorker { impl Worker for StreamableHttpClientWorker { type Role = RoleClient; type Error = StreamableHttpError; + fn is_control_message(message: &ClientJsonRpcMessage) -> bool { + matches!( + message, + ClientJsonRpcMessage::Response(_) | ClientJsonRpcMessage::Error(_) + ) || matches!( + message, + ClientJsonRpcMessage::Notification(notification) + if matches!( + ¬ification.notification, + ClientNotification::CancelledNotification(_) + ) + ) + } + fn supports_request_cancellation() -> bool { + true + } fn err_closed() -> Self::Error { StreamableHttpError::TransportChannelClosed } @@ -844,6 +967,7 @@ impl Worker for StreamableHttpClientWorker { let WorkerSendRequest { responder, message: startup_request, + .. } = context.recv_from_handler().await?; let is_legacy_startup = matches!( &startup_request, @@ -953,17 +1077,31 @@ impl Worker for StreamableHttpClientWorker { clippy::large_enum_variant, reason = "the event is short-lived and boxing would add allocation in the event loop" )] - enum Event { - ClientMessage(WorkerSendRequest), + enum Event { + ClientMessage(WorkerSendRequest>), + ControlMessage(WorkerSendRequest>), + StartPost(WorkerSendRequest>), + PostResult(PostResult), + RecoveryTimeout, ServerMessage(ServerJsonRpcMessage), StreamResult { request_id: Option, - result: Result<(), StreamableHttpError>, + result: Result<(), StreamableHttpError>, }, } let mut streams = tokio::task::JoinSet::new(); let mut pending_stream_response_ids = HashSet::new(); - let mut request_stream_cancellations = HashMap::::new(); + let mut request_stream_cancellations = + HashMap::>::new(); + let mut posts = FuturesUnordered::>>::new(); + let mut control_posts = FuturesUnordered::>>::new(); + let mut session_cancellation = CancellationToken::new(); + let mut pending_message: Option> = None; + let mut recovery_posts = VecDeque::>::new(); + let mut recovery_deadline: Option = None; + let mut retrying_recovery = false; + let mut barrier_in_flight = false; + let max_concurrent_requests = config.max_concurrent_requests.max(1); let mut awaiting_fallback_initialized = false; if let Some(session_id) = &session_id { Self::spawn_common_stream( @@ -976,19 +1114,147 @@ impl Worker for StreamableHttpClientWorker { transport_task_ct.clone(), ); } - // Main event loop - capture exit reason so we can do cleanup before returning + // Each POST uses the session and headers chosen when it starts. + // Only this loop updates the current session and protocol version. let loop_result: Result<(), WorkerQuitReason> = 'main_loop: loop { + if retrying_recovery && recovery_posts.is_empty() && posts.is_empty() { + retrying_recovery = false; + } + if !retrying_recovery + && !recovery_posts.is_empty() + && posts.is_empty() + && control_posts.is_empty() + { + // Old POSTs have finished or reached the drain deadline. + // Retry only ordinary POSTs that returned SessionExpired, at most once each. + session_cancellation.cancel(); + recovery_deadline = None; + let recovery = tokio::select! { + _ = transport_task_ct.cancelled() => { + break 'main_loop Err(WorkerQuitReason::Cancelled); + } + result = tokio::time::timeout( + config.session_recovery_timeout, + Self::perform_reinitialization( + self.client.clone(), + saved_init_request.clone().expect("session recovery requires an initialize request"), + config.uri.clone(), + config.auth_header.clone(), + config.custom_headers.clone(), + config.max_sse_event_size, + ), + ) => result.unwrap_or(Err(StreamableHttpError::SessionRecoveryTimeout)), + }; + match recovery { + Ok((new_session_id, new_version, new_headers)) => { + streams.abort_all(); + while streams.join_next().await.is_some() {} + request_stream_cancellations.clear(); + Self::drain_queued_stream_messages( + &mut sse_worker_rx, + &mut context, + &mut pending_stream_response_ids, + ) + .await?; + Self::fail_pending_stream_responses( + &mut context, + &mut pending_stream_response_ids, + ) + .await?; + session_id = new_session_id; + negotiated_version = new_version; + protocol_headers = new_headers; + session_cleanup_info = session_id.as_ref().map(|sid| SessionCleanupInfo { + client: self.client.clone(), + uri: config.uri.clone(), + session_id: sid.clone(), + auth_header: config.auth_header.clone(), + protocol_headers: protocol_headers.clone(), + }); + // Do not send controls queued during recovery to the new session. + context.advance_control_generation(); + session_cancellation = CancellationToken::new(); + if let Some(session_id) = &session_id { + Self::spawn_common_stream( + &mut streams, + self.client.clone(), + session_id.clone(), + &config, + protocol_headers.clone(), + sse_worker_tx.clone(), + transport_task_ct.clone(), + ); + } + retrying_recovery = true; + } + Err(error) => { + session_cancellation = CancellationToken::new(); + // The backend error cannot be cloned. Return it to one caller + // and return the original session-expired error to the others. + if let Some(send_request) = recovery_posts.pop_front() { + let _ = send_request.responder.send(Err(error)); + } + for send_request in recovery_posts.drain(..) { + let _ = send_request + .responder + .send(Err(StreamableHttpError::SessionExpired)); + } + } + } + continue; + } + + let may_start = (retrying_recovery || recovery_posts.is_empty()) + && !barrier_in_flight + && posts.len() < max_concurrent_requests; + let queued = if retrying_recovery { + recovery_posts.front() + } else { + pending_message.as_ref() + }; + let can_dispatch = may_start + && queued.is_some_and(|request| { + !Self::is_ordering_barrier(&request.message, &negotiated_version) + || (posts.is_empty() && control_posts.is_empty()) + }); let event = tokio::select! { + _ = std::future::ready(()), if can_dispatch => { + let request = if retrying_recovery { + recovery_posts.pop_front() + } else { + pending_message.take() + }; + Event::StartPost(request.expect("a POST is ready to start")) + } _ = transport_task_ct.cancelled() => { tracing::debug!("cancelled"); break 'main_loop Err(WorkerQuitReason::Cancelled); } - message = context.recv_from_handler() => { + message = context.from_handler_rx.recv(), + if may_start && pending_message.is_none() && !retrying_recovery => { match message { - Ok(msg) => Event::ClientMessage(msg), - Err(e) => break 'main_loop Err(e), + Some(msg) => Event::ClientMessage(msg), + None => break 'main_loop Err(WorkerQuitReason::HandlerTerminated), } }, + message = context.control_from_handler_rx.recv(), + if control_posts.is_empty() && !session_cancellation.is_cancelled() => { + match message { + Some(msg) => Event::ControlMessage(msg), + None => break 'main_loop Err(WorkerQuitReason::HandlerTerminated), + } + }, + Some(result) = posts.next(), if !posts.is_empty() => { + Event::PostResult(result) + }, + Some(result) = control_posts.next(), if !control_posts.is_empty() => { + Event::PostResult(result) + }, + _ = async { + if let Some(deadline) = recovery_deadline { + tokio::time::sleep_until(deadline).await; + } + }, if recovery_deadline.is_some() => Event::RecoveryTimeout, message = sse_worker_rx.recv() => { let Some(message) = message else { tracing::trace!("transport dropped, exiting"); @@ -998,48 +1264,90 @@ impl Worker for StreamableHttpClientWorker { }, terminated_stream = streams.join_next(), if !streams.is_empty() => { match terminated_stream { - Some(result) => { - match result { - Ok((request_id, result)) => { - Event::StreamResult { request_id, result } - } - Err(error) => Event::StreamResult { - request_id: None, - result: Err(StreamableHttpError::TokioJoinError(error)), - }, - } - } - None => { - continue + Some(Ok((request_id, result))) => { + Event::StreamResult { request_id, result } } + Some(Err(error)) => Event::StreamResult { + request_id: None, + result: Err(StreamableHttpError::TokioJoinError(error)), + }, + None => continue, } } }; match event { Event::ClientMessage(send_request) => { - let WorkerSendRequest { message, responder } = send_request; - let cancellation_request_id = match &message { - ClientJsonRpcMessage::Notification(notification) => { - match ¬ification.notification { - ClientNotification::CancelledNotification(cancelled) => { - cancelled.params.request_id.clone() - } - _ => None, - } - } - _ => None, - }; - if uses_modern_http && let Some(request_id) = cancellation_request_id { - if let Some(stream_ct) = request_stream_cancellations.remove(&request_id) { - stream_ct.cancel(); - } - pending_stream_response_ids.remove(&request_id); - let _ = responder.send(Ok(())); + pending_message = Some(send_request); + } + Event::ControlMessage(send_request) => { + if send_request.responder.is_closed() { continue; } + let cancellation_request_id = + Self::cancellation_request_id(&send_request.message); + if let Some(request_id) = &cancellation_request_id + && let Some(registration) = crate::service::remove_pending_request( + &mut request_stream_cancellations, + request_id, + ) + { + drop(registration); + } + if let Some(request_id) = cancellation_request_id + && !pending_stream_response_ids.remove(request_id) + && let Some(id) = request_id.numeric_string_value() + { + pending_stream_response_ids.remove(&RequestId::Number(id)); + } + let stale = send_request.control_generation() != context.control_generation(); + if stale || (uses_modern_http && cancellation_request_id.is_some()) { + // Local cancellation has already been signalled. Do not send an + // old cancellation or reply to a replacement session. + let result = if stale && cancellation_request_id.is_none() { + Err(StreamableHttpError::SessionExpired) + } else { + Ok(()) + }; + let _ = send_request.responder.send(result); + continue; + } + let (version, headers) = request_version_headers( + &protocol_headers, + &send_request.message, + &negotiated_version, + &tool_header_cache, + ); + control_posts.push(Self::post_request( + self.client.clone(), + &config, + send_request, + PostSession { + id: session_id.clone(), + headers, + version, + cancellation: session_cancellation.clone(), + }, + transport_task_ct.clone(), + )); + } + Event::RecoveryTimeout => { + recovery_deadline = None; + session_cancellation.cancel(); + tracing::warn!("old-session POSTs did not finish before the recovery deadline"); + } + Event::StartPost(send_request) => { + if send_request.responder.is_closed() + || send_request + .cancellation_token() + .is_some_and(|token| token.is_cancelled()) + { + let _ = send_request.responder.send(Ok(())); + continue; + } + let message = &send_request.message; let is_fallback_initialize = saved_init_request.is_none() && matches!( - &message, + message, ClientJsonRpcMessage::Request(request) if matches!( &request.request, @@ -1048,6 +1356,9 @@ impl Worker for StreamableHttpClientWorker { ); if is_fallback_initialize { saved_init_request = Some(message.clone()); + let WorkerSendRequest { + message, responder, .. + } = send_request; // Servers do not assign sessions to `server/discover`, so a // fallback initialize starts from a clean slate: no session // ID, no cleanup state, and no streams to tear down. @@ -1110,27 +1421,17 @@ impl Worker for StreamableHttpClientWorker { continue; } - let request_id = Self::client_request_id(&message); - let inline_version = match &message { + let barrier = Self::is_ordering_barrier(message, &negotiated_version); + debug_assert!(!barrier || (posts.is_empty() && control_posts.is_empty())); + let inline_version = match message { ClientJsonRpcMessage::Request(request) => { request.request.get_meta().protocol_version() } _ => None, }; - let is_initialized_notification = matches!( - &message, - ClientJsonRpcMessage::Notification(notification) - if matches!( - ¬ification.notification, - ClientNotification::InitializedNotification(_) - ) - ); - // Pass a clone to the first attempt so `message` is retained for a - // potential re-init retry. `post_message` takes ownership and the - // trait cannot be changed, so the clone is unavoidable. let (request_version, request_headers) = request_version_headers( &protocol_headers, - &message, + message, &negotiated_version, &tool_header_cache, ); @@ -1144,180 +1445,61 @@ impl Worker for StreamableHttpClientWorker { cleanup.protocol_headers = protocol_headers.clone(); } } - let response = self - .client - .post_message_with_max_sse_event_size( - config.uri.clone(), - message.clone(), - session_id.clone(), - config.auth_header.clone(), - request_headers, - config.max_sse_event_size, - ) - .await; - let send_result = match response { - Err(StreamableHttpError::SessionExpired) => { - if let Some(saved_init_request) = saved_init_request - .as_ref() - .filter(|_| config.reinit_on_expired_session) - { - // The server discarded the session (HTTP 404). Perform a - // fresh handshake once and replay the original message. - tracing::info!( - "session expired (HTTP 404), attempting transparent re-initialization" - ); - match Self::perform_reinitialization( - self.client.clone(), - saved_init_request.clone(), - config.uri.clone(), - config.auth_header.clone(), - config.custom_headers.clone(), - config.max_sse_event_size, - ) - .await - { - Ok(( - new_session_id, - new_negotiated_version, - new_protocol_headers, - )) => { - // Old streams hold the stale session ID. Stop them first - // so no late stale-session messages can arrive after the - // pending requests below are completed. - streams.abort_all(); - while streams.join_next().await.is_some() {} - - // Forward any already queued response messages and fail - // the remaining accepted requests so callers do not wait - // forever for responses that can no longer arrive. - Self::drain_queued_stream_messages( - &mut sse_worker_rx, - &mut context, - &mut pending_stream_response_ids, - ) - .await?; - Self::fail_pending_stream_responses( - &mut context, - &mut pending_stream_response_ids, - ) - .await?; - - session_id = new_session_id; - negotiated_version = new_negotiated_version; - protocol_headers = new_protocol_headers; - session_cleanup_info = - session_id.as_ref().map(|sid| SessionCleanupInfo { - client: self.client.clone(), - uri: config.uri.clone(), - session_id: sid.clone(), - auth_header: config.auth_header.clone(), - protocol_headers: protocol_headers.clone(), - }); - - if let Some(new_sid) = &session_id { - Self::spawn_common_stream( - &mut streams, - self.client.clone(), - new_sid.clone(), - &config, - protocol_headers.clone(), - sse_worker_tx.clone(), - transport_task_ct.clone(), - ); - } - - let (_, retry_headers) = request_version_headers( - &protocol_headers, - &message, - &negotiated_version, - &tool_header_cache, - ); - let retry_response = self - .client - .post_message_with_max_sse_event_size( - config.uri.clone(), - message, - session_id.clone(), - config.auth_header.clone(), - retry_headers, - config.max_sse_event_size, - ) - .await; - match retry_response { - Err(e) => Err(e), - Ok(StreamableHttpPostResponse::Accepted) => { - Self::mark_stream_response_pending( - &mut pending_stream_response_ids, - request_id, - ); - tracing::trace!( - "client message accepted after re-init" - ); - Ok(()) - } - Ok(StreamableHttpPostResponse::Json(mut msg, ..)) => { - cache_tools_from_response( - &mut tool_header_cache, - &mut msg, - &negotiated_version, - ); - context.send_to_handler(msg).await?; - Ok(()) - } - Ok(StreamableHttpPostResponse::Sse(stream, ..)) => { - let stream_request_id = request_id.clone(); - Self::mark_stream_response_pending( - &mut pending_stream_response_ids, - request_id, - ); - let sse_stream = Self::response_sse_to_jsonrpc( - stream, - session_id.clone(), - self.client.clone(), - config.uri.clone(), - config.auth_header.clone(), - protocol_headers.clone(), - config.max_sse_event_size, - self.config.retry_config.clone(), - ); - let stream_ct = transport_task_ct.child_token(); - if uses_modern_http - && let Some(request_id) = - stream_request_id.as_ref() - { - request_stream_cancellations.insert( - request_id.clone(), - stream_ct.clone(), - ); - } - let stream_tx = sse_worker_tx.clone(); - let origin = match &stream_request_id { - Some(id) => { - InboundStreamOrigin::OutboundRequest( - id.clone(), - ) - } - None => InboundStreamOrigin::Unassociated, - }; - streams.spawn(async move { - let result = Self::execute_sse_stream( - sse_stream, stream_tx, origin, true, - stream_ct, - ) - .await; - (stream_request_id, result) - }); - tracing::trace!("got new sse stream after re-init"); - Ok(()) - } - } - } - Err(reinit_err) => Err(reinit_err), - } - } else { - Err(StreamableHttpError::SessionExpired) - } + barrier_in_flight = barrier; + posts.push(Self::post_request( + self.client.clone(), + &config, + send_request, + PostSession { + id: session_id.clone(), + headers: request_headers, + version: request_version, + cancellation: session_cancellation.clone(), + }, + transport_task_ct.clone(), + )); + } + Event::PostResult(PostResult { + send_request, + response, + version, + }) => { + let is_control = Self::is_control_message(&send_request.message); + if !is_control { + // An ordering barrier runs without other ordinary POSTs. + barrier_in_flight = false; + } + let request_id = Self::client_request_id(&send_request.message); + let Some(response) = response else { + let _ = send_request.responder.send(Ok(())); + continue; + }; + if matches!(&response, Err(StreamableHttpError::SessionExpired)) + && !is_control + && !retrying_recovery + && config.reinit_on_expired_session + && saved_init_request.is_some() + { + if recovery_posts.is_empty() { + recovery_deadline = + Some(tokio::time::Instant::now() + config.session_recovery_timeout); } + recovery_posts.push_back(send_request); + continue; + } + let request_cancellation = send_request.cancellation_registration(); + let WorkerSendRequest { + message, responder, .. + } = send_request; + let is_initialized_notification = matches!( + &message, + ClientJsonRpcMessage::Notification(notification) + if matches!( + ¬ification.notification, + ClientNotification::InitializedNotification(_) + ) + ); + let send_result = match response { Err(e) => Err(e), Ok(StreamableHttpPostResponse::Accepted) => { Self::mark_stream_response_pending( @@ -1331,7 +1513,7 @@ impl Worker for StreamableHttpClientWorker { cache_tools_from_response( &mut tool_header_cache, &mut message, - &negotiated_version, + &version, ); context.send_to_handler(message).await?; Ok(()) @@ -1352,11 +1534,16 @@ impl Worker for StreamableHttpClientWorker { config.max_sse_event_size, self.config.retry_config.clone(), ); - let stream_ct = transport_task_ct.child_token(); - if uses_modern_http && let Some(request_id) = stream_request_id.as_ref() + // Keep the request cancellable until its response stream ends. + let stream_ct = request_cancellation + .as_ref() + .map(|registration| registration.token()) + .unwrap_or_else(|| transport_task_ct.child_token()); + if let (Some(request_id), Some(registration)) = + (stream_request_id.as_ref(), request_cancellation) { request_stream_cancellations - .insert(request_id.clone(), stream_ct.clone()); + .insert(request_id.clone(), registration); } let stream_tx = sse_worker_tx.clone(); let origin = match &stream_request_id { @@ -1395,12 +1582,12 @@ impl Worker for StreamableHttpClientWorker { } Event::ServerMessage(mut json_rpc_message) => { if let Some(response_id) = Self::server_response_id(&json_rpc_message) - && let Some(stream_ct) = crate::service::remove_pending_request( + && let Some(registration) = crate::service::remove_pending_request( &mut request_stream_cancellations, response_id, ) { - stream_ct.cancel(); + drop(registration); } Self::clear_stream_response_pending( &mut pending_stream_response_ids, @@ -1424,8 +1611,10 @@ impl Worker for StreamableHttpClientWorker { &mut pending_stream_response_ids, ) .await?; - request_stream_cancellations.remove(&request_id); - if pending_stream_response_ids.remove(&request_id) { + let cancelled = request_stream_cancellations + .remove(&request_id) + .is_some_and(|registration| registration.token().is_cancelled()); + if pending_stream_response_ids.remove(&request_id) && !cancelled { context .send_to_handler(ServerJsonRpcMessage::error( ErrorData::transport_closed( @@ -1446,6 +1635,14 @@ impl Worker for StreamableHttpClientWorker { } }; + // Stop outstanding http requests before deleting their session. + transport_task_ct.cancel(); + drop(posts); + drop(control_posts); + drop(pending_message); + drop(recovery_posts); + streams.abort_all(); + // Cleanup session before returning (ensures close() waits for session deletion) // Use a timeout to prevent indefinite hangs if the server is unresponsive if let Some(cleanup) = session_cleanup_info { @@ -1678,6 +1875,17 @@ pub struct StreamableHttpClientTransportConfig { pub uri: Arc, pub retry_config: Arc, pub channel_buffer_capacity: usize, + /// Maximum number of ordinary http POSTs in progress (default: 16). + /// A POST stops counting when it completes or opens an sse response stream. + /// Zero is treated as one. Cancellation and replies use a separate queue + /// with one extra POST slot. Each control POST has a five-second timeout + /// after it starts. + pub max_concurrent_requests: usize, + /// Maximum wait for old POSTs to finish before session recovery (default: five seconds). + /// The new initialization handshake has a separate timeout of the same length. + /// An unfinished old POST returns [`StreamableHttpError::SessionRecoveryTimeout`] + /// and is not retried because the server may have processed it. + pub session_recovery_timeout: Duration, /// if true, the transport will not require a session to be established pub allow_stateless: bool, /// The value to send in the authorization header @@ -1690,15 +1898,17 @@ pub struct StreamableHttpClientTransportConfig { /// [`StreamableHttpClient`] implementations must override the corresponding /// `*_with_max_sse_event_size` methods to enforce it. pub max_sse_event_size: usize, - /// Enables transparent recovery when the server reports an expired session (`HTTP 404`). + /// Automatically creates a new session when the server reports an expired + /// session (`http 404`). /// - /// When enabled, the transport performs one automatic recovery attempt: - /// 1. Replays the original `initialize` handshake to create a new session. - /// 2. Re-establishes streaming state for that session. - /// 3. Retries the in-flight request that failed with `SessionExpired`. + /// Ordinary POSTs that fail with `SessionExpired` in the same session share one + /// recovery attempt: + /// 1. Wait for old POSTs, up to [`Self::session_recovery_timeout`]. + /// 2. Repeat the original `initialize` handshake and open new streams. + /// 3. Retry each ordinary POST that failed with `SessionExpired` once. /// - /// This recovery is best-effort and bounded to a single attempt. If recovery fails, - /// the original failure path is preserved and the error is returned to the caller. + /// Control POSTs and other POST failures are not retried. If recovery or a retry + /// fails, the transport returns an error to the caller. pub reinit_on_expired_session: bool, } @@ -1710,6 +1920,18 @@ impl StreamableHttpClientTransportConfig { } } + /// Set how many ordinary POSTs can run at once. One keeps them serial; zero also means one. + pub fn max_concurrent_requests(mut self, limit: usize) -> Self { + self.max_concurrent_requests = limit.max(1); + self + } + + /// Set the separate timeouts for waiting for old POSTs and creating a replacement session. + pub fn session_recovery_timeout(mut self, timeout: Duration) -> Self { + self.session_recovery_timeout = timeout; + self + } + /// Set the authorization header to send with requests /// /// # Arguments @@ -1774,6 +1996,8 @@ impl Default for StreamableHttpClientTransportConfig { uri: "localhost".into(), retry_config: Arc::new(ExponentialBackoff::default()), channel_buffer_capacity: 16, + max_concurrent_requests: 16, + session_recovery_timeout: Duration::from_secs(5), allow_stateless: true, auth_header: None, custom_headers: HashMap::new(), diff --git a/crates/rmcp/src/transport/streamable_http_server/session/local.rs b/crates/rmcp/src/transport/streamable_http_server/session/local.rs index cc9e14893..e03e5b736 100644 --- a/crates/rmcp/src/transport/streamable_http_server/session/local.rs +++ b/crates/rmcp/src/transport/streamable_http_server/session/local.rs @@ -1127,7 +1127,9 @@ impl Worker for LocalSessionWorker { } }; match event { - InnerEvent::FromHandler(WorkerSendRequest { message, responder }) => { + InnerEvent::FromHandler(WorkerSendRequest { + message, responder, .. + }) => { // catch response let to_unregister = match &message { crate::model::JsonRpcMessage::Response(json_rpc_response) => { diff --git a/crates/rmcp/src/transport/worker.rs b/crates/rmcp/src/transport/worker.rs index 5294640e5..e7cac5f59 100644 --- a/crates/rmcp/src/transport/worker.rs +++ b/crates/rmcp/src/transport/worker.rs @@ -1,10 +1,21 @@ -use std::{borrow::Cow, time::Duration}; +use std::{ + borrow::Cow, + collections::HashMap, + sync::{ + Arc, Mutex, PoisonError, + atomic::{AtomicU64, Ordering}, + }, + time::Duration, +}; use tokio_util::sync::CancellationToken; use tracing::{Instrument, Level}; use super::{IntoTransport, Transport}; -use crate::service::{RxJsonRpcMessage, ServiceRole, TxJsonRpcMessage}; +use crate::{ + model::{CancelledNotification, JsonRpcMessage, RequestId}, + service::{RxJsonRpcMessage, ServiceRole, TxJsonRpcMessage}, +}; #[derive(Debug, thiserror::Error)] #[non_exhaustive] @@ -53,17 +64,100 @@ pub trait Worker: Sized + Send + 'static { fn config(&self) -> WorkerConfig { WorkerConfig::default() } + /// Return true to send this message through the separate control queue. + /// + /// Workers that opt in must read [`WorkerContext::control_from_handler_rx`] + /// and preserve any required ordering with ordinary messages. + fn is_control_message(_message: &TxJsonRpcMessage) -> bool { + false + } + /// Return true to register outgoing requests for cancellation before they enter a queue. + /// + /// Workers that opt in must honor [`WorkerSendRequest::cancellation_token`]. + fn supports_request_cancellation() -> bool { + false + } +} + +type RequestCancellations = Arc>>>; + +/// Keeps a request's cancellation token registered for a chosen lifetime. +pub(crate) struct RequestCancellationRegistration { + id: RequestId, + cancellation: Arc, + pending: RequestCancellations, +} + +impl RequestCancellationRegistration { + fn new(id: RequestId, token: CancellationToken, pending: RequestCancellations) -> Self { + let cancellation = Arc::new(token); + pending + .lock() + .unwrap_or_else(PoisonError::into_inner) + .insert(id.clone(), cancellation.clone()); + Self { + id, + cancellation, + pending, + } + } + + /// Return the token kept alive by this registration. + pub(crate) fn token(&self) -> CancellationToken { + self.cancellation.as_ref().clone() + } +} + +impl Drop for RequestCancellationRegistration { + fn drop(&mut self) { + self.cancellation.cancel(); + let mut pending = self.pending.lock().unwrap_or_else(PoisonError::into_inner); + if pending + .get(&self.id) + .is_some_and(|current| Arc::ptr_eq(current, &self.cancellation)) + { + pending.remove(&self.id); + } + } } #[non_exhaustive] pub struct WorkerSendRequest { pub message: TxJsonRpcMessage, pub responder: tokio::sync::oneshot::Sender>, + cancellation: Option>, + control_generation: u64, +} + +impl WorkerSendRequest { + /// Return the token registered before this request entered the send queue. + /// + /// This is present only for requests sent to a worker that enables + /// [`Worker::supports_request_cancellation`]. It is not sent over the wire. + pub fn cancellation_token(&self) -> Option { + self.cancellation + .as_deref() + .map(RequestCancellationRegistration::token) + } + + /// Keep the same cancellation registration alive after the POST completes. + #[cfg(feature = "transport-streamable-http-client")] + pub(crate) fn cancellation_registration(&self) -> Option> { + self.cancellation.clone() + } + + /// Return the control generation when the send was created. + pub fn control_generation(&self) -> u64 { + self.control_generation + } } pub struct WorkerTransport { rx: tokio::sync::mpsc::Receiver>, send_service: tokio::sync::mpsc::Sender>, + control_send_service: tokio::sync::mpsc::Sender>, + request_cancellations: RequestCancellations, + control_generation: Arc, join_handle: Option>>>, _drop_guard: tokio_util::sync::DropGuard, ct: CancellationToken, @@ -104,11 +198,17 @@ impl WorkerTransport { let worker_name = config.name; let (to_transport_tx, from_handler_rx) = tokio::sync::mpsc::channel::>(config.channel_buffer_capacity); + let (control_to_transport_tx, control_from_handler_rx) = + tokio::sync::mpsc::channel::>(config.channel_buffer_capacity); let (to_handler_tx, from_transport_rx) = tokio::sync::mpsc::channel::>(config.channel_buffer_capacity); + let request_cancellations = RequestCancellations::default(); + let control_generation = Arc::new(AtomicU64::new(0)); let context = WorkerContext { to_handler_tx, from_handler_rx, + control_from_handler_rx, + control_generation: control_generation.clone(), cancellation_token: transport_task_ct.clone(), }; @@ -142,6 +242,9 @@ impl WorkerTransport { Self { rx: from_transport_rx, send_service: to_transport_tx, + control_send_service: control_to_transport_tx, + request_cancellations, + control_generation, join_handle: Some(join_handle), ct: transport_task_ct.clone(), _drop_guard: transport_task_ct.drop_guard(), @@ -159,10 +262,25 @@ pub struct SendRequest { pub struct WorkerContext { pub to_handler_tx: tokio::sync::mpsc::Sender>, pub from_handler_rx: tokio::sync::mpsc::Receiver>, + /// Messages selected by [`Worker::is_control_message`]. + pub control_from_handler_rx: tokio::sync::mpsc::Receiver>, pub cancellation_token: CancellationToken, + control_generation: Arc, } impl WorkerContext { + /// Return the generation assigned to newly created control sends. + pub fn control_generation(&self) -> u64 { + self.control_generation.load(Ordering::SeqCst) + } + + /// Advance the generation so the worker can reject older control sends. + pub fn advance_control_generation(&self) -> u64 { + self.control_generation + .fetch_add(1, Ordering::SeqCst) + .wrapping_add(1) + } + pub async fn send_to_handler( &mut self, item: RxJsonRpcMessage, @@ -190,15 +308,62 @@ impl Transport for WorkerTransport { &mut self, item: TxJsonRpcMessage, ) -> impl Future> + Send + 'static { - let tx = self.send_service.clone(); + let control_generation = self.control_generation.load(Ordering::SeqCst); + let registration = if W::supports_request_cancellation() { + match &item { + JsonRpcMessage::Request(request) => { + Some(Arc::new(RequestCancellationRegistration::new( + request.id.clone(), + self.ct.child_token(), + self.request_cancellations.clone(), + ))) + } + JsonRpcMessage::Notification(notification) => { + let cancelled: Result = + notification.notification.clone().try_into(); + if let Ok(cancelled) = cancelled + && let Some(id) = cancelled.params.request_id.as_ref() + { + let pending = self + .request_cancellations + .lock() + .unwrap_or_else(PoisonError::into_inner); + if let Some(cancellation) = pending.get(id).or_else(|| { + id.numeric_string_value() + .and_then(|id| pending.get(&RequestId::Number(id))) + }) { + // Signal cancellation even if the control queue is full. + cancellation.cancel(); + } + } + None + } + _ => None, + } + } else { + None + }; + let tx = if W::is_control_message(&item) { + self.control_send_service.clone() + } else { + self.send_service.clone() + }; + let cancellation_guard = registration + .as_ref() + .map(|registration| registration.token().drop_guard()); let (responder, receiver) = tokio::sync::oneshot::channel(); let request = WorkerSendRequest { message: item, responder, + cancellation: registration, + control_generation, }; async move { tx.send(request).await.map_err(|_| W::err_closed())?; receiver.await.map_err(|_| W::err_closed())??; + if let Some(guard) = cancellation_guard { + let _ = guard.disarm(); + } Ok(()) } } diff --git a/crates/rmcp/tests/test_streamable_http_client_concurrency.rs b/crates/rmcp/tests/test_streamable_http_client_concurrency.rs new file mode 100644 index 000000000..0a2399039 --- /dev/null +++ b/crates/rmcp/tests/test_streamable_http_client_concurrency.rs @@ -0,0 +1,722 @@ +//! Independent http POSTs may overlap. A POST that reports an expired session +//! is retried at most once. +#![cfg(not(feature = "local"))] + +use std::{ + collections::HashMap, + io, + sync::{ + Arc, + atomic::{AtomicBool, AtomicUsize, Ordering::SeqCst}, + }, + time::Duration, +}; + +use futures::{StreamExt, stream::BoxStream}; +use http::{HeaderName, HeaderValue}; +use rmcp::{ + model::{ + CallToolRequestParams, CancelledNotificationParam, ClientInfo, ClientJsonRpcMessage, + ClientRequest, DiscoverResult, ProtocolVersion, Request, RequestId, RequestMetaObject, + ServerJsonRpcMessage, + }, + service::{ + ClientLifecycleMode, PeerRequestOptions, RequestHandle, RoleClient, RunningService, + serve_client_with_lifecycle, + }, + transport::streamable_http_client::{ + StreamableHttpClient, StreamableHttpClientTransport, StreamableHttpClientTransportConfig, + StreamableHttpError, StreamableHttpPostResponse, + }, +}; +use serde_json::{Value, json}; +use sse_stream::{Error as SseError, Sse}; +use tokio::{ + sync::{Mutex, mpsc, oneshot}, + task::JoinHandle, + time::timeout, +}; +use tokio_stream::wrappers::UnboundedReceiverStream; + +const TEST_TIMEOUT: Duration = Duration::from_secs(5); +type PostResult = Result>; +type Call = JoinHandle>; +type SseReceiver = mpsc::UnboundedReceiver>; + +#[derive(Default)] +struct Counts { + initialized: AtomicUsize, + hold_reinitialization: AtomicBool, + manual_controls: AtomicBool, + deleted: AtomicUsize, + cancelled: AtomicUsize, + posted: AtomicUsize, + active: AtomicUsize, + peak: AtomicUsize, +} + +struct ActivePost(Arc); + +impl Drop for ActivePost { + fn drop(&mut self) { + self.0.active.fetch_sub(1, SeqCst); + } +} + +struct Posted { + id: Value, + name: String, + session: Option>, + reply: oneshot::Sender, + returned: oneshot::Receiver<()>, +} + +struct ControlPost { + message: Value, + session: Option>, + reply: oneshot::Sender, +} + +fn response(id: Value, result: Value) -> ServerJsonRpcMessage { + serde_json::from_value(json!({ "jsonrpc": "2.0", "id": id, "result": result })) + .expect("valid scripted response") +} + +fn sse(message: Value) -> Result { + Ok(Sse { + event: Some("message".into()), + data: Some(message.to_string()), + id: None, + retry: None, + }) +} + +async fn next_event(receiver: &mut mpsc::UnboundedReceiver) -> T { + timeout(TEST_TIMEOUT, receiver.recv()) + .await + .expect("expected scripted event") + .expect("scripted client remains connected") +} + +impl Posted { + fn result(&self) -> ServerJsonRpcMessage { + response( + self.id.clone(), + json!({ "content": [{ "type": "text", "text": self.name }] }), + ) + } + + fn finish(self, result: PostResult) { + self.reply.send(result).expect("POST is still waiting"); + } + + fn succeed(self) { + let result = StreamableHttpPostResponse::Json(self.result(), None); + self.finish(Ok(result)); + } + + fn expire(self) { + self.finish(Err(StreamableHttpError::SessionExpired)); + } + + async fn finish_and_wait(self, result: PostResult) -> anyhow::Result<()> { + let Self { + reply, returned, .. + } = self; + reply.send(result).expect("POST is still waiting"); + timeout(TEST_TIMEOUT, returned).await??; + Ok(()) + } + + async fn expire_and_wait(self) -> anyhow::Result<()> { + self.finish_and_wait(Err(StreamableHttpError::SessionExpired)) + .await + } + + async fn start_sse(self) -> anyhow::Result> { + let message = serde_json::to_value(self.result()).unwrap(); + let (release, released) = oneshot::channel(); + let stream = futures::stream::once(async move { + released.await.expect("release the SSE response"); + sse(message) + }) + .boxed(); + self.finish_and_wait(Ok(StreamableHttpPostResponse::Sse(stream, None))) + .await?; + Ok(release) + } +} + +#[derive(Clone)] +struct ScriptedClient { + started: mpsc::UnboundedSender, + controls: mpsc::UnboundedSender, + incoming: Arc>>, + reinitializing: mpsc::UnboundedSender>, + counts: Arc, +} + +impl ScriptedClient { + async fn control_post(&self, message: Value, session: Option>) -> PostResult { + if !self.counts.manual_controls.load(SeqCst) { + return Ok(StreamableHttpPostResponse::Accepted); + } + let (reply, result) = oneshot::channel(); + self.controls + .send(ControlPost { + message, + session, + reply, + }) + .expect("test remains connected"); + result.await.expect("test answers the control POST") + } +} + +impl StreamableHttpClient for ScriptedClient { + type Error = io::Error; + + async fn post_message( + &self, + _uri: Arc, + message: ClientJsonRpcMessage, + session: Option>, + _auth_header: Option, + _custom_headers: HashMap, + ) -> PostResult { + let value = serde_json::to_value(message).unwrap(); + match value["method"].as_str() { + Some("server/discover") => Ok(StreamableHttpPostResponse::Json( + response( + value["id"].clone(), + serde_json::to_value(DiscoverResult::new( + vec![ProtocolVersion::V_2026_07_28], + serde_json::from_value(json!({ "tools": {} })).unwrap(), + )) + .unwrap(), + ), + None, + )), + Some("initialize") => { + let generation = self.counts.initialized.fetch_add(1, SeqCst) + 1; + if generation > 1 && self.counts.hold_reinitialization.load(SeqCst) { + let (release, released) = oneshot::channel(); + self.reinitializing + .send(release) + .expect("test remains connected"); + released.await.expect("test releases reinitialization"); + } + Ok(StreamableHttpPostResponse::Json( + response( + value["id"].clone(), + json!({ + "protocolVersion": "2025-11-25", + "capabilities": { "tools": {} }, + "serverInfo": { "name": "scripted", "version": "1" }, + }), + ), + Some(format!("session-{generation}")), + )) + } + Some("notifications/initialized") => Ok(StreamableHttpPostResponse::Accepted), + Some("notifications/cancelled") => { + self.counts.cancelled.fetch_add(1, SeqCst); + self.control_post(value, session).await + } + Some("tools/call") => { + self.counts.posted.fetch_add(1, SeqCst); + let active = self.counts.active.fetch_add(1, SeqCst) + 1; + self.counts.peak.fetch_max(active, SeqCst); + let _active = ActivePost(self.counts.clone()); + let (reply, response) = oneshot::channel(); + let (finished, returned) = oneshot::channel(); + self.started + .send(Posted { + id: value["id"].clone(), + name: value["params"]["name"].as_str().unwrap().to_owned(), + session, + reply, + returned, + }) + .expect("test remains connected"); + let response = response.await.expect("test answers each POST"); + let _ = finished.send(()); + response + } + None if value.get("result").is_some() || value.get("error").is_some() => { + self.control_post(value, session).await + } + method => panic!("unexpected scripted method: {method:?}"), + } + } + + async fn delete_session( + &self, + _uri: Arc, + _session: Arc, + _auth_header: Option, + _custom_headers: HashMap, + ) -> Result<(), StreamableHttpError> { + assert_eq!( + self.counts.active.load(SeqCst), + 0, + "POSTs must stop before deleting the session" + ); + self.counts.deleted.fetch_add(1, SeqCst); + Ok(()) + } + + async fn get_stream( + &self, + _uri: Arc, + _session: Option>, + _last_event_id: Option, + _auth_header: Option, + _custom_headers: HashMap, + ) -> Result>, StreamableHttpError> { + Ok(match self.incoming.lock().await.take() { + Some(incoming) => UnboundedReceiverStream::new(incoming).boxed(), + None => futures::stream::pending().boxed(), + }) + } +} + +struct Harness { + client: RunningService, + started: mpsc::UnboundedReceiver, + controls: mpsc::UnboundedReceiver, + incoming: mpsc::UnboundedSender>, + reinitializations: mpsc::UnboundedReceiver>, + counts: Arc, +} + +fn config() -> StreamableHttpClientTransportConfig { + StreamableHttpClientTransportConfig::with_uri("http://scripted/mcp") +} + +fn transport_error(error: &anyhow::Error) -> &StreamableHttpError { + let service_error = error + .downcast_ref::() + .expect("expected a service error"); + let rmcp::ServiceError::TransportSend(transport_error) = service_error else { + panic!("expected a transport error, got {service_error:?}"); + }; + transport_error + .error + .downcast_ref::>() + .expect("expected a streamable http error") +} + +fn assert_recovery_timeout(error: anyhow::Error) { + assert!(matches!( + transport_error(&error), + StreamableHttpError::SessionRecoveryTimeout + )); +} + +impl Harness { + async fn start(config: StreamableHttpClientTransportConfig) -> anyhow::Result { + Self::with_lifecycle(config, ClientLifecycleMode::Initialize).await + } + + async fn with_lifecycle( + config: StreamableHttpClientTransportConfig, + lifecycle: ClientLifecycleMode, + ) -> anyhow::Result { + let (started, requests) = mpsc::unbounded_channel(); + let (control_tx, controls) = mpsc::unbounded_channel(); + let (incoming, incoming_rx) = mpsc::unbounded_channel(); + let (reinitializing, reinitializations) = mpsc::unbounded_channel(); + let counts = Arc::new(Counts::default()); + let transport = StreamableHttpClientTransport::with_client( + ScriptedClient { + started, + controls: control_tx, + incoming: Arc::new(Mutex::new(Some(incoming_rx))), + reinitializing, + counts: counts.clone(), + }, + config, + ); + let client = + serve_client_with_lifecycle(ClientInfo::default(), transport, lifecycle).await?; + Ok(Self { + client, + started: requests, + controls, + incoming, + reinitializations, + counts, + }) + } + + fn call(&self, name: impl Into) -> Call { + let name = name.into(); + let peer = self.client.peer().clone(); + tokio::spawn(async move { + let result = peer + .call_tool(CallToolRequestParams::new(name.clone())) + .await?; + anyhow::ensure!(serde_json::to_value(result)?["content"][0]["text"] == name); + Ok(()) + }) + } + + async fn cancellable(&self, name: &'static str) -> anyhow::Result> { + self.cancellable_with_options(name, PeerRequestOptions::no_options()) + .await + } + + async fn cancellable_with_options( + &self, + name: &'static str, + options: PeerRequestOptions, + ) -> anyhow::Result> { + Ok(self + .client + .peer() + .send_cancellable_request( + ClientRequest::CallToolRequest(Request::new(CallToolRequestParams::new(name))), + options, + ) + .await?) + } + + fn notify_cancellation(&self, id: RequestId) -> JoinHandle> { + let peer = self.client.peer().clone(); + tokio::spawn(async move { + peer.notify_cancelled(CancelledNotificationParam::new(Some(id), None)) + .await + }) + } + + async fn next(&mut self) -> Posted { + next_event(&mut self.started).await + } + + async fn next_control(&mut self) -> ControlPost { + next_event(&mut self.controls).await + } + + async fn exchange_ping(&mut self, id: &str) { + self.incoming + .send(sse(json!({ "jsonrpc": "2.0", "id": id, "method": "ping" }))) + .expect("common SSE stream remains open"); + let control = self.next_control().await; + assert_eq!(control.message["id"], id); + assert!(control.message["result"].is_object()); + assert_eq!(control.session.as_deref(), Some("session-1")); + control + .reply + .send(Ok(StreamableHttpPostResponse::Accepted)) + .expect("reply POST is still waiting"); + } + + async fn finish( + self, + calls: Vec, + posted: usize, + initialized: usize, + ) -> anyhow::Result<()> { + for call in calls { + timeout(TEST_TIMEOUT, call).await???; + } + assert_eq!(self.counts.posted.load(SeqCst), posted); + assert_eq!(self.counts.initialized.load(SeqCst), initialized); + assert_eq!(self.counts.active.load(SeqCst), 0); + self.client.cancel().await?; + Ok(()) + } +} + +#[tokio::test] +async fn json_limits_allow_overlap_and_preserve_response_ids() -> anyhow::Result<()> { + let mut zero = config(); + zero.max_concurrent_requests = 0; + for (config, limit, total) in [ + (config().max_concurrent_requests(2), 2, 5), + (config().max_concurrent_requests(1), 1, 3), + (zero, 1, 2), + (config(), 16, 17), + ] { + let mut harness = Harness::start(config).await?; + let calls = (0..total) + .map(|index| harness.call(format!("request-{index}"))) + .collect(); + let mut pending = Vec::new(); + for _ in 0..limit { + pending.push(harness.next().await); + } + assert_eq!(harness.counts.active.load(SeqCst), limit); + // Keep the oldest response blocked while newer requests finish first. + for _ in limit..total { + pending.pop().unwrap().succeed(); + pending.push(harness.next().await); + } + for request in pending.into_iter().rev() { + request.succeed(); + } + let counts = harness.counts.clone(); + harness.finish(calls, total, 1).await?; + assert_eq!(counts.peak.load(SeqCst), limit); + } + Ok(()) +} + +#[tokio::test] +async fn early_sse_response_releases_the_post_slot() -> anyhow::Result<()> { + let mut harness = Harness::start(config().max_concurrent_requests(1)).await?; + let first = harness.call("first"); + let release = harness.next().await.start_sse().await?; + let second = harness.call("second"); + harness.next().await.succeed(); + timeout(TEST_TIMEOUT, second).await???; + assert!(!first.is_finished(), "the SSE response is still blocked"); + release.send(()).unwrap(); + assert_eq!(harness.counts.peak.load(SeqCst), 1); + harness.finish(vec![first], 2, 1).await +} + +#[tokio::test] +async fn concurrent_session_expiry_shares_one_reinitialization() -> anyhow::Result<()> { + let mut harness = Harness::start(config().max_concurrent_requests(2)).await?; + let calls = vec![harness.call("first"), harness.call("second")]; + // Hold both requests before releasing either expired response. + let first = harness.next().await; + let second = harness.next().await; + let mut originals = HashMap::new(); + for request in [first, second] { + assert_eq!(request.session.as_deref(), Some("session-1")); + originals.insert(request.name.clone(), request.id.clone()); + request.expire(); + } + for _ in 0..2 { + let retry = harness.next().await; + assert_eq!(retry.session.as_deref(), Some("session-2")); + assert_eq!(originals.remove(&retry.name), Some(retry.id.clone())); + retry.succeed(); + } + assert!(originals.is_empty()); + harness.finish(calls, 4, 2).await +} + +#[tokio::test] +async fn cancellation_still_runs_while_recovery_waits_for_old_posts() -> anyhow::Result<()> { + let mut harness = Harness::start(config().max_concurrent_requests(3)).await?; + let hanging = harness.cancellable("hanging").await?; + let mut blocked = harness.next().await; + let expired = harness.call("expired"); + harness.next().await.expire_and_wait().await?; + assert_eq!(harness.counts.initialized.load(SeqCst), 1); + assert_eq!(harness.counts.active.load(SeqCst), 1); + + timeout(TEST_TIMEOUT, hanging.cancel(None)) + .await + .expect("cancellation must bypass the session recovery wait")?; + timeout(TEST_TIMEOUT, blocked.reply.closed()).await?; + let retry = harness.next().await; + assert_eq!(retry.name, "expired"); + assert_eq!(retry.session.as_deref(), Some("session-2")); + retry.succeed(); + harness.finish(vec![expired], 3, 2).await +} + +#[tokio::test] +async fn server_replies_still_run_while_recovery_waits_for_old_posts() -> anyhow::Result<()> { + let mut harness = Harness::start(config().max_concurrent_requests(2)).await?; + harness.counts.manual_controls.store(true, SeqCst); + let waiting = harness.call("waiting-for-client"); + let blocked = harness.next().await; + let expired = harness.call("expired"); + harness.next().await.expire_and_wait().await?; + + harness.exchange_ping("recovery-ping").await; + blocked.succeed(); + let retry = harness.next().await; + assert_eq!(retry.name, "expired"); + assert_eq!(retry.session.as_deref(), Some("session-2")); + retry.succeed(); + harness.finish(vec![waiting, expired], 3, 2).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?; + harness.counts.manual_controls.store(true, SeqCst); + let mut meta = RequestMetaObject::new(); + meta.set_protocol_version(ProtocolVersion::V_2025_06_18); + let barrier = harness + .cancellable_with_options("barrier", PeerRequestOptions::no_options().with_meta(meta)) + .await?; + let blocked = harness.next().await; + let queued = harness.cancellable("after-barrier").await?; + + harness.exchange_ping("barrier-ping").await; + assert_eq!(harness.counts.posted.load(SeqCst), 1); + blocked.succeed(); + timeout(TEST_TIMEOUT, barrier.await_response()).await??; + let next = harness.next().await; + assert_eq!(next.name, "after-barrier"); + next.succeed(); + timeout(TEST_TIMEOUT, queued.await_response()).await??; + harness.finish(vec![], 2, 1).await +} + +#[tokio::test] +async fn recovery_deadline_drops_ambiguous_posts_without_retrying_them() -> anyhow::Result<()> { + assert_eq!(config().session_recovery_timeout, Duration::from_secs(5)); + let mut harness = Harness::start( + config() + .max_concurrent_requests(3) + .session_recovery_timeout(Duration::from_millis(50)), + ) + .await?; + let ambiguous = harness.call("possibly-applied"); + let mut blocked = harness.next().await; + let expired = harness.call("expired"); + let rejected = harness.next().await; + let rejected_id = rejected.id.clone(); + rejected.expire_and_wait().await?; + + assert_recovery_timeout(timeout(TEST_TIMEOUT, ambiguous).await??.unwrap_err()); + timeout(TEST_TIMEOUT, blocked.reply.closed()).await?; + let retry = harness.next().await; + assert_eq!(retry.name, "expired"); + assert_eq!(retry.id, rejected_id); + assert_eq!(retry.session.as_deref(), Some("session-2")); + retry.succeed(); + harness.finish(vec![expired], 3, 2).await +} + +#[tokio::test] +async fn reinitialization_has_its_own_deadline() -> anyhow::Result<()> { + let mut harness = + Harness::start(config().session_recovery_timeout(Duration::from_millis(50))).await?; + harness.counts.hold_reinitialization.store(true, SeqCst); + let expired = harness.call("expired"); + harness.next().await.expire_and_wait().await?; + let mut reinitialization = next_event(&mut harness.reinitializations).await; + + assert_recovery_timeout(timeout(TEST_TIMEOUT, expired).await??.unwrap_err()); + timeout(TEST_TIMEOUT, reinitialization.closed()).await?; + harness.finish(vec![], 1, 2).await +} + +#[tokio::test] +async fn an_expired_retry_is_not_retried_again() -> anyhow::Result<()> { + let mut harness = Harness::start(config().max_concurrent_requests(2)).await?; + let call = harness.call("expires-twice"); + let first = harness.next().await; + let id = first.id.clone(); + first.expire(); + let retry = harness.next().await; + assert_eq!(retry.id, id); + assert_eq!(retry.session.as_deref(), Some("session-2")); + retry.expire(); + let error = timeout(TEST_TIMEOUT, call).await??.unwrap_err(); + assert!(error.to_string().contains("Session expired")); + harness.finish(vec![], 2, 2).await +} + +#[tokio::test] +async fn a_lost_post_response_is_not_retried() -> anyhow::Result<()> { + let mut harness = Harness::start(config().max_concurrent_requests(2)).await?; + let call = harness.call("possibly-applied"); + harness + .next() + .await + .finish(Err(StreamableHttpError::Client(io::Error::new( + io::ErrorKind::ConnectionReset, + "scripted response lost", + )))); + let error = timeout(TEST_TIMEOUT, call).await??.unwrap_err(); + assert!(error.to_string().contains("scripted response lost")); + harness.finish(vec![], 1, 1).await +} + +#[tokio::test] +async fn cancellation_bypasses_queued_posts_at_capacity() -> anyhow::Result<()> { + for (lifecycle, legacy_notifications, initializations) in [ + (ClientLifecycleMode::Initialize, 2, 1), + ( + ClientLifecycleMode::Discover { + preferred_versions: vec![ProtocolVersion::V_2026_07_28], + }, + 0, + 0, + ), + ] { + let mut harness = + Harness::with_lifecycle(config().max_concurrent_requests(1), lifecycle).await?; + let request = harness.cancellable("cancel-me").await?; + let mut blocked = harness.next().await; + let cancelled = harness.cancellable("never-send").await?; + timeout(TEST_TIMEOUT, cancelled.cancel(None)).await??; + let queued = harness.cancellable("queued").await?; + timeout(TEST_TIMEOUT, request.cancel(None)).await??; + timeout(TEST_TIMEOUT, blocked.reply.closed()).await?; + let next = harness.next().await; + assert_eq!(next.name, "queued"); + next.succeed(); + timeout(TEST_TIMEOUT, queued.await_response()).await??; + assert_eq!(harness.counts.cancelled.load(SeqCst), legacy_notifications); + harness.finish(vec![], 2, initializations).await?; + } + Ok(()) +} + +#[tokio::test] +async fn a_hanging_legacy_control_does_not_delay_local_cancellation() -> anyhow::Result<()> { + let mut harness = Harness::start(config().max_concurrent_requests(1)).await?; + harness.counts.manual_controls.store(true, SeqCst); + let stale = harness.notify_cancellation(RequestId::Number(999)); + let mut held = harness.next_control().await; + assert_eq!(held.message["params"]["requestId"], 999); + harness.counts.manual_controls.store(false, SeqCst); + + let live = harness.cancellable("live-post").await?; + let mut blocked = harness.next().await; + let cancel_post = tokio::spawn(async move { live.cancel(None).await }); + timeout(Duration::from_secs(1), blocked.reply.closed()).await?; + + let streaming = harness.cancellable("live-stream").await?; + let mut stream = harness.next().await.start_sse().await?; + let cancel_stream = harness.notify_cancellation(streaming.id.clone()); + timeout(Duration::from_secs(1), stream.closed()).await?; + assert!( + !held.reply.is_closed(), + "the old control POST is still held" + ); + assert_eq!(harness.counts.cancelled.load(SeqCst), 1); + + // The private control timeout is five seconds; give its watchdog headroom. + let error = timeout(Duration::from_secs(10), stale).await??.unwrap_err(); + assert!(matches!( + transport_error(&error.into()), + StreamableHttpError::ControlRequestTimeout + )); + timeout(TEST_TIMEOUT, held.reply.closed()).await?; + timeout(TEST_TIMEOUT, cancel_post).await???; + timeout(TEST_TIMEOUT, cancel_stream).await???; + assert!(matches!( + timeout(TEST_TIMEOUT, streaming.await_response()).await?, + Err(rmcp::ServiceError::Cancelled { .. }) + )); + assert_eq!(harness.counts.cancelled.load(SeqCst), 3); + harness.finish(vec![], 2, 1).await +} + +#[tokio::test] +async fn close_drops_blocked_posts_before_deleting_the_session() -> anyhow::Result<()> { + let mut harness = Harness::start(config().max_concurrent_requests(2)).await?; + let calls = [harness.call("first"), harness.call("second")]; + let posts = [harness.next().await, harness.next().await]; + timeout(TEST_TIMEOUT, harness.client.cancel()).await??; + assert!(posts.iter().all(|post| post.reply.is_closed())); + assert_eq!(harness.counts.active.load(SeqCst), 0); + assert_eq!(harness.counts.deleted.load(SeqCst), 1); + for call in calls { + assert!(timeout(TEST_TIMEOUT, call).await??.is_err()); + } + Ok(()) +}