From 697c4977245573a319ec63ff2a9f64f7acb7f923 Mon Sep 17 00:00:00 2001 From: Nick Cooper <142361983+nickcoai@users.noreply.github.com> Date: Mon, 24 Aug 2026 09:43:35 -0400 Subject: [PATCH 01/12] fix: allow concurrent streamable http requests (#1186) * fix: allow concurrent streamable http requests * fix: keep streamable http recovery responsive Keep cancellation and replies available while old session POSTs finish. Bound the wait for old POSTs and the replacement initialization handshake. Do not retry interrupted POSTs because the server may have processed them. Add regressions for recovery, queued cancellation, control timeouts, and server replies needed by active requests. * fix: preserve response and cancellation ordering * refactor: clarify streamable http control flow * fix: match stream responses against pending request ids Match responses against all pending requests before removing a stream registration. Keep distinct numeric and string ids separate while preserving the existing fallback for servers that stringify numeric ids. Add a mixed-id regression and keep a separate registration owner alive in the abandoned-cancellation test. * feat: make streamable http control timeouts configurable --- README.md | 18 + crates/rmcp/Cargo.toml | 5 + .../src/transport/streamable_http_client.rs | 909 ++++++++----- .../streamable_http_server/session/local.rs | 4 +- crates/rmcp/src/transport/worker.rs | 330 ++++- ...test_streamable_http_client_concurrency.rs | 1133 +++++++++++++++++ 6 files changed, 2103 insertions(+), 296 deletions(-) create mode 100644 crates/rmcp/tests/test_streamable_http_client_concurrency.rs diff --git a/README.md b/README.md index 6ab22d67c..6e90827bc 100644 --- a/README.md +++ b/README.md @@ -1631,6 +1631,24 @@ 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. Configure their timeout with +`control_request_timeout` (default: five seconds). The timeout starts when the +POST starts, excluding time in the queue. Cancellation stops a queued or active +POST immediately. +For an open legacy response stream, the client stops reading but keeps the stream +alive until the cancellation send finishes or is dropped. This lets custom http +adapters handle cancellation before their stream state is removed. + +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 +reinitialization 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 852ac3b20..e554e8606 100644 --- a/crates/rmcp/Cargo.toml +++ b/crates/rmcp/Cargo.toml @@ -296,6 +296,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..a4275ab86 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,7 +30,10 @@ use crate::{ service::InboundStreamOrigin, transport::{ common::{client_side_sse::SseAutoReconnectStream, mcp_headers}, - worker::{Worker, WorkerQuitReason, WorkerSendRequest, WorkerTransport}, + worker::{ + RequestCancellationRegistration, Worker, WorkerQuitReason, WorkerSendRequest, + WorkerTransport, + }, }, }; @@ -207,6 +214,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 { @@ -325,6 +338,10 @@ impl StreamableHttpPostResponse { /// [`Self::post_message_with_max_sse_event_size`] and /// [`Self::get_stream_with_max_sse_event_size`] to enforce the transport's /// configured event-size limit. +/// +/// For legacy http, the transport keeps an open response stream alive until +/// its cancellation send finishes or is dropped. This lets a custom client +/// handle the cancellation using state owned by that stream. pub trait StreamableHttpClient: Clone + Send + 'static { type Error: std::error::Error + Send + Sync + 'static; fn post_message( @@ -475,6 +492,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 +526,85 @@ 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 control_request_timeout = config.control_request_timeout; + 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_request_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()), @@ -509,28 +620,16 @@ impl StreamableHttpClientWorker { } } - fn mark_stream_response_pending( - pending_stream_response_ids: &mut HashSet, - request_id: Option, - ) { - if let Some(request_id) = request_id { - pending_stream_response_ids.insert(request_id); - } - } - fn clear_stream_response_pending( pending_stream_response_ids: &mut HashSet, message: &ServerJsonRpcMessage, - ) { - let Some(response_id) = Self::server_response_id(message) else { - return; - }; - if pending_stream_response_ids.remove(response_id) { - return; - } - if let Some(id) = response_id.numeric_string_value() { - pending_stream_response_ids.remove(&RequestId::Number(id)); + ) -> Option { + let response_id = Self::server_response_id(message)?; + if let Some(id) = pending_stream_response_ids.take(response_id) { + return Some(id); } + let id = RequestId::Number(response_id.numeric_string_value()?); + pending_stream_response_ids.take(&id) } async fn drain_queued_stream_messages( @@ -541,7 +640,8 @@ impl StreamableHttpClientWorker { loop { match sse_worker_rx.try_recv() { Ok(message) => { - Self::clear_stream_response_pending(pending_stream_response_ids, &message); + let _ = + Self::clear_stream_response_pending(pending_stream_response_ids, &message); context.send_to_handler(message).await?; } Err(tokio::sync::mpsc::error::TryRecvError::Empty) => return Ok(()), @@ -573,6 +673,44 @@ impl StreamableHttpClientWorker { Ok(()) } + async fn fail_pending_responses_except_retries( + context: &mut super::worker::WorkerContext, + pending_stream_response_ids: &mut HashSet, + recovery_posts: &VecDeque>, + ) -> Result<(), WorkerQuitReason>> { + // Keep only retries that have not already received a stream response. + let retry_ids: Vec<_> = recovery_posts + .iter() + .filter_map(|request| Self::client_request_id(&request.message)) + .filter(|id| pending_stream_response_ids.remove(id)) + .collect(); + Self::fail_pending_stream_responses(context, pending_stream_response_ids).await?; + pending_stream_response_ids.extend(retry_ids); + Ok(()) + } + + fn fail_recovery_posts( + recovery_posts: &mut VecDeque>, + pending_stream_response_ids: &mut HashSet, + error: StreamableHttpError, + ) { + // The backend error cannot be cloned. Return it to one caller + // and return the original session-expired error to the others. + let mut recovery_error = Some(error); + for send_request in recovery_posts.drain(..) { + let pending = Self::client_request_id(&send_request.message) + .is_none_or(|id| pending_stream_response_ids.remove(&id)); + let result = if pending { + Err(recovery_error + .take() + .unwrap_or(StreamableHttpError::SessionExpired)) + } else { + Ok(()) + }; + let _ = send_request.responder.send(result); + } + } + /// Convert an SSE stream into JSON-RPC messages with reconnect semantics. /// /// This is used for request-scoped SSE responses as well as the standalone @@ -631,10 +769,34 @@ impl StreamableHttpClientWorker { .boxed() } + async fn run_response_stream( + mut sse_stream: BoxStream< + 'static, + Result>, + >, + sse_worker_tx: tokio::sync::mpsc::Sender, + origin: InboundStreamOrigin, + request_ct: CancellationToken, + stream_ct: CancellationToken, + uses_modern_http: bool, + ) -> Result<(), StreamableHttpError> { + tokio::select! { + biased; + _ = request_ct.cancelled(), if !uses_modern_http => { + // Stop reading, but keep the stream until the adapter + // handles cancellation or the send is dropped. + stream_ct.cancelled().await; + Ok(()) + } + result = Self::execute_sse_stream( + sse_stream.as_mut(), sse_worker_tx, origin, true, stream_ct.clone(), + ) => result, + } + } + async fn execute_sse_stream( sse_stream: impl Stream>> - + Send - + 'static, + + Send, sse_worker_tx: tokio::sync::mpsc::Sender, origin: InboundStreamOrigin, close_on_response: bool, @@ -819,6 +981,19 @@ impl StreamableHttpClientWorker { impl Worker for StreamableHttpClientWorker { type Role = RoleClient; type Error = StreamableHttpError; + fn is_control_message(message: &ClientJsonRpcMessage) -> bool { + match message { + ClientJsonRpcMessage::Response(_) | ClientJsonRpcMessage::Error(_) => true, + ClientJsonRpcMessage::Notification(notification) => matches!( + notification.notification, + ClientNotification::CancelledNotification(_) + ), + ClientJsonRpcMessage::Request(_) => false, + } + } + fn supports_request_cancellation() -> bool { + true + } fn err_closed() -> Self::Error { StreamableHttpError::TransportChannelClosed } @@ -844,6 +1019,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 +1129,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 +1166,159 @@ 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_responses_except_retries( + &mut context, + &mut pending_stream_response_ids, + &recovery_posts, + ) + .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(); + Self::fail_recovery_posts( + &mut recovery_posts, + &mut pending_stream_response_ids, + error, + ); + } + } + continue; + } + + let has_post_capacity = posts.len() < max_concurrent_requests; + let may_start = (retrying_recovery || recovery_posts.is_empty()) + && !barrier_in_flight + && has_post_capacity; + let may_receive = may_start && pending_message.is_none() && !retrying_recovery; + let queued = if retrying_recovery { + recovery_posts.front_mut() + } else { + pending_message.as_mut() + }; + let has_queued = queued.is_some(); + let can_process_queued = queued.as_ref().is_some_and(|request| { + let retry_completed = retrying_recovery + && Self::client_request_id(&request.message) + .is_some_and(|id| !pending_stream_response_ids.contains(&id)); + let ordering_satisfied = + !Self::is_ordering_barrier(&request.message, &negotiated_version) + || (posts.is_empty() && control_posts.is_empty()); + retry_completed || (may_start && ordering_satisfied) + }); let event = tokio::select! { + _ = async { + if can_process_queued { + return; + } + let request = queued.expect("a POST is queued"); + let cancellation = request.cancellation_token().unwrap_or_default(); + tokio::select! { + _ = request.responder.closed() => {} + _ = cancellation.cancelled() => {} + } + }, if has_queued => { + 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_receive => { + match message { + 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 { - Ok(msg) => Event::ClientMessage(msg), - Err(e) => break 'main_loop Err(e), + 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 +1328,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, - } + 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); + let stale = send_request.control_generation() != context.control_generation(); + if stale { + // Do not send old controls to a replacement session. + let result = match cancellation_request_id { + Some(_) => Ok(()), + None => Err(StreamableHttpError::SessionExpired), + }; + let _ = send_request.responder.send(result); + continue; + } + if let Some(request_id) = cancellation_request_id { + drop(request_stream_cancellations.remove(request_id)); + pending_stream_response_ids.remove(request_id); + if uses_modern_http { + let _ = send_request.responder.send(Ok(())); + continue; } - _ => 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(); + } + 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) => { + let request_id = Self::client_request_id(&send_request.message); + let send_cancelled = send_request.responder.is_closed() + || send_request + .cancellation_token() + .is_some_and(|token| token.is_cancelled()); + let retry_completed = retrying_recovery + && request_id + .as_ref() + .is_some_and(|id| !pending_stream_response_ids.contains(id)); + if send_cancelled || retry_completed { + if retrying_recovery && let Some(id) = &request_id { + pending_stream_response_ids.remove(id); } - pending_stream_response_ids.remove(&request_id); - let _ = responder.send(Ok(())); + 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 +1420,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 +1485,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,186 +1509,87 @@ 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; + // The common stream can return a response before this POST finishes. + if let Some(request_id) = request_id { + pending_stream_response_ids.insert(request_id); + } + 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); + if request_id + .as_ref() + .is_some_and(|id| !pending_stream_response_ids.contains(id)) + { + // A stream response or cancellation already completed this request. + let _ = send_request.responder.send(Ok(())); + continue; + } + let will_retry = + matches!(&response, Some(Err(StreamableHttpError::SessionExpired))) + && !is_control + && !retrying_recovery + && config.reinit_on_expired_session + && saved_init_request.is_some(); + let awaits_stream_response = matches!( + &response, + Some(Ok(StreamableHttpPostResponse::Accepted + | StreamableHttpPostResponse::Sse(..))) + ); + if !will_retry + && !awaits_stream_response + && let Some(id) = &request_id + { + pending_stream_response_ids.remove(id); + } + let Some(response) = response else { + let _ = send_request.responder.send(Ok(())); + continue; + }; + if will_retry { + 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( - &mut pending_stream_response_ids, - request_id, - ); tracing::trace!("client message accepted"); Ok(()) } @@ -1331,17 +1597,13 @@ impl Worker for StreamableHttpClientWorker { cache_tools_from_response( &mut tool_header_cache, &mut message, - &negotiated_version, + &version, ); context.send_to_handler(message).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 stream_request_id = request_id; let sse_stream = Self::response_sse_to_jsonrpc( stream, session_id.clone(), @@ -1352,11 +1614,24 @@ 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() + let request_ct = request_cancellation + .as_ref() + .map(|registration| registration.token()) + .unwrap_or_else(|| transport_task_ct.child_token()); + // A legacy adapter may need the open stream to handle cancellation. + let stream_ct = if uses_modern_http { + request_ct.clone() + } else { + request_cancellation + .as_ref() + .map(|registration| registration.lifetime_token()) + .unwrap_or_else(|| request_ct.clone()) + }; + 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 { @@ -1364,8 +1639,13 @@ impl Worker for StreamableHttpClientWorker { None => InboundStreamOrigin::Unassociated, }; streams.spawn(async move { - let result = Self::execute_sse_stream( - sse_stream, stream_tx, origin, true, stream_ct, + let result = Self::run_response_stream( + sse_stream, + stream_tx, + origin, + request_ct, + stream_ct, + uses_modern_http, ) .await; (stream_request_id, result) @@ -1394,18 +1674,13 @@ impl Worker for StreamableHttpClientWorker { let _ = responder.send(send_result); } 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( - &mut request_stream_cancellations, - response_id, - ) - { - stream_ct.cancel(); - } - Self::clear_stream_response_pending( + // Match against all pending requests, not just open response streams. + if let Some(request_id) = Self::clear_stream_response_pending( &mut pending_stream_response_ids, &json_rpc_message, - ); + ) { + drop(request_stream_cancellations.remove(&request_id)); + } cache_tools_from_response( &mut tool_header_cache, &mut json_rpc_message, @@ -1424,8 +1699,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 +1723,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 +1963,19 @@ 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 and use [`Self::control_request_timeout`]. + pub max_concurrent_requests: usize, + /// Maximum time a cancellation or reply POST can run after it starts (default: five seconds). + /// Time spent waiting in the control queue does not count toward this timeout. + pub control_request_timeout: Duration, + /// 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 +1988,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 +2010,24 @@ 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 timeout for cancellation and reply POSTs, starting when each POST starts. + pub fn control_request_timeout(mut self, timeout: Duration) -> Self { + self.control_request_timeout = timeout; + 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 +2092,9 @@ impl Default for StreamableHttpClientTransportConfig { uri: "localhost".into(), retry_config: Arc::new(ExponentialBackoff::default()), channel_buffer_capacity: 16, + max_concurrent_requests: 16, + control_request_timeout: Duration::from_secs(5), + session_recovery_timeout: Duration::from_secs(5), allow_stateless: true, auth_header: None, custom_headers: HashMap::new(), @@ -2158,11 +2479,13 @@ mod tests { NumberOrString::String("1".into()), ); - StreamableHttpClientWorker::::clear_stream_response_pending( - &mut pending, - &response, - ); + let matched_id = + StreamableHttpClientWorker::::clear_stream_response_pending( + &mut pending, + &response, + ); + assert_eq!(matched_id, Some(NumberOrString::Number(1))); assert!(pending.is_empty()); } @@ -2173,14 +2496,16 @@ mod tests { let mut pending = HashSet::from([NumberOrString::Number(1), string_id.clone()]); let response = ServerJsonRpcMessage::response( ServerResult::ListToolsResult(ListToolsResult::default()), - string_id, + string_id.clone(), ); - StreamableHttpClientWorker::::clear_stream_response_pending( - &mut pending, - &response, - ); + let matched_id = + StreamableHttpClientWorker::::clear_stream_response_pending( + &mut pending, + &response, + ); + assert_eq!(matched_id, Some(string_id)); assert_eq!(pending, HashSet::from([NumberOrString::Number(1)])); } } 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..e32da7d7d 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, Weak, + 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,114 @@ 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, + lifetime: CancellationToken, + cancellation: CancellationToken, + pending: RequestCancellations, +} + +impl RequestCancellationRegistration { + fn new(id: RequestId, token: CancellationToken, pending: RequestCancellations) -> Arc { + let registration = Arc::new(Self { + id: id.clone(), + cancellation: token.child_token(), + lifetime: token, + pending, + }); + registration + .pending + .lock() + .unwrap_or_else(PoisonError::into_inner) + .insert(id, Arc::downgrade(®istration)); + registration + } + + /// Return the token kept alive by this registration. + pub(crate) fn token(&self) -> CancellationToken { + self.cancellation.clone() + } + + /// Return the token cancelled when the request lifetime ends. + pub(crate) fn lifetime_token(&self) -> CancellationToken { + self.lifetime.clone() + } +} + +impl Drop for RequestCancellationRegistration { + fn drop(&mut self) { + self.lifetime.cancel(); + let mut pending = self.pending.lock().unwrap_or_else(PoisonError::into_inner); + if pending + .get(&self.id) + .is_some_and(|current| std::ptr::eq(current.as_ptr(), self)) + { + 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. + /// Keep this request alive while its work is active. Cloning the token does + /// not keep its cancellation registration alive. + pub fn cancellation_token(&self) -> Option { + self.cancellation + .as_deref() + .map(RequestCancellationRegistration::token) + } + + /// Keep the same cancellation registration alive after the send completes. + #[cfg(feature = "transport-streamable-http-client")] + pub(crate) fn cancellation_registration(&self) -> Option> { + self.cancellation.clone() + } + + /// Return the local generation captured when [`Transport::send`] created its future. + /// + /// 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. + 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 +212,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,11 +256,32 @@ 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(), } } + + fn cancel_request_from_notification( + &self, + notification: &::Not, + ) -> Option> { + let cancelled: CancelledNotification = notification.clone().try_into().ok()?; + let id = cancelled.params.request_id.as_ref()?; + let target = { + let pending = self + .request_cancellations + .lock() + .unwrap_or_else(PoisonError::into_inner); + pending.get(id).and_then(Weak::upgrade) + }?; + // Signal cancellation even if the control queue is full. + target.token().cancel(); + Some(target) + } } #[non_exhaustive] @@ -159,10 +294,32 @@ 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 local generation that newly created sends will capture. + /// + /// Workers may use this value to check messages from an earlier connection + /// or session. The generic transport does not check it 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. + /// 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 { + self.control_generation + .fetch_add(1, Ordering::SeqCst) + .wrapping_add(1) + } + pub async fn send_to_handler( &mut self, item: RxJsonRpcMessage, @@ -190,15 +347,52 @@ 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 mut cancellation_target = None; + let registration = if W::supports_request_cancellation() { + match &item { + JsonRpcMessage::Request(request) => Some(RequestCancellationRegistration::new( + request.id.clone(), + self.ct.child_token(), + self.request_cancellations.clone(), + )), + JsonRpcMessage::Notification(notification) => { + cancellation_target = + self.cancel_request_from_notification(¬ification.notification); + 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.lifetime_token().drop_guard()); + let target_guard = cancellation_target + .as_ref() + .map(|target| target.lifetime_token().drop_guard()); let (responder, receiver) = tokio::sync::oneshot::channel(); let request = WorkerSendRequest { message: item, responder, + cancellation: registration, + control_generation, }; async move { + // Keep the stream alive until its cancellation is handled or abandoned. + let _cancellation_target = cancellation_target; + let _target_guard = target_guard; 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(()) } } @@ -215,3 +409,133 @@ impl Transport for WorkerTransport { } } } + +#[cfg(all(test, feature = "client"))] +mod tests { + use std::io; + + use super::*; + use crate::{model::ClientJsonRpcMessage, service::RoleClient}; + + struct TestWorker(tokio::sync::oneshot::Sender>); + + impl Worker for TestWorker { + type Error = io::Error; + type Role = RoleClient; + + fn err_closed() -> Self::Error { + io::Error::other("worker closed") + } + + fn err_join(error: tokio::task::JoinError) -> Self::Error { + io::Error::other(error) + } + + async fn run( + self, + context: WorkerContext, + ) -> Result<(), WorkerQuitReason> { + let cancellation = context.cancellation_token.clone(); + self.0 + .send(context) + .map_err(|_| WorkerQuitReason::HandlerTerminated)?; + cancellation.cancelled().await; + Ok(()) + } + + fn is_control_message(message: &ClientJsonRpcMessage) -> bool { + matches!(message, JsonRpcMessage::Notification(_)) + } + + fn supports_request_cancellation() -> bool { + true + } + } + + fn cancellation_message(id: RequestId) -> ClientJsonRpcMessage { + serde_json::from_value(serde_json::json!({ + "jsonrpc": "2.0", + "method": "notifications/cancelled", + "params": { "requestId": id }, + })) + .unwrap() + } + + #[tokio::test] + async fn cancellation_matches_request_id_exactly() { + let (context_tx, context_rx) = tokio::sync::oneshot::channel(); + let mut transport = WorkerTransport::spawn(TestWorker(context_tx)); + let _context = context_rx.await.unwrap(); + let registration = RequestCancellationRegistration::new( + RequestId::Number(7), + CancellationToken::new(), + transport.request_cancellations.clone(), + ); + + drop(transport.send(cancellation_message(RequestId::String("7".into())))); + assert!(!registration.token().is_cancelled()); + + drop(transport.send(cancellation_message(RequestId::Number(7)))); + assert!(registration.token().is_cancelled()); + transport.close().await.unwrap(); + } + + #[tokio::test] + async fn abandoned_cancellation_send_ends_request_lifetime() { + 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 admitted in [false, true] { + let id = RequestId::Number(7); + let registration = RequestCancellationRegistration::new( + id.clone(), + CancellationToken::new(), + transport.request_cancellations.clone(), + ); + let lifetime = registration.lifetime_token(); + let weak = Arc::downgrade(®istration); + let mut send = Box::pin(transport.send(cancellation_message(id))); + assert!(registration.token().is_cancelled()); + assert!(!lifetime.is_cancelled()); + + let queued = if admitted { + assert!(futures::poll!(send.as_mut()).is_pending()); + Some(context.control_from_handler_rx.recv().await.unwrap()) + } else { + None + }; + drop(send); + + // An open stream can still own the registration after cancellation is abandoned. + assert!(lifetime.is_cancelled()); + assert!(weak.upgrade().is_some()); + drop(registration); + assert!(weak.upgrade().is_none()); + assert!(queued.is_none_or(|request| request.responder.is_closed())); + assert!(transport.request_cancellations.lock().unwrap().is_empty()); + } + transport.close().await.unwrap(); + } + + #[test] + fn dropping_an_old_registration_preserves_a_reused_id() { + let pending = RequestCancellations::default(); + let id = RequestId::Number(7); + let old = RequestCancellationRegistration::new( + id.clone(), + CancellationToken::new(), + pending.clone(), + ); + let current = RequestCancellationRegistration::new( + id.clone(), + CancellationToken::new(), + pending.clone(), + ); + drop(old); + let registered = pending.lock().unwrap().get(&id).cloned().unwrap(); + assert!(Weak::ptr_eq(®istered, &Arc::downgrade(¤t))); + drop(current); + assert!(pending.lock().unwrap().is_empty()); + } +} 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..20b725831 --- /dev/null +++ b/crates/rmcp/tests/test_streamable_http_client_concurrency.rs @@ -0,0 +1,1133 @@ +//! 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, Mutex as StdMutex, + 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; +use tokio_util::sync::CancellationToken; + +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, + reject_unmatched_cancellation: AtomicBool, + local_streams: StdMutex>, + local_cancellations: AtomicUsize, + 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 LocalStreamRegistration { + id: RequestId, + counts: Arc, + dropped: Option>, +} + +impl Drop for LocalStreamRegistration { + fn drop(&mut self) { + self.counts.local_streams.lock().unwrap().remove(&self.id); + let _ = self.dropped.take().unwrap().send(()); + } +} + +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) + } + + async fn start_local_sse( + self, + counts: Arc, + ) -> anyhow::Result<(CancellationToken, oneshot::Receiver<()>)> { + let id: RequestId = serde_json::from_value(self.id.clone())?; + let cancellation = CancellationToken::new(); + let eof = cancellation.clone(); + counts + .local_streams + .lock() + .unwrap() + .insert(id.clone(), cancellation.clone()); + let (dropped, closed) = oneshot::channel(); + let registration = LocalStreamRegistration { + id, + counts, + dropped: Some(dropped), + }; + let (started, listening) = oneshot::channel(); + let stream = futures::stream::once(async move { + let _registration = registration; + let _ = started.send(()); + cancellation.cancelled().await; + None::> + }) + .filter_map(futures::future::ready) + .boxed(); + self.finish_and_wait(Ok(StreamableHttpPostResponse::Sse(stream, None))) + .await?; + timeout(TEST_TIMEOUT, listening).await??; + Ok((eof, closed)) + } +} + +#[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") => { + let id: RequestId = + serde_json::from_value(value["params"]["requestId"].clone()).unwrap(); + let local = self.counts.local_streams.lock().unwrap().get(&id).cloned(); + if let Some(cancellation) = local { + self.counts.local_cancellations.fetch_add(1, SeqCst); + cancellation.cancel(); + return Ok(StreamableHttpPostResponse::Accepted); + } + self.counts.cancelled.fetch_add(1, SeqCst); + if self.counts.reject_unmatched_cancellation.load(SeqCst) { + return Err(StreamableHttpError::Client(io::Error::other( + "local event cancellation reached the http POST path", + ))); + } + 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 a_common_stream_response_finishes_the_request_before_its_post_returns() +-> anyhow::Result<()> { + let (stream, incoming) = mpsc::unbounded_channel(); + for (late, orphan) in [ + (Ok(StreamableHttpPostResponse::Accepted), None), + (Err(StreamableHttpError::SessionExpired), None), + ( + Ok(StreamableHttpPostResponse::Sse( + UnboundedReceiverStream::new(incoming).boxed(), + None, + )), + Some(stream), + ), + ] { + let mut harness = Harness::start(config().max_concurrent_requests(1)).await?; + let request = harness.cancellable("completed-on-common-stream").await?; + let post = harness.next().await; + harness + .incoming + .send(sse(serde_json::to_value(post.result())?))?; + timeout(TEST_TIMEOUT, request.await_response()).await??; + + post.finish_and_wait(late).await?; + let next = harness.call("after-completed-request"); + let post = harness.next().await; + assert_eq!( + post.name, "after-completed-request", + "do not retry completed work" + ); + assert_eq!(post.session.as_deref(), Some("session-1")); + post.succeed(); + timeout(TEST_TIMEOUT, next).await???; + if let Some(stream) = orphan { + timeout(Duration::from_secs(1), stream.closed()) + .await + .expect("a late SSE response must not leave an orphan stream"); + } + harness.finish(vec![], 2, 1).await?; + } + Ok(()) +} + +#[tokio::test] +async fn common_stream_response_keeps_numeric_and_string_ids_distinct() -> anyhow::Result<()> { + use rmcp::transport::Transport; + + let (started, mut requests) = mpsc::unbounded_channel(); + let (incoming, incoming_rx) = mpsc::unbounded_channel(); + let counts = Arc::new(Counts::default()); + let mut transport = StreamableHttpClientTransport::with_client( + ScriptedClient { + started, + controls: mpsc::unbounded_channel().0, + incoming: Arc::new(Mutex::new(Some(incoming_rx))), + reinitializing: mpsc::unbounded_channel().0, + counts: counts.clone(), + }, + config(), + ); + for message in [ + json!({ "jsonrpc": "2.0", "id": 0, "method": "initialize", "params": ClientInfo::default() }), + json!({ "jsonrpc": "2.0", "method": "notifications/initialized" }), + ] { + timeout( + TEST_TIMEOUT, + transport.send(serde_json::from_value(message)?), + ) + .await??; + } + assert!(timeout(TEST_TIMEOUT, transport.receive()).await?.is_some()); + let call = |id, name| { + ClientJsonRpcMessage::request( + ClientRequest::CallToolRequest(Request::new(CallToolRequestParams::new(name))), + id, + ) + }; + + let mut numeric_send = Box::pin(transport.send(call(RequestId::Number(7), "numeric"))); + assert!(futures::poll!(numeric_send.as_mut()).is_pending()); + let numeric_post = next_event(&mut requests).await; + let numeric_result = numeric_post.result(); + let release_numeric = numeric_post.start_sse().await?; + timeout(TEST_TIMEOUT, numeric_send).await??; + + // This exact string id is pending, but has no response stream registration. + let mut string_send = Box::pin(transport.send(call(RequestId::String("7".into()), "string"))); + assert!(futures::poll!(string_send.as_mut()).is_pending()); + let string_post = next_event(&mut requests).await; + let string_result = serde_json::to_value(string_post.result())?; + string_post.finish(Ok(StreamableHttpPostResponse::Accepted)); + timeout(TEST_TIMEOUT, string_send).await??; + + incoming.send(sse(string_result.clone()))?; + let received = timeout(TEST_TIMEOUT, transport.receive()) + .await? + .expect("string-id response"); + assert_eq!(serde_json::to_value(received)?, string_result); + + release_numeric + .send(()) + .expect("the distinct numeric-id stream must remain open"); + let received = timeout(TEST_TIMEOUT, transport.receive()) + .await? + .expect("numeric-id response"); + assert_eq!( + serde_json::to_value(received)?, + serde_json::to_value(numeric_result)? + ); + assert_eq!(counts.posted.load(SeqCst), 2); + assert_eq!(counts.cancelled.load(SeqCst), 0); + transport.close().await?; + Ok(()) +} + +#[tokio::test] +async fn a_common_stream_response_prevents_replay_during_session_recovery() -> anyhow::Result<()> { + let mut harness = Harness::start(config().max_concurrent_requests(2)).await?; + let blocking = harness.call("blocking-recovery"); + let blocked = harness.next().await; + let request = harness.cancellable("completed-during-recovery").await?; + let expired = harness.next().await; + let final_response = sse(serde_json::to_value(expired.result())?); + expired.expire_and_wait().await?; + harness.incoming.send(final_response)?; + timeout(TEST_TIMEOUT, request.await_response()).await??; + blocked.succeed(); + timeout(TEST_TIMEOUT, blocking).await???; + + let next = harness.call("after-completed-request"); + let post = harness.next().await; + assert_eq!( + post.name, "after-completed-request", + "do not retry completed work" + ); + let initializations = harness.counts.initialized.load(SeqCst); + assert!((1..=2).contains(&initializations)); + post.succeed(); + timeout(TEST_TIMEOUT, next).await???; + harness.finish(vec![], 3, initializations).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(Duration::from_secs(1), 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 a_cancelled_version_barrier_does_not_block_unrelated_posts() -> anyhow::Result<()> { + let mut harness = Harness::start(config().max_concurrent_requests(2)).await?; + let blocking = harness.call("blocking"); + let blocked = harness.next().await; + let mut meta = RequestMetaObject::new(); + meta.set_protocol_version(ProtocolVersion::V_2025_06_18); + let barrier = harness + .cancellable_with_options( + "cancelled-barrier", + PeerRequestOptions::no_options().with_meta(meta), + ) + .await?; + timeout(TEST_TIMEOUT, barrier.cancel(None)).await??; + + let next = harness.call("after-cancelled-barrier"); + let post = harness.next().await; + assert_eq!(post.name, "after-cancelled-barrier"); + post.succeed(); + timeout(TEST_TIMEOUT, next).await???; + assert!(!blocked.reply.is_closed(), "the older POST is still held"); + blocked.succeed(); + harness.finish(vec![blocking], 2, 1).await +} + +#[tokio::test] +async fn dropping_a_parked_barrier_send_unblocks_the_ordinary_queue() -> anyhow::Result<()> { + use std::task::{Context, Wake, Waker}; + + use rmcp::transport::Transport; + + struct WakeSignal(mpsc::UnboundedSender<()>); + impl Wake for WakeSignal { + fn wake(self: Arc) { + let _ = self.0.send(()); + } + } + + let (started, mut requests) = mpsc::unbounded_channel(); + let counts = Arc::new(Counts::default()); + let mut config = config().max_concurrent_requests(2); + config.channel_buffer_capacity = 1; + let mut transport = StreamableHttpClientTransport::with_client( + ScriptedClient { + started, + controls: mpsc::unbounded_channel().0, + incoming: Arc::new(Mutex::new(None)), + reinitializing: mpsc::unbounded_channel().0, + counts: counts.clone(), + }, + config, + ); + for message in [ + json!({ "jsonrpc": "2.0", "id": 0, "method": "initialize", "params": ClientInfo::default() }), + json!({ "jsonrpc": "2.0", "method": "notifications/initialized" }), + ] { + timeout( + TEST_TIMEOUT, + transport.send(serde_json::from_value(message)?), + ) + .await??; + } + assert!(timeout(TEST_TIMEOUT, transport.receive()).await?.is_some()); + let call = |id, params| { + ClientJsonRpcMessage::request( + ClientRequest::CallToolRequest(Request::new(params)), + RequestId::Number(id), + ) + }; + let mut first = Box::pin(transport.send(call(1, CallToolRequestParams::new("held")))); + assert!(futures::poll!(first.as_mut()).is_pending()); + let blocked = next_event(&mut requests).await; + + let mut meta = RequestMetaObject::new(); + meta.set_protocol_version(ProtocolVersion::V_2025_06_18); + let mut barrier_request = Request::new(CallToolRequestParams::new("abandoned-barrier")); + barrier_request.extensions.insert(meta); + let mut barrier = Box::pin(transport.send(ClientJsonRpcMessage::request( + ClientRequest::CallToolRequest(barrier_request), + RequestId::Number(2), + ))); + assert!(futures::poll!(barrier.as_mut()).is_pending()); + let mut next = Box::pin(transport.send(call(3, CallToolRequestParams::new("after-barrier")))); + let (wake, mut woke) = mpsc::unbounded_channel(); + let waker = Waker::from(Arc::new(WakeSignal(wake))); + assert!( + next.as_mut() + .poll(&mut Context::from_waker(&waker)) + .is_pending() + ); + // The barrier fills the queue. The next send wakes only after the worker parks it. + next_event(&mut woke).await; + assert_eq!(counts.posted.load(SeqCst), 1, "the barrier must be parked"); + drop(barrier); + assert!(futures::poll!(next.as_mut()).is_pending()); + + let post = timeout(Duration::from_secs(1), requests.recv()) + .await + .expect("dropping the parked send must wake the worker") + .unwrap(); + assert_eq!(post.name, "after-barrier"); + post.succeed(); + timeout(TEST_TIMEOUT, next).await??; + assert!(timeout(TEST_TIMEOUT, transport.receive()).await?.is_some()); + assert!(!blocked.reply.is_closed(), "the first POST remains held"); + blocked.succeed(); + timeout(TEST_TIMEOUT, first).await??; + assert_eq!(counts.posted.load(SeqCst), 2); + assert_eq!(counts.cancelled.load(SeqCst), 0); + transport.close().await?; + Ok(()) +} + +#[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 legacy_cancellation_reaches_the_adapter_before_its_local_stream_drops() +-> 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 held = harness.next_control().await; + harness.counts.manual_controls.store(false, SeqCst); + harness + .counts + .reject_unmatched_cancellation + .store(true, SeqCst); + + let event = harness.cancellable("local-events").await?; + let (eof, dropped) = harness + .next() + .await + .start_local_sse(harness.counts.clone()) + .await?; + let peer = harness.client.peer().clone(); + let mut cancel = Box::pin(peer.notify_cancelled(CancelledNotificationParam::new( + Some(event.id.clone()), + None, + ))); + assert!(futures::poll!(cancel.as_mut()).is_pending()); + // The ordinary request follows the cancellation through the service queue. + let probe = harness.call("after-cancellation"); + harness.next().await.succeed(); + timeout(TEST_TIMEOUT, probe).await???; + + // EOF is ready, but queued cancellation must pause reads until dispatch. + eof.cancel(); + let probe = harness.call("after-eof"); + harness.next().await.succeed(); + timeout(TEST_TIMEOUT, probe).await???; + assert!( + harness + .counts + .local_streams + .lock() + .unwrap() + .contains_key(&event.id), + "the adapter's stream registration must survive queued cancellation" + ); + assert!(!held.reply.is_closed(), "the control slot is still held"); + + held.reply + .send(Ok(StreamableHttpPostResponse::Accepted)) + .unwrap(); + timeout(TEST_TIMEOUT, stale).await???; + timeout(TEST_TIMEOUT, cancel).await??; + timeout(TEST_TIMEOUT, dropped).await??; + assert!(matches!( + timeout(TEST_TIMEOUT, event.await_response()).await?, + Err(rmcp::ServiceError::Cancelled { .. }) + )); + assert_eq!(harness.counts.local_cancellations.load(SeqCst), 1); + assert_eq!(harness.counts.cancelled.load(SeqCst), 1); + harness.finish(vec![], 3, 1).await +} + +#[tokio::test] +async fn control_timeout_is_configurable_and_releases_its_slot() -> anyhow::Result<()> { + assert_eq!(config().control_request_timeout, Duration::from_secs(5)); + let mut harness = + Harness::start(config().control_request_timeout(Duration::from_millis(50))).await?; + harness.counts.manual_controls.store(true, SeqCst); + let cancellation = harness.notify_cancellation(RequestId::Number(999)); + let mut cancelled_post = harness.next_control().await; + harness.incoming.send(sse(json!({ + "jsonrpc": "2.0", "id": "timed-reply", "method": "ping" + })))?; + + let error = timeout(Duration::from_secs(1), cancellation) + .await?? + .unwrap_err(); + assert!(matches!( + transport_error(&error.into()), + StreamableHttpError::ControlRequestTimeout + )); + timeout(Duration::from_secs(1), cancelled_post.reply.closed()).await?; + + let mut reply_post = harness.next_control().await; + assert_eq!(reply_post.message["id"], "timed-reply"); + assert!(reply_post.message["result"].is_object()); + harness.counts.manual_controls.store(false, SeqCst); + let next = harness.notify_cancellation(RequestId::Number(1000)); + timeout(Duration::from_secs(1), reply_post.reply.closed()).await?; + timeout(Duration::from_secs(1), next).await???; + assert_eq!(harness.counts.cancelled.load(SeqCst), 2); + harness.finish(vec![], 0, 1).await +} + +#[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()); + assert!( + !held.reply.is_closed(), + "the old control POST is still held" + ); + assert_eq!(harness.counts.cancelled.load(SeqCst), 1); + + // The default 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???; + timeout(TEST_TIMEOUT, stream.closed()).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(()) +} From 02decfb682ed299f9f439bd7cface84f3488cf29 Mon Sep 17 00:00:00 2001 From: King Star Date: Wed, 26 Aug 2026 00:16:49 +0800 Subject: [PATCH 02/12] fix(transport): fall back after sessionless HTTP discover rejections (#1211) * fix(transport): fall back after HTTP discover rejection * fix(transport): limit legacy fallback to sessionless probes * style: apply nightly rustfmt import ordering --- .../common/reqwest/streamable_http_client.rs | 5 + .../rmcp/src/transport/common/unix_socket.rs | 5 + .../src/transport/streamable_http_client.rs | 39 ++++++- .../test_discover_http_client_startup.rs | 103 +++++++++++++++++- .../test_streamable_http_4xx_error_body.rs | 35 +++++- 5 files changed, 184 insertions(+), 3 deletions(-) diff --git a/crates/rmcp/src/transport/common/reqwest/streamable_http_client.rs b/crates/rmcp/src/transport/common/reqwest/streamable_http_client.rs index 7f08b6c25..b0b5090b5 100644 --- a/crates/rmcp/src/transport/common/reqwest/streamable_http_client.rs +++ b/crates/rmcp/src/transport/common/reqwest/streamable_http_client.rs @@ -293,6 +293,11 @@ impl StreamableHttpClient for reqwest::Client { ), } } + if let Some(response) = + legacy_discover_response(&message, session_was_attached, status, &body) + { + return Ok(response); + } return Err(StreamableHttpError::UnexpectedServerResponse(Cow::Owned( format!("HTTP {status}: {body}"), ))); diff --git a/crates/rmcp/src/transport/common/unix_socket.rs b/crates/rmcp/src/transport/common/unix_socket.rs index 5f995db2b..e3b34898f 100644 --- a/crates/rmcp/src/transport/common/unix_socket.rs +++ b/crates/rmcp/src/transport/common/unix_socket.rs @@ -274,6 +274,11 @@ impl StreamableHttpClient for UnixSocketHttpClient { .await .map(|c| String::from_utf8_lossy(&c.to_bytes()).into_owned()) .unwrap_or_else(|_| "".to_owned()); + if let Some(response) = + legacy_discover_response(&message, session_was_attached, status, &body) + { + return Ok(response); + } return Err(StreamableHttpError::UnexpectedServerResponse(Cow::Owned( format!("HTTP {status}: {body}"), ))); diff --git a/crates/rmcp/src/transport/streamable_http_client.rs b/crates/rmcp/src/transport/streamable_http_client.rs index a4275ab86..d702bc1ca 100644 --- a/crates/rmcp/src/transport/streamable_http_client.rs +++ b/crates/rmcp/src/transport/streamable_http_client.rs @@ -10,7 +10,7 @@ use futures::{ future::BoxFuture, stream::{BoxStream, FuturesUnordered}, }; -use http::{HeaderName, HeaderValue}; +use http::{HeaderName, HeaderValue, StatusCode}; pub use sse_stream::Error as SseError; use sse_stream::Sse; use thiserror::Error; @@ -332,6 +332,43 @@ impl StreamableHttpPostResponse { } } +/// Convert a sessionless discovery rejection into a response the lifecycle +/// layer can classify as a legacy-server signal. +/// +/// Some legacy streamable-HTTP servers reject `server/discover` in middleware +/// before it reaches JSON-RPC dispatch. Their response may be an empty or +/// plain-text 4xx, so there is no JSON-RPC error for the client to forward. +/// Keep authentication failures and server errors on their original paths. +pub(super) fn legacy_discover_response( + message: &ClientJsonRpcMessage, + session_was_attached: bool, + status: StatusCode, + body: &str, +) -> Option { + if session_was_attached + || !status.is_client_error() + || matches!(status, StatusCode::UNAUTHORIZED | StatusCode::FORBIDDEN) + { + return None; + } + + let ClientJsonRpcMessage::Request(request) = message else { + return None; + }; + if !matches!(request.request, ClientRequest::DiscoverRequest(_)) { + return None; + } + + let error = ErrorData::invalid_request( + format!("server/discover rejected with HTTP {status}: {body}"), + None, + ); + Some(StreamableHttpPostResponse::Json( + ServerJsonRpcMessage::error(error, Some(request.id.clone())), + None, + )) +} + /// HTTP backend used by [`StreamableHttpClientTransport`]. /// /// Custom implementations that parse SSE responses must override diff --git a/crates/rmcp/tests/test_discover_http_client_startup.rs b/crates/rmcp/tests/test_discover_http_client_startup.rs index c6051a047..c2825f892 100644 --- a/crates/rmcp/tests/test_discover_http_client_startup.rs +++ b/crates/rmcp/tests/test_discover_http_client_startup.rs @@ -5,8 +5,15 @@ feature = "transport-streamable-http-server" ))] -use std::borrow::Cow; +use std::{borrow::Cow, sync::Arc}; +use axum::{ + Router, + body::{Body, Bytes}, + extract::State, + http::{Response, StatusCode}, + routing::post, +}; use rmcp::{ ClientLifecycleMode, ClientServiceExt, ServerHandler, model::{ClientInfo, DiscoverResult, ErrorCode, ErrorData, ProtocolVersion}, @@ -19,6 +26,8 @@ use rmcp::{ }, }, }; +use serde_json::json; +use tokio::sync::Mutex; use tokio_util::sync::CancellationToken; #[derive(Clone, Default)] @@ -46,6 +55,54 @@ impl ServerHandler for LegacyHttpServer { } } +#[derive(Clone, Default)] +struct PlainTextLegacyHttpState { + methods: Arc>>, +} + +async fn plain_text_legacy_http_handler( + State(state): State, + body: Bytes, +) -> Response { + let request: serde_json::Value = serde_json::from_slice(&body).expect("valid JSON-RPC body"); + let method = request["method"] + .as_str() + .expect("request method") + .to_owned(); + state.methods.lock().await.push(method.clone()); + + if method == "server/discover" { + return Response::builder() + .status(StatusCode::UNPROCESSABLE_ENTITY) + .body(Body::from("Unexpected message, expect initialize request")) + .expect("build rejection response"); + } + + if method == "initialize" { + return Response::builder() + .status(StatusCode::OK) + .header("content-type", "application/json") + .body(Body::from( + json!({ + "jsonrpc": "2.0", + "id": request["id"], + "result": { + "protocolVersion": "2025-11-25", + "capabilities": {}, + "serverInfo": {"name": "legacy", "version": "1.0"} + } + }) + .to_string(), + )) + .expect("build initialize response"); + } + + Response::builder() + .status(StatusCode::ACCEPTED) + .body(Body::empty()) + .expect("build notification response") +} + #[tokio::test] async fn discover_http_client_bootstraps_headers_without_initialize() { let ct = CancellationToken::new(); @@ -135,3 +192,47 @@ async fn auto_http_client_falls_back_to_stateful_legacy_startup() { ct.cancel(); server.await.expect("server task"); } + +#[tokio::test] +async fn auto_http_client_falls_back_after_plain_text_4xx_rejection() { + let ct = CancellationToken::new(); + let state = PlainTextLegacyHttpState::default(); + let methods = state.methods.clone(); + let router = Router::new() + .route("/mcp", post(plain_text_legacy_http_handler)) + .with_state(state); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0") + .await + .expect("listener should bind"); + let address = listener.local_addr().expect("listener address"); + let server = tokio::spawn({ + let ct = ct.clone(); + async move { + let _ = axum::serve(listener, router) + .with_graceful_shutdown(async move { ct.cancelled_owned().await }) + .await; + } + }); + + let transport = StreamableHttpClientTransport::from_config( + StreamableHttpClientTransportConfig::with_uri(format!("http://{address}/mcp")), + ); + let client = ClientInfo::default() + .serve_with_lifecycle( + transport, + ClientLifecycleMode::Auto { + preferred_versions: vec![ProtocolVersion::V_2026_07_28], + legacy_version: Some(ProtocolVersion::V_2025_11_25), + }, + ) + .await + .expect("auto HTTP client should fall back after a transport-level rejection"); + client.cancel().await.expect("cancel client"); + + assert_eq!( + methods.lock().await.as_slice(), + &["server/discover", "initialize", "notifications/initialized"] + ); + ct.cancel(); + server.await.expect("server task"); +} diff --git a/crates/rmcp/tests/test_streamable_http_4xx_error_body.rs b/crates/rmcp/tests/test_streamable_http_4xx_error_body.rs index ea49a4172..13deeeb60 100644 --- a/crates/rmcp/tests/test_streamable_http_4xx_error_body.rs +++ b/crates/rmcp/tests/test_streamable_http_4xx_error_body.rs @@ -7,7 +7,10 @@ use std::{collections::HashMap, sync::Arc}; use rmcp::{ - model::{ClientJsonRpcMessage, ClientRequest, PingRequest, RequestId}, + model::{ + ClientJsonRpcMessage, ClientRequest, DiscoverRequest, DiscoverRequestParams, PingRequest, + RequestId, + }, transport::streamable_http_client::{ StreamableHttpClient, StreamableHttpError, StreamableHttpPostResponse, }, @@ -45,6 +48,13 @@ fn ping_message() -> ClientJsonRpcMessage { ) } +fn discover_message() -> ClientJsonRpcMessage { + ClientJsonRpcMessage::request( + ClientRequest::DiscoverRequest(DiscoverRequest::new(DiscoverRequestParams {})), + RequestId::Number(1), + ) +} + /// HTTP 4xx with Content-Type: application/json and a valid JSON-RPC error body /// must be surfaced as `StreamableHttpPostResponse::Json`, not swallowed as a /// transport error. @@ -119,3 +129,26 @@ async fn http_4xx_malformed_json_body_falls_back_to_unexpected_server_response() other => panic!("expected UnexpectedServerResponse, got: {other:?}"), } } + +/// A discovery request attached to an existing session must retain the +/// transport error path; legacy fallback only applies to sessionless probes. +#[tokio::test] +async fn sessionful_discover_4xx_does_not_trigger_legacy_fallback() { + let url = spawn_mock_server(422, "text/plain", "Session-bound rejection").await; + + let client = reqwest::Client::new(); + let result = client + .post_message( + Arc::from(url.as_str()), + discover_message(), + Some(Arc::from("session-1")), + None, + HashMap::new(), + ) + .await; + + match result { + Err(StreamableHttpError::UnexpectedServerResponse(_)) => {} + other => panic!("expected UnexpectedServerResponse, got: {other:?}"), + } +} From 78cf9f1cb654bcc7b6c8d60e40a8ff527bd0e542 Mon Sep 17 00:00:00 2001 From: Jack Amadeo Date: Tue, 25 Aug 2026 14:50:29 -0400 Subject: [PATCH 03/12] ci: pin GitHub Actions to commit SHAs (#1216) --- .github/workflows/auto-label-pr.yml | 2 +- .github/workflows/ci.yml | 92 ++++++++++++++--------------- .github/workflows/codeql.yml | 8 +-- .github/workflows/conformance.yml | 16 ++--- .github/workflows/release-plz.yml | 8 +-- .github/workflows/stale.yml | 2 +- .github/workflows/triage.yml | 2 +- 7 files changed, 65 insertions(+), 65 deletions(-) diff --git a/.github/workflows/auto-label-pr.yml b/.github/workflows/auto-label-pr.yml index 9850dafa5..7259750e2 100644 --- a/.github/workflows/auto-label-pr.yml +++ b/.github/workflows/auto-label-pr.yml @@ -20,7 +20,7 @@ jobs: PR_URL: ${{ github.event.pull_request.html_url }} steps: - - uses: actions/labeler@v7 + - uses: actions/labeler@bf12e9b00b37c5c0ca2b87b79b2daf7891dbda13 # v7 with: # Auto-include paths starting with dot (e.g. .github) dot: true diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 26b812055..ab19c7b9a 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -18,12 +18,12 @@ jobs: runs-on: ubuntu-latest if: github.event_name == 'pull_request' steps: - - uses: actions/checkout@v7 + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: fetch-depth: 0 - name: Setup Node.js - uses: actions/setup-node@v7 + uses: actions/setup-node@820762786026740c76f36085b0efc47a31fe5020 # v7 with: node-version: '22' @@ -39,7 +39,7 @@ jobs: name: Code Formatting runs-on: ubuntu-latest steps: - - uses: actions/checkout@v7 + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 - name: Install Rust fmt run: rustup toolchain install nightly --component rustfmt @@ -51,12 +51,12 @@ jobs: name: Lint with Clippy runs-on: ubuntu-latest steps: - - uses: actions/checkout@v7 + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 - name: Install Rust - uses: dtolnay/rust-toolchain@stable + uses: dtolnay/rust-toolchain@4360b52568e2003a75bf9bc1d59f33a8e3fc893c # stable - - uses: Swatinem/rust-cache@v2 + - uses: Swatinem/rust-cache@6323deb102c322ba6fcbdcafc7e3dddab59af2b6 # v2 - name: Run clippy run: cargo clippy --all-targets --all-features -- -D warnings @@ -66,17 +66,17 @@ jobs: runs-on: ubuntu-latest if: github.event_name == 'pull_request' steps: - - uses: actions/checkout@v7 + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: fetch-depth: 0 - name: Install Rust - uses: dtolnay/rust-toolchain@stable + uses: dtolnay/rust-toolchain@4360b52568e2003a75bf9bc1d59f33a8e3fc893c # stable - - uses: Swatinem/rust-cache@v2 + - uses: Swatinem/rust-cache@6323deb102c322ba6fcbdcafc7e3dddab59af2b6 # v2 - name: Install cargo-semver-checks - uses: taiki-e/install-action@v2.85.13 + uses: taiki-e/install-action@82cd3e7658a6f96c86c0234aeeda1748937cb0a1 # v2.85.13 with: tool: cargo-semver-checks @@ -117,7 +117,7 @@ jobs: runs-on: ubuntu-latest if: github.event_name == 'pull_request' steps: - - uses: actions/checkout@v7 + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: fetch-depth: 0 @@ -125,12 +125,12 @@ jobs: # to be installed (it does not need to be the default; the tool invokes it # via `cargo +nightly`). - name: Install Rust - uses: dtolnay/rust-toolchain@nightly + uses: dtolnay/rust-toolchain@7c8d7d138f5c09cef361f8214cf96882cd029cdb # nightly - - uses: Swatinem/rust-cache@v2 + - uses: Swatinem/rust-cache@6323deb102c322ba6fcbdcafc7e3dddab59af2b6 # v2 - name: Install cargo-public-api - uses: taiki-e/install-action@v2.85.13 + uses: taiki-e/install-action@82cd3e7658a6f96c86c0234aeeda1748937cb0a1 # v2.85.13 with: tool: cargo-public-api @@ -201,9 +201,9 @@ jobs: name: spell check with typos runs-on: ubuntu-latest steps: - - uses: actions/checkout@v7 + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 - name: Spell Check Repo - uses: crate-ci/typos@master + uses: crate-ci/typos@1a51d4b5a03bb97576af186c813af67e9137ba7c # master msrv: name: Check MSRV @@ -211,7 +211,7 @@ jobs: permissions: contents: read steps: - - uses: actions/checkout@v7 + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 - name: Read MSRV from Cargo.toml run: | @@ -220,11 +220,11 @@ jobs: echo "MSRV=$MSRV" >> "$GITHUB_ENV" - name: Install MSRV Rust - uses: dtolnay/rust-toolchain@master + uses: dtolnay/rust-toolchain@6c977a6ca4077a0ceb28ffbe03f59d46e9ac8772 # master with: toolchain: ${{ env.MSRV }} - - uses: Swatinem/rust-cache@v2 + - uses: Swatinem/rust-cache@6323deb102c322ba6fcbdcafc7e3dddab59af2b6 # v2 - name: Check default workspace members with MSRV run: cargo +${{ env.MSRV }} check --all-targets --all-features @@ -233,19 +233,19 @@ jobs: name: Run Tests runs-on: ubuntu-latest steps: - - uses: actions/checkout@v7 + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 # install nodejs - name: Setup Node.js - uses: actions/setup-node@v7 + uses: actions/setup-node@820762786026740c76f36085b0efc47a31fe5020 # v7 with: node-version: '22' - name: Install uv - uses: astral-sh/setup-uv@v7 + uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7 - name: Install Rust - uses: dtolnay/rust-toolchain@stable + uses: dtolnay/rust-toolchain@4360b52568e2003a75bf9bc1d59f33a8e3fc893c # stable - name: Set up Python run: uv python install @@ -253,7 +253,7 @@ jobs: - name: Create venv for python run: uv venv - - uses: Swatinem/rust-cache@v2 + - uses: Swatinem/rust-cache@6323deb102c322ba6fcbdcafc7e3dddab59af2b6 # v2 - name: Run tests run: cargo test --all-features @@ -262,19 +262,19 @@ jobs: name: Run Tests (no local feature) runs-on: ubuntu-latest steps: - - uses: actions/checkout@v7 + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 # install nodejs - name: Setup Node.js - uses: actions/setup-node@v7 + uses: actions/setup-node@820762786026740c76f36085b0efc47a31fe5020 # v7 with: node-version: '22' - name: Install uv - uses: astral-sh/setup-uv@v7 + uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7 - name: Install Rust - uses: dtolnay/rust-toolchain@stable + uses: dtolnay/rust-toolchain@4360b52568e2003a75bf9bc1d59f33a8e3fc893c # stable - name: Set up Python run: uv python install @@ -282,7 +282,7 @@ jobs: - name: Create venv for python run: uv venv - - uses: Swatinem/rust-cache@v2 + - uses: Swatinem/rust-cache@6323deb102c322ba6fcbdcafc7e3dddab59af2b6 # v2 - name: Run tests without local feature run: | @@ -298,19 +298,19 @@ jobs: permissions: contents: write steps: - - uses: actions/checkout@v7 + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 # install nodejs - name: Setup Node.js - uses: actions/setup-node@v7 + uses: actions/setup-node@820762786026740c76f36085b0efc47a31fe5020 # v7 with: node-version: '22' - name: Install uv - uses: astral-sh/setup-uv@v7 + uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7 - name: Install Rust - uses: dtolnay/rust-toolchain@stable + uses: dtolnay/rust-toolchain@4360b52568e2003a75bf9bc1d59f33a8e3fc893c # stable - name: Set up Python run: uv python install @@ -318,7 +318,7 @@ jobs: - name: Create venv for python run: uv venv - - uses: Swatinem/rust-cache@v2 + - uses: Swatinem/rust-cache@6323deb102c322ba6fcbdcafc7e3dddab59af2b6 # v2 - name: Install cargo-llvm-cov run: cargo install cargo-llvm-cov @@ -333,19 +333,19 @@ jobs: name: Example test runs-on: ubuntu-latest steps: - - uses: actions/checkout@v7 + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 # install nodejs - name: Setup Node.js - uses: actions/setup-node@v7 + uses: actions/setup-node@820762786026740c76f36085b0efc47a31fe5020 # v7 with: node-version: '22' - name: Install uv - uses: astral-sh/setup-uv@v7 + uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7 - name: Install Rust - uses: dtolnay/rust-toolchain@stable + uses: dtolnay/rust-toolchain@4360b52568e2003a75bf9bc1d59f33a8e3fc893c # stable - name: Set up Python run: uv python install @@ -353,7 +353,7 @@ jobs: - name: Create venv for python run: uv venv - - uses: Swatinem/rust-cache@v2 + - uses: Swatinem/rust-cache@6323deb102c322ba6fcbdcafc7e3dddab59af2b6 # v2 - name: Add target WASI preview 2 run: | @@ -389,12 +389,12 @@ jobs: name: Security Audit runs-on: ubuntu-latest steps: - - uses: actions/checkout@v7 + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 - name: Install Rust - uses: dtolnay/rust-toolchain@stable + uses: dtolnay/rust-toolchain@4360b52568e2003a75bf9bc1d59f33a8e3fc893c # stable - - uses: Swatinem/rust-cache@v2 + - uses: Swatinem/rust-cache@6323deb102c322ba6fcbdcafc7e3dddab59af2b6 # v2 - name: Install cargo-audit run: cargo install cargo-audit @@ -406,12 +406,12 @@ jobs: name: Generate Documentation runs-on: ubuntu-latest steps: - - uses: actions/checkout@v7 + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 - name: Install Rust - uses: dtolnay/rust-toolchain@nightly + uses: dtolnay/rust-toolchain@7c8d7d138f5c09cef361f8214cf96882cd029cdb # nightly - - uses: Swatinem/rust-cache@v2 + - uses: Swatinem/rust-cache@6323deb102c322ba6fcbdcafc7e3dddab59af2b6 # v2 - name: Generate documentation run: | @@ -432,7 +432,7 @@ jobs: # This happened recently in the attack on `tj-actions/changed-files`, but # has happened many times before as well. - - uses: actions/checkout@v7 + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 - name: Update Rust run: | diff --git a/.github/workflows/codeql.yml b/.github/workflows/codeql.yml index c1f2788a3..7c3a77129 100644 --- a/.github/workflows/codeql.yml +++ b/.github/workflows/codeql.yml @@ -21,16 +21,16 @@ jobs: language: [rust, javascript-typescript, python, actions] steps: - name: Checkout repository - uses: actions/checkout@v7 + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 - name: Initialize CodeQL - uses: github/codeql-action/init@v4.37.7 + uses: github/codeql-action/init@ff2f1c621b7f889edc0d3c761ac2e6a3f8cdb0dd # v4.37.7 with: languages: ${{ matrix.language }} config-file: ./.github/codeql/codeql-config.yml - name: Autobuild - uses: github/codeql-action/autobuild@v4.37.7 + uses: github/codeql-action/autobuild@ff2f1c621b7f889edc0d3c761ac2e6a3f8cdb0dd # v4.37.7 - name: Perform CodeQL Analysis - uses: github/codeql-action/analyze@v4.37.7 + uses: github/codeql-action/analyze@ff2f1c621b7f889edc0d3c761ac2e6a3f8cdb0dd # v4.37.7 diff --git a/.github/workflows/conformance.yml b/.github/workflows/conformance.yml index ced1875d8..857f40074 100644 --- a/.github/workflows/conformance.yml +++ b/.github/workflows/conformance.yml @@ -21,12 +21,12 @@ jobs: permissions: contents: read steps: - - uses: actions/checkout@v7 + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 - name: Install Rust toolchain - uses: dtolnay/rust-toolchain@stable + uses: dtolnay/rust-toolchain@4360b52568e2003a75bf9bc1d59f33a8e3fc893c # stable - - uses: Swatinem/rust-cache@v2 + - uses: Swatinem/rust-cache@6323deb102c322ba6fcbdcafc7e3dddab59af2b6 # v2 - name: Build conformance binaries run: cargo build -p mcp-conformance @@ -91,7 +91,7 @@ jobs: - name: Upload results if: always() - uses: actions/upload-artifact@v7 + uses: actions/upload-artifact@043fb46d1a93c77aae656e7c1c64a875d1fc6a0a # v7.0.1 with: name: conformance-server-results path: | @@ -104,12 +104,12 @@ jobs: permissions: contents: read steps: - - uses: actions/checkout@v7 + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 - name: Install Rust toolchain - uses: dtolnay/rust-toolchain@stable + uses: dtolnay/rust-toolchain@4360b52568e2003a75bf9bc1d59f33a8e3fc893c # stable - - uses: Swatinem/rust-cache@v2 + - uses: Swatinem/rust-cache@6323deb102c322ba6fcbdcafc7e3dddab59af2b6 # v2 - name: Build conformance binaries run: cargo build -p mcp-conformance @@ -140,7 +140,7 @@ jobs: - name: Upload results if: always() - uses: actions/upload-artifact@v7 + uses: actions/upload-artifact@043fb46d1a93c77aae656e7c1c64a875d1fc6a0a # v7.0.1 with: name: conformance-client-results path: conformance-client-results diff --git a/.github/workflows/release-plz.yml b/.github/workflows/release-plz.yml index 04f51190e..4c4356e24 100644 --- a/.github/workflows/release-plz.yml +++ b/.github/workflows/release-plz.yml @@ -20,11 +20,11 @@ jobs: contents: write steps: - name: Checkout repository - uses: actions/checkout@v7 + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: fetch-depth: 0 - name: Install Rust toolchain - uses: dtolnay/rust-toolchain@stable + uses: dtolnay/rust-toolchain@4360b52568e2003a75bf9bc1d59f33a8e3fc893c # stable with: toolchain: "1.92" # Using fork until semver_check_features support is released upstream. @@ -51,11 +51,11 @@ jobs: cancel-in-progress: false steps: - name: Checkout repository - uses: actions/checkout@v7 + uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 with: fetch-depth: 0 - name: Install Rust toolchain - uses: dtolnay/rust-toolchain@stable + uses: dtolnay/rust-toolchain@4360b52568e2003a75bf9bc1d59f33a8e3fc893c # stable with: toolchain: "1.92" # Using fork until semver_check_features support is released upstream. diff --git a/.github/workflows/stale.yml b/.github/workflows/stale.yml index 58794d517..6140f0610 100644 --- a/.github/workflows/stale.yml +++ b/.github/workflows/stale.yml @@ -12,7 +12,7 @@ jobs: issues: write pull-requests: write steps: - - uses: actions/stale@v11 + - uses: actions/stale@4391f3da665fdf50b6810c1a66712fb9ba21aa93 # v11 with: days-before-stale: 60 days-before-close: 14 diff --git a/.github/workflows/triage.yml b/.github/workflows/triage.yml index d0d9186a5..4f1b9abfe 100644 --- a/.github/workflows/triage.yml +++ b/.github/workflows/triage.yml @@ -39,7 +39,7 @@ jobs: TRIAGE_MODEL: ${{ vars.TRIAGE_MODEL || 'gpt-4o-mini' }} steps: - - uses: actions/checkout@v7 + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 - name: Install jq run: sudo apt-get install -y jq From 103bf239b3c2fb24bdd6db2164854ab7a058c8a8 Mon Sep 17 00:00:00 2001 From: Jack Amadeo Date: Thu, 27 Aug 2026 05:52:18 -0400 Subject: [PATCH 04/12] ci: scope release token permissions to jobs (#1220) --- .github/workflows/release-plz.yml | 5 ++--- 1 file changed, 2 insertions(+), 3 deletions(-) diff --git a/.github/workflows/release-plz.yml b/.github/workflows/release-plz.yml index 4c4356e24..41570a624 100644 --- a/.github/workflows/release-plz.yml +++ b/.github/workflows/release-plz.yml @@ -1,8 +1,6 @@ name: Release-plz -permissions: - pull-requests: write - contents: write +permissions: {} on: push: @@ -18,6 +16,7 @@ jobs: if: ${{ github.repository_owner == 'modelcontextprotocol' }} permissions: contents: write + pull-requests: read steps: - name: Checkout repository uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 From 3501f3e646d2f74c9c241e0bcea2621981ebffa6 Mon Sep 17 00:00:00 2001 From: Jack Amadeo Date: Thu, 27 Aug 2026 06:01:29 -0400 Subject: [PATCH 05/12] ci: default workflow tokens to read-only contents (#1218) --- .github/workflows/auto-label-pr.yml | 3 +++ .github/workflows/ci.yml | 3 +++ .github/workflows/codeql.yml | 3 +++ .github/workflows/conformance.yml | 3 +++ .github/workflows/stale.yml | 3 +++ 5 files changed, 15 insertions(+) diff --git a/.github/workflows/auto-label-pr.yml b/.github/workflows/auto-label-pr.yml index 7259750e2..d91ce04b3 100644 --- a/.github/workflows/auto-label-pr.yml +++ b/.github/workflows/auto-label-pr.yml @@ -3,6 +3,9 @@ on: # Runs workflow when activity on a PR in the workflow's repository occurs. pull_request_target: +permissions: + contents: read + jobs: auto-label: permissions: diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index ab19c7b9a..efa921857 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -12,6 +12,9 @@ env: CARGO_TERM_COLOR: always ARTIFACT_DIR: release-artifacts +permissions: + contents: read + jobs: commit-lint: name: Lint Commit Messages diff --git a/.github/workflows/codeql.yml b/.github/workflows/codeql.yml index 7c3a77129..4ecf3484b 100644 --- a/.github/workflows/codeql.yml +++ b/.github/workflows/codeql.yml @@ -8,6 +8,9 @@ on: schedule: - cron: '0 0 * * 1' # Weekly on Monday +permissions: + contents: read + jobs: analyze: name: Analyze diff --git a/.github/workflows/conformance.yml b/.github/workflows/conformance.yml index 857f40074..1233deb56 100644 --- a/.github/workflows/conformance.yml +++ b/.github/workflows/conformance.yml @@ -14,6 +14,9 @@ concurrency: env: CONFORMANCE_VERSION: "0.2.0-alpha.10" +permissions: + contents: read + jobs: server: runs-on: ubuntu-latest diff --git a/.github/workflows/stale.yml b/.github/workflows/stale.yml index 6140f0610..fe1baa96f 100644 --- a/.github/workflows/stale.yml +++ b/.github/workflows/stale.yml @@ -4,6 +4,9 @@ on: - cron: "0 0 * * *" workflow_dispatch: +permissions: + contents: read + jobs: stale: name: Clean Up From c8bb1c6f7d22d38e575166f364fcab3f9b0fe72c Mon Sep 17 00:00:00 2001 From: Jack Amadeo Date: Thu, 27 Aug 2026 06:02:06 -0400 Subject: [PATCH 06/12] ci: remove coverage job write permission (#1219) --- .github/workflows/ci.yml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index efa921857..c72680b33 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -299,7 +299,7 @@ jobs: name: Code Coverage runs-on: ubuntu-latest permissions: - contents: write + contents: read steps: - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 From 12db02833f69f5793004ed218ca44ca7b1d8a15c Mon Sep 17 00:00:00 2001 From: Jack Amadeo Date: Thu, 27 Aug 2026 12:33:32 -0400 Subject: [PATCH 07/12] ci: pin release-plz fork revision (#1221) --- .github/workflows/release-plz.yml | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/.github/workflows/release-plz.yml b/.github/workflows/release-plz.yml index 41570a624..311b74ec0 100644 --- a/.github/workflows/release-plz.yml +++ b/.github/workflows/release-plz.yml @@ -29,7 +29,7 @@ jobs: # Using fork until semver_check_features support is released upstream. # See: https://github.com/release-plz/release-plz/pull/2757 - name: Install release-plz from fork - run: cargo install --locked --git https://github.com/DaleSeo/release-plz --branch feat/semver-check-features release-plz + run: cargo install --locked --git https://github.com/DaleSeo/release-plz --rev c2ea743bacf9bbd56e344df68219e60b0ad35d04 release-plz - name: Run release-plz release run: release-plz release env: @@ -60,7 +60,7 @@ jobs: # Using fork until semver_check_features support is released upstream. # See: https://github.com/release-plz/release-plz/pull/2757 - name: Install release-plz from fork - run: cargo install --locked --git https://github.com/DaleSeo/release-plz --branch feat/semver-check-features release-plz + run: cargo install --locked --git https://github.com/DaleSeo/release-plz --rev c2ea743bacf9bbd56e344df68219e60b0ad35d04 release-plz - name: Run release-plz release-pr run: release-plz release-pr env: From 3ef34e64e7f65526037ce650ea446be475839056 Mon Sep 17 00:00:00 2001 From: camillelawrence Date: Fri, 28 Aug 2026 09:19:23 -0400 Subject: [PATCH 08/12] feat: add request-state key rotation (#1128) * feat: add request-state key rotation * docs: streamline request-state codec documentation * fix: harden request-state fallback verification * refactor: refine request-state keyring API * docs: make request-state rotation guidance self-contained * docs: streamline request-state keyring rustdocs * test: streamline request-state keyring coverage * docs: restore request-state key rotation doctest * chore: remove manual changelog entry --- README.md | 20 +- crates/rmcp/src/model/request_state.rs | 938 +++++++++++++++++++++++-- 2 files changed, 897 insertions(+), 61 deletions(-) diff --git a/README.md b/README.md index 6e90827bc..390f96c8a 100644 --- a/README.md +++ b/README.md @@ -1363,11 +1363,21 @@ async fn call_tool(&self, request: CallToolRequestParams, _ctx: RequestContext **`requestState` is untrusted.** The client echoes it back verbatim, so a -> stateless server that stores meaningful data in it MUST verify integrity -> first. Enable the `request-state` feature and use `RequestStateCodec` to seal -> and open it (HMAC-tagged), or keep state server-side and use `requestState` -> only as an opaque handle. +> **`requestState` is untrusted.** [SEP-2322 requires servers to validate +> it](https://modelcontextprotocol.io/seps/2322-MRTR#protocol-requirements-for-ephemeral-workflow) +> because the client echoes it back verbatim. A stateless server that stores +> meaningful data in it MUST verify integrity first. Enable the `request-state` +> feature and use `RequestStateCodec` to seal and open it (HMAC-tagged), or keep +> state server-side and use `requestState` only as an opaque handle. + +For multi-replica deployments, use `RequestStateCodec::new_with_keyring` to +rotate signing keys without invalidating in-flight requests: + +1. Deploy the old and new keys everywhere, continuing to emit `rs1` with the + old key via `with_rs1_signing("old")`. +2. Start emitting `rs2` with the new key while retaining the old key via + `with_rs1_fallback("old")`. +3. After the maximum `requestState` lifetime has elapsed, remove the old key. ### Client-side diff --git a/crates/rmcp/src/model/request_state.rs b/crates/rmcp/src/model/request_state.rs index c6415aa22..9f45572c7 100644 --- a/crates/rmcp/src/model/request_state.rs +++ b/crates/rmcp/src/model/request_state.rs @@ -63,7 +63,10 @@ //! # } //! ``` -use std::time::Duration; +use std::{ + collections::{HashMap, hash_map::Entry}, + time::Duration, +}; use base64::{Engine, engine::general_purpose::URL_SAFE_NO_PAD}; use hmac::{Hmac, KeyInit, Mac}; @@ -74,19 +77,28 @@ use zeroize::Zeroizing; type HmacSha256 = Hmac; -/// Version tag prefixing every sealed value, so the wire format can evolve. -const VERSION: &str = "rs1"; +const VERSION_V1: &str = "rs1"; + +const VERSION_V2: &str = "rs2"; /// Domain-separation label mixed into the HMAC so a `requestState` tag can never /// be confused with an HMAC computed for some other purpose using the same key. -const DOMAIN: &[u8] = b"rmcp/mrtr/request-state/v1"; +const DOMAIN_V1: &[u8] = b"rmcp/mrtr/request-state/v1"; + +/// Domain-separation label for keyed request-state tags. +const DOMAIN_V2: &[u8] = b"rmcp/mrtr/request-state/v2"; /// Length of the big-endian expiry prefix (unix milliseconds) stored at the /// front of every sealed body. `0` means "no expiry". const EXPIRY_LEN: usize = 8; -/// Errors returned when constructing a [`RequestStateCodec`] or processing a -/// sealed value. +const MAX_KID_LEN: usize = 255; + +/// Maximum unpadded base64url length for [`MAX_KID_LEN`] bytes. +const MAX_ENCODED_KID_LEN: usize = 340; + +/// Errors returned when constructing or configuring a +/// [`RequestStateCodec`], or when processing a sealed value. #[derive(Debug, Error)] #[non_exhaustive] pub enum RequestStateError { @@ -114,6 +126,18 @@ pub enum RequestStateError { #[error("request state has expired")] Expired, + /// The token names (or, for rs1, requires) a key this codec does not hold. + #[error("request state was sealed with an unknown key")] + UnknownKeyId, + + /// The decoded rs2 key id was empty, too long, or not valid UTF-8. + #[error("request state contains an invalid key identifier")] + InvalidKeyId, + + /// An invalid keyring configuration; the message is diagnostic only. + #[error("invalid keyring configuration: {0}")] + InvalidKeyring(&'static str), + /// The sealed payload could not be serialized to JSON. #[error("failed to serialize request state payload: {0}")] Serialization(#[source] serde_json::Error), @@ -159,28 +183,95 @@ impl<'a> SealOptions<'a> { } } -/// A keyed codec that seals and opens SEP-2322 `requestState` values with -/// HMAC-SHA256 integrity protection. +/// A keyed codec for integrity-protected SEP-2322 `requestState` values. /// -/// Construct one codec per signing key and reuse it for the lifetime of the -/// key. The same key must be used to [`seal`](Self::seal) and -/// [`open`](Self::open) a value, so it has to survive across the rounds of a -/// single MRTR exchange (e.g. a stable per-process or per-deployment secret). +/// [`new`](Self::new) preserves `rs1`; [`new_with_keyring`](Self::new_with_keyring) +/// emits `rs2`. +/// Use [`with_rs1_signing`](Self::with_rs1_signing) and +/// [`with_rs1_fallback`](Self::with_rs1_fallback) for rolling migrations. /// -/// Use [`try_new`](Self::try_new) to require at least +/// Use [`try_new`](Self::try_new) and +/// [`new_with_keyring`](Self::new_with_keyring) to require at least /// [`MIN_KEY_LENGTH`](Self::MIN_KEY_LENGTH) bytes of high-entropy key material. -/// The stored key material is zeroized when the codec is dropped. +/// Configure the same keys on every replica. Stored key material is zeroized +/// when the codec is dropped. +/// +/// Values are authenticated, not encrypted; key ids and payloads remain +/// readable. Retain old keys until every value they signed has expired. +/// +/// # Key rotation +/// +/// Rotate in stages so every replica can open states emitted by its peers: +/// +/// ``` +/// # use rmcp::model::{RequestStateCodec, RequestStateError}; +/// # fn main() -> Result<(), RequestStateError> { +/// let old_key = b"old-request-state-key-at-least-32b".as_slice(); +/// let new_key = b"new-request-state-key-at-least-32b".as_slice(); +/// let keys = [("old", old_key), ("new", new_key)]; +/// +/// // 1. Deploy both keys, but continue emitting rs1 with the old key. +/// let transitional = RequestStateCodec::new_with_keyring("new", keys)? +/// .with_rs1_signing("old")?; +/// +/// // 2. Emit rs2 with the new key, while accepting in-flight rs1 values. +/// let promoted = RequestStateCodec::new_with_keyring("new", keys)? +/// .with_rs1_fallback("old")?; +/// +/// let old_state = transitional.seal(b"old state"); +/// assert_eq!(promoted.open(&old_state)?, b"old state"); +/// let new_state = promoted.seal(b"new state"); +/// assert_eq!(transitional.open(&new_state)?, b"new state"); +/// +/// // 3. After old values expire, remove the old key. +/// let retired = RequestStateCodec::new_with_keyring("new", [("new", new_key)])?; +/// assert_eq!(retired.open(&new_state)?, b"new state"); +/// # Ok(()) +/// # } +/// ``` #[derive(Clone)] pub struct RequestStateCodec { - key: Zeroizing>, + keys: Keys, +} + +#[derive(Clone)] +enum Keys { + Single(Zeroizing>), + Ring { + keys: HashMap>>, + seal_mode: SealMode, + rs1_fallbacks: Vec, + }, +} + +#[derive(Clone, Debug)] +enum SealMode { + Rs1 { key_id: String }, + Rs2 { key_id: String }, } impl std::fmt::Debug for RequestStateCodec { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - // Never leak the signing key through Debug output. - f.debug_struct("RequestStateCodec") - .field("key", &"") - .finish() + // Never leak signing or verification keys through Debug output. + match &self.keys { + Keys::Single(_) => f + .debug_struct("RequestStateCodec") + .field("mode", &"single") + .field("key", &"") + .finish(), + Keys::Ring { + keys, + seal_mode, + rs1_fallbacks, + } => f + .debug_struct("RequestStateCodec") + .field("mode", &"ring") + .field("key_count", &keys.len()) + .field("keys", &"") + .field("seal_mode", seal_mode) + .field("rs1_fallbacks", rs1_fallbacks) + .finish(), + } } } @@ -188,7 +279,8 @@ impl RequestStateCodec { /// Minimum accepted signing-key length in bytes. pub const MIN_KEY_LENGTH: usize = 32; - /// Creates a codec from a signing key without validating its length. + /// Creates a legacy single-key codec that seals and opens `rs1` values + /// without validating the signing-key length. #[deprecated( since = "3.1.4", note = "use RequestStateCodec::try_new to enforce the minimum signing-key length" @@ -203,11 +295,11 @@ impl RequestStateCodec { /// validated. pub fn new_unchecked(key: impl Into>) -> Self { Self { - key: Zeroizing::new(key.into()), + keys: Keys::Single(Zeroizing::new(key.into())), } } - /// Creates a codec from a signing key after validating its length. + /// Creates a legacy single-key codec after validating its signing-key length. /// /// # Errors /// @@ -215,14 +307,138 @@ impl RequestStateCodec { /// than [`MIN_KEY_LENGTH`](Self::MIN_KEY_LENGTH) bytes. pub fn try_new(key: impl Into>) -> Result { let key = Zeroizing::new(key.into()); - if key.len() < Self::MIN_KEY_LENGTH { - return Err(RequestStateError::KeyTooShort { - minimum: Self::MIN_KEY_LENGTH, - actual: key.len(), - }); + Self::validate_key_length(&key)?; + Ok(Self { + keys: Keys::Single(key), + }) + } + + /// Creates a keyring codec that seals `rs2` with `active_kid`. + /// + /// All configured keys can open matching `rs2` values. Legacy `rs1` values + /// require [`with_rs1_fallback`](Self::with_rs1_fallback) or + /// [`with_rs1_signing`](Self::with_rs1_signing). + /// Key ids are case-sensitive UTF-8 strings of 1 to 255 bytes. + /// + /// # Errors + /// + /// Returns [`RequestStateError::KeyTooShort`] if a configured key contains + /// fewer than [`MIN_KEY_LENGTH`](Self::MIN_KEY_LENGTH) bytes. Returns + /// [`RequestStateError::InvalidKeyring`] if the keyring is empty, contains + /// duplicate or invalid key ids, or does not contain `active_kid`. + pub fn new_with_keyring( + active_kid: impl Into, + keys: impl IntoIterator, + ) -> Result + where + K: Into, + V: Into>, + { + let mut ring = HashMap::new(); + for (kid, key) in keys { + let kid = kid.into(); + let key = Zeroizing::new(key.into()); + Self::validate_config_kid(&kid)?; + Self::validate_key_length(&key)?; + match ring.entry(kid) { + Entry::Vacant(entry) => { + entry.insert(key); + } + Entry::Occupied(_) => { + return Err(RequestStateError::InvalidKeyring("duplicate key id")); + } + } + } + + if ring.is_empty() { + return Err(RequestStateError::InvalidKeyring( + "keyring must contain at least one key", + )); } - Ok(Self { key }) + let active_kid = active_kid.into(); + Self::validate_config_kid(&active_kid)?; + if !ring.contains_key(&active_kid) { + return Err(RequestStateError::InvalidKeyring( + "active key id is not present in the keyring", + )); + } + + Ok(Self { + keys: Keys::Ring { + keys: ring, + seal_mode: SealMode::Rs2 { key_id: active_kid }, + rs1_fallbacks: Vec::new(), + }, + }) + } + + /// Uses `kid` to sign legacy `rs1` values during a rolling migration. + /// The codec accepts its own `rs1` output and configured `rs2` values. + /// + /// # Errors + /// + /// Returns [`RequestStateError::InvalidKeyring`] when called on a single-key + /// codec, or if the keyring does not contain `kid`. + pub fn with_rs1_signing(mut self, kid: impl AsRef) -> Result { + let kid = kid.as_ref(); + match &mut self.keys { + Keys::Single(_) => Err(RequestStateError::InvalidKeyring( + "rs1 transitional signing requires a keyring", + )), + Keys::Ring { + keys, + seal_mode, + rs1_fallbacks, + } => { + if !keys.contains_key(kid) { + return Err(RequestStateError::InvalidKeyring( + "rs1 signing key id is not present in the keyring", + )); + } + + let kid = kid.to_owned(); + *seal_mode = SealMode::Rs1 { + key_id: kid.clone(), + }; + if !rs1_fallbacks.iter().any(|existing| existing == &kid) { + rs1_fallbacks.push(kid); + } + Ok(self) + } + } + } + + /// Accepts legacy `rs1` values signed with `kid`. + /// Adding the same id more than once has no effect; remove fallbacks after + /// old values have drained. + /// + /// # Errors + /// + /// Returns [`RequestStateError::InvalidKeyring`] when called on a single-key + /// codec, or if the keyring does not contain `kid`. + pub fn with_rs1_fallback(mut self, kid: impl AsRef) -> Result { + let kid = kid.as_ref(); + match &mut self.keys { + Keys::Single(_) => Err(RequestStateError::InvalidKeyring( + "rs1 fallbacks require a keyring", + )), + Keys::Ring { + keys, + rs1_fallbacks, + .. + } => { + if !keys.contains_key(kid) { + return Err(RequestStateError::InvalidKeyring( + "rs1 fallback key id is not present in the keyring", + )); + } + if !rs1_fallbacks.iter().any(|existing| existing == kid) { + rs1_fallbacks.push(kid.to_owned()); + } + Ok(self) + } + } } /// Seals raw bytes into an opaque, integrity-protected string suitable for @@ -280,12 +496,12 @@ impl RequestStateCodec { /// /// # Errors /// - /// - [`RequestStateError::IntegrityCheckFailed`] if the value was not - /// produced by this key or the associated data differs. - /// - [`RequestStateError::Expired`] if the value's TTL has elapsed. - /// - [`RequestStateError::MalformedFormat`] or - /// [`RequestStateError::InvalidEncoding`] if it is not a well-formed sealed - /// value. + /// Returns a [`RequestStateError`] if the value is malformed, cannot be + /// verified, names an unavailable key, or has expired. + /// + /// Applications MUST map all token-opening failures to a single + /// client-facing error and reserve the detailed variants for internal + /// diagnostics. pub fn open_with( &self, sealed: &str, @@ -332,24 +548,76 @@ impl RequestStateCodec { body.extend_from_slice(&expiry.to_be_bytes()); body.extend_from_slice(payload); - let tag = self - .mac_for(options.associated_data, &body) + match &self.keys { + Keys::Single(key) => Self::seal_rs1(key, options.associated_data, &body), + Keys::Ring { + keys, seal_mode, .. + } => match seal_mode { + SealMode::Rs1 { key_id } => Self::seal_rs1( + keys.get(key_id).expect("validated rs1 signing key"), + options.associated_data, + &body, + ), + SealMode::Rs2 { key_id } => Self::seal_rs2( + key_id, + keys.get(key_id).expect("validated rs2 signing key"), + options.associated_data, + &body, + ), + }, + } + } + + fn open_at( + &self, + sealed: &str, + associated_data: &[u8], + now_ms: i64, + ) -> Result, RequestStateError> { + match sealed.split('.').next() { + Some(VERSION_V1) => self.open_rs1_at(sealed, associated_data, now_ms), + Some(VERSION_V2) => self.open_rs2_at(sealed, associated_data, now_ms), + _ => Err(RequestStateError::MalformedFormat), + } + } + + fn seal_rs1(key: &[u8], associated_data: &[u8], body: &[u8]) -> String { + let tag = Self::mac_v1(key, associated_data, body) .finalize() .into_bytes(); + let mut out = String::with_capacity( + VERSION_V1.len() + 2 + Self::b64_len(body.len()) + Self::b64_len(tag.len()), + ); + out.push_str(VERSION_V1); + out.push('.'); + URL_SAFE_NO_PAD.encode_string(body, &mut out); + out.push('.'); + URL_SAFE_NO_PAD.encode_string(tag.as_slice(), &mut out); + out + } - // base64url without padding encodes 3 bytes as 4 chars, rounding up. - let b64_len = |n: usize| n.div_ceil(3) * 4; - let mut out = - String::with_capacity(VERSION.len() + 2 + b64_len(body.len()) + b64_len(tag.len())); - out.push_str(VERSION); + fn seal_rs2(kid: &str, key: &[u8], associated_data: &[u8], body: &[u8]) -> String { + let tag = Self::mac_v2(key, kid.as_bytes(), associated_data, body) + .finalize() + .into_bytes(); + let mut out = String::with_capacity( + VERSION_V2.len() + + 3 + + Self::b64_len(kid.len()) + + Self::b64_len(body.len()) + + Self::b64_len(tag.len()), + ); + out.push_str(VERSION_V2); + out.push('.'); + URL_SAFE_NO_PAD.encode_string(kid.as_bytes(), &mut out); out.push('.'); - URL_SAFE_NO_PAD.encode_string(&body, &mut out); + URL_SAFE_NO_PAD.encode_string(body, &mut out); out.push('.'); URL_SAFE_NO_PAD.encode_string(tag.as_slice(), &mut out); out } - fn open_at( + fn open_rs1_at( &self, sealed: &str, associated_data: &[u8], @@ -359,10 +627,83 @@ impl RequestStateCodec { let version = parts.next().ok_or(RequestStateError::MalformedFormat)?; let body_b64 = parts.next().ok_or(RequestStateError::MalformedFormat)?; let tag_b64 = parts.next().ok_or(RequestStateError::MalformedFormat)?; - if parts.next().is_some() || version != VERSION { + if parts.next().is_some() || version != VERSION_V1 { return Err(RequestStateError::MalformedFormat); } + if matches!( + &self.keys, + Keys::Ring { rs1_fallbacks, .. } if rs1_fallbacks.is_empty() + ) { + return Err(RequestStateError::UnknownKeyId); + } + + let body = URL_SAFE_NO_PAD + .decode(body_b64) + .map_err(|_| RequestStateError::InvalidEncoding)?; + let tag = URL_SAFE_NO_PAD + .decode(tag_b64) + .map_err(|_| RequestStateError::InvalidEncoding)?; + + // Authentication must not short-circuit based on tag bytes. + match &self.keys { + Keys::Single(key) => Self::mac_v1(key, associated_data, &body) + .verify_slice(&tag) + .map_err(|_| RequestStateError::IntegrityCheckFailed)?, + Keys::Ring { + keys, + rs1_fallbacks, + .. + } => { + // Try every fallback so timing does not identify the matching key. + let mut verified = false; + for kid in rs1_fallbacks { + let key = keys.get(kid).expect("validated rs1 fallback key"); + let matches = Self::mac_v1(key, associated_data, &body) + .verify_slice(&tag) + .is_ok(); + verified |= matches; + } + if !verified { + return Err(RequestStateError::IntegrityCheckFailed); + } + } + } + + Self::open_authenticated_body(body, now_ms) + } + + fn open_rs2_at( + &self, + sealed: &str, + associated_data: &[u8], + now_ms: i64, + ) -> Result, RequestStateError> { + let mut parts = sealed.split('.'); + let version = parts.next().ok_or(RequestStateError::MalformedFormat)?; + let kid_b64 = parts.next().ok_or(RequestStateError::MalformedFormat)?; + let body_b64 = parts.next().ok_or(RequestStateError::MalformedFormat)?; + let tag_b64 = parts.next().ok_or(RequestStateError::MalformedFormat)?; + if parts.next().is_some() || version != VERSION_V2 { + return Err(RequestStateError::MalformedFormat); + } + + if kid_b64.len() > MAX_ENCODED_KID_LEN { + return Err(RequestStateError::InvalidKeyId); + } + let kid = URL_SAFE_NO_PAD + .decode(kid_b64) + .map_err(|_| RequestStateError::InvalidEncoding)?; + if kid.is_empty() || kid.len() > MAX_KID_LEN { + return Err(RequestStateError::InvalidKeyId); + } + let kid = String::from_utf8(kid).map_err(|_| RequestStateError::InvalidKeyId)?; + + let key = match &self.keys { + Keys::Single(_) => return Err(RequestStateError::UnknownKeyId), + Keys::Ring { keys, .. } => keys.get(&kid).ok_or(RequestStateError::UnknownKeyId)?, + }; + let body = URL_SAFE_NO_PAD .decode(body_b64) .map_err(|_| RequestStateError::InvalidEncoding)?; @@ -370,11 +711,15 @@ impl RequestStateCodec { .decode(tag_b64) .map_err(|_| RequestStateError::InvalidEncoding)?; - // `verify_slice` compares in constant time and rejects wrong-length tags. - self.mac_for(associated_data, &body) + // Authentication must not short-circuit based on tag bytes. + Self::mac_v2(key, kid.as_bytes(), associated_data, &body) .verify_slice(&tag) .map_err(|_| RequestStateError::IntegrityCheckFailed)?; + Self::open_authenticated_body(body, now_ms) + } + + fn open_authenticated_body(body: Vec, now_ms: i64) -> Result, RequestStateError> { // The body is now authenticated, so its framing can be trusted. if body.len() < EXPIRY_LEN { return Err(RequestStateError::MalformedFormat); @@ -387,20 +732,55 @@ impl RequestStateCodec { Ok(body[EXPIRY_LEN..].to_vec()) } - /// Builds an HMAC keyed for request-state tags, pre-fed with the - /// domain-separation label, a length-prefixed `associated_data`, and the - /// body. The length prefix keeps the `associated_data`/`body` boundary - /// unambiguous so distinct inputs cannot collide. - fn mac_for(&self, associated_data: &[u8], body: &[u8]) -> HmacSha256 { - let mut mac = HmacSha256::new_from_slice(self.key.as_slice()) - .expect("HMAC accepts keys of any length"); - mac.update(DOMAIN); + fn mac_v1(key: &[u8], associated_data: &[u8], body: &[u8]) -> HmacSha256 { + let mut mac = HmacSha256::new_from_slice(key).expect("HMAC accepts keys of any length"); + mac.update(DOMAIN_V1); mac.update(&(associated_data.len() as u64).to_be_bytes()); mac.update(associated_data); mac.update(body); mac } + fn mac_v2(key: &[u8], kid: &[u8], associated_data: &[u8], body: &[u8]) -> HmacSha256 { + let mut mac = HmacSha256::new_from_slice(key).expect("HMAC accepts keys of any length"); + mac.update(DOMAIN_V2); + mac.update(&(kid.len() as u64).to_be_bytes()); + mac.update(kid); + mac.update(&(associated_data.len() as u64).to_be_bytes()); + mac.update(associated_data); + mac.update(body); + mac + } + + fn validate_key_length(key: &[u8]) -> Result<(), RequestStateError> { + if key.len() < Self::MIN_KEY_LENGTH { + return Err(RequestStateError::KeyTooShort { + minimum: Self::MIN_KEY_LENGTH, + actual: key.len(), + }); + } + Ok(()) + } + + fn validate_config_kid(kid: &str) -> Result<(), RequestStateError> { + if kid.is_empty() { + return Err(RequestStateError::InvalidKeyring( + "key id must not be empty", + )); + } + if kid.len() > MAX_KID_LEN { + return Err(RequestStateError::InvalidKeyring( + "key id exceeds 255 UTF-8 bytes", + )); + } + Ok(()) + } + + // Base64url without padding encodes at most three bytes as four characters. + fn b64_len(len: usize) -> usize { + len.div_ceil(3) * 4 + } + fn now_ms() -> i64 { chrono::Utc::now().timestamp_millis() } @@ -518,11 +898,11 @@ mod tests { } #[test] - fn wrong_version_prefix_is_malformed() { + fn unsupported_version_is_malformed() { let codec = RequestStateCodec::try_new(b"wrong-version-test-signing-key!!".to_vec()).unwrap(); let sealed = codec.seal(b"state"); - let bumped = sealed.replacen("rs1.", "rs2.", 1); + let bumped = sealed.replacen("rs1.", "rs3.", 1); assert!(matches!( codec.open(&bumped), Err(RequestStateError::MalformedFormat) @@ -665,4 +1045,450 @@ mod tests { )); } } + + mod key_rotation { + use super::*; + + const KEY_A: &[u8] = b"aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa"; + const KEY_B: &[u8] = b"bbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbb"; + const KEY_C: &[u8] = b"cccccccccccccccccccccccccccccccc"; + const KAT_KEY: &[u8] = b"0123456789abcdef0123456789abcdef"; + const KAT_KID: &str = "rotation-key-2026-08"; + const KAT_AD: &[u8] = b"user:alice|request:weather"; + const KAT_NOW_MS: i64 = 1_700_000_000_000; + + fn two_key_ring(active: &str) -> RequestStateCodec { + RequestStateCodec::new_with_keyring(active, [("a", KEY_A), ("b", KEY_B)]).unwrap() + } + + fn body(expiry: i64, payload: &[u8]) -> Vec { + let mut body = expiry.to_be_bytes().to_vec(); + body.extend_from_slice(payload); + body + } + + fn replace_segment(token: &str, index: usize, replacement: &str) -> String { + let mut parts: Vec = token.split('.').map(str::to_owned).collect(); + parts[index] = replacement.to_owned(); + parts.join(".") + } + + fn test_mac( + key: &[u8], + domain: &[u8], + kid: Option<&[u8]>, + associated_data: &[u8], + body: &[u8], + ) -> Vec { + let mut mac = HmacSha256::new_from_slice(key).expect("HMAC accepts keys of any length"); + mac.update(domain); + if let Some(kid) = kid { + mac.update(&(kid.len() as u64).to_be_bytes()); + mac.update(kid); + } + mac.update(&(associated_data.len() as u64).to_be_bytes()); + mac.update(associated_data); + mac.update(body); + mac.finalize().into_bytes().to_vec() + } + + fn raw_rs1(body: &[u8], tag: &[u8]) -> String { + format!( + "rs1.{}.{}", + URL_SAFE_NO_PAD.encode(body), + URL_SAFE_NO_PAD.encode(tag) + ) + } + + fn raw_rs2(kid: &[u8], body: &[u8], tag: &[u8]) -> String { + format!( + "rs2.{}.{}.{}", + URL_SAFE_NO_PAD.encode(kid), + URL_SAFE_NO_PAD.encode(body), + URL_SAFE_NO_PAD.encode(tag) + ) + } + + #[test] + fn rs1_known_answer_is_unchanged() { + // Independently checked with Python stdlib HMAC and Ruby OpenSSL. + let codec = RequestStateCodec::try_new(KAT_KEY).unwrap(); + let sealed = codec.seal_at( + b"step=2", + &SealOptions::new() + .associated_data(KAT_AD) + .ttl(Duration::from_secs(90)), + KAT_NOW_MS, + ); + assert_eq!( + sealed, + "rs1.AAABi8_mx5BzdGVwPTI.GQgS0X7mtSz8ZOy_kld2Zjuc4gAMGpBL74EghWI36IQ" + ); + } + + #[test] + fn rs2_known_answer_matches_wire_specification() { + // Independently checked with Python stdlib HMAC and Ruby OpenSSL. + let codec = RequestStateCodec::new_with_keyring(KAT_KID, [(KAT_KID, KAT_KEY)]).unwrap(); + let sealed = codec.seal_at( + b"step=2", + &SealOptions::new() + .associated_data(KAT_AD) + .ttl(Duration::from_secs(90)), + KAT_NOW_MS, + ); + assert_eq!( + sealed, + "rs2.cm90YXRpb24ta2V5LTIwMjYtMDg.AAABi8_mx5BzdGVwPTI.twv2acu7lKXqebyrmit-JrHHZm-BkKlBQEMMtL3lHk8" + ); + } + + #[test] + fn rs2_roundtrips_with_associated_data_and_ttl() { + let codec = two_key_ring("a"); + let options = SealOptions::new() + .associated_data(b"user:alice") + .ttl(Duration::from_secs(60)); + + let sealed = codec.seal_at(b"state", &options, 1_000); + assert!(sealed.starts_with("rs2.YQ.")); + assert_eq!( + codec.open_at(&sealed, b"user:alice", 30_000).unwrap(), + b"state" + ); + assert!(matches!( + codec.open_at(&sealed, b"user:bob", 30_000), + Err(RequestStateError::IntegrityCheckFailed) + )); + assert!(matches!( + codec.open_at(&sealed, b"user:alice", 70_000), + Err(RequestStateError::Expired) + )); + } + + #[test] + fn rotation_keeps_previous_rs2_key_verifiable() { + let signer_a = two_key_ring("a"); + let token_a = signer_a.seal(b"state-a"); + + let signer_b = two_key_ring("b"); + assert_eq!(signer_b.open(&token_a).unwrap(), b"state-a"); + let token_b = signer_b.seal(b"state-b"); + assert!(token_b.starts_with("rs2.Yg.")); + assert_eq!(signer_b.open(&token_b).unwrap(), b"state-b"); + + let without_a = RequestStateCodec::new_with_keyring("b", [("b", KEY_B)]).unwrap(); + assert!(matches!( + without_a.open(&token_a), + Err(RequestStateError::UnknownKeyId) + )); + } + + #[test] + fn rolling_migration_is_bidirectionally_compatible() { + let old = RequestStateCodec::try_new(KEY_A).unwrap(); + let transitional = + RequestStateCodec::new_with_keyring("new", [("old", KEY_A), ("new", KEY_B)]) + .unwrap() + .with_rs1_signing("old") + .unwrap(); + let promoted = + RequestStateCodec::new_with_keyring("new", [("old", KEY_A), ("new", KEY_B)]) + .unwrap() + .with_rs1_fallback("old") + .unwrap(); + let retired = RequestStateCodec::new_with_keyring("new", [("new", KEY_B)]).unwrap(); + + let old_token = old.seal(b"old"); + assert_eq!(transitional.open(&old_token).unwrap(), b"old"); + + let transitional_token = transitional.seal(b"transition"); + assert!(transitional_token.starts_with("rs1.")); + assert_eq!(old.open(&transitional_token).unwrap(), b"transition"); + assert_eq!(promoted.open(&transitional_token).unwrap(), b"transition"); + + let promoted_token = promoted.seal(b"promoted"); + assert!(promoted_token.starts_with("rs2.")); + assert_eq!(transitional.open(&promoted_token).unwrap(), b"promoted"); + assert_eq!(retired.open(&promoted_token).unwrap(), b"promoted"); + assert!(matches!( + retired.open(&old_token), + Err(RequestStateError::UnknownKeyId) + )); + } + + #[test] + fn multiple_legacy_fallbacks_are_supported() { + let legacy_a = RequestStateCodec::try_new(KEY_A).unwrap().seal(b"a"); + let legacy_b = RequestStateCodec::try_new(KEY_B).unwrap().seal(b"b"); + let legacy_c = RequestStateCodec::try_new(KEY_C).unwrap().seal(b"c"); + let ring = RequestStateCodec::new_with_keyring( + "new", + [("a", KEY_A), ("b", KEY_B), ("new", KEY_C)], + ) + .unwrap() + .with_rs1_fallback("a") + .unwrap() + .with_rs1_fallback("b") + .unwrap(); + + assert_eq!(ring.open(&legacy_a).unwrap(), b"a"); + assert_eq!(ring.open(&legacy_b).unwrap(), b"b"); + assert!(matches!( + ring.open(&legacy_c), + Err(RequestStateError::IntegrityCheckFailed) + )); + } + + #[test] + fn kid_is_authenticated_even_when_ids_share_key_bytes() { + let codec = + RequestStateCodec::new_with_keyring("a", [("a", KEY_A), ("b", KEY_A)]).unwrap(); + let sealed = codec.seal(b"state"); + let swapped = replace_segment(&sealed, 1, &URL_SAFE_NO_PAD.encode(b"b")); + assert!(matches!( + codec.open(&swapped), + Err(RequestStateError::IntegrityCheckFailed) + )); + } + + #[test] + fn independent_rs2_segment_tampering_is_rejected() { + let codec = two_key_ring("a"); + let sealed = codec.seal(b"state"); + + let swapped_kid = replace_segment(&sealed, 1, &URL_SAFE_NO_PAD.encode(b"b")); + let tampered_body = replace_segment(&sealed, 2, &URL_SAFE_NO_PAD.encode(b"changed")); + let tampered_tag = replace_segment(&sealed, 3, &URL_SAFE_NO_PAD.encode([0_u8; 32])); + + for tampered in [swapped_kid, tampered_body, tampered_tag] { + assert!(matches!( + codec.open(&tampered), + Err(RequestStateError::IntegrityCheckFailed) + )); + } + } + + #[test] + fn version_domains_are_cryptographically_separate() { + let token_body = body(0, b"state"); + + let wrong_v2_tag = test_mac(KEY_A, DOMAIN_V1, Some(b"a"), b"", &token_body); + let wrong_v2 = raw_rs2(b"a", &token_body, &wrong_v2_tag); + assert!(matches!( + two_key_ring("a").open(&wrong_v2), + Err(RequestStateError::IntegrityCheckFailed) + )); + + let wrong_v1_tag = test_mac(KEY_A, DOMAIN_V2, None, b"", &token_body); + let wrong_v1 = raw_rs1(&token_body, &wrong_v1_tag); + assert!(matches!( + RequestStateCodec::try_new(KEY_A).unwrap().open(&wrong_v1), + Err(RequestStateError::IntegrityCheckFailed) + )); + } + + #[test] + fn constructor_invariants_are_enforced() { + let empty = + RequestStateCodec::new_with_keyring("a", std::iter::empty::<(&str, &[u8])>()); + assert!(matches!(empty, Err(RequestStateError::InvalidKeyring(_)))); + + let duplicate = RequestStateCodec::new_with_keyring("a", [("a", KEY_A), ("a", KEY_B)]); + assert!(matches!( + duplicate, + Err(RequestStateError::InvalidKeyring(_)) + )); + + let empty_kid = RequestStateCodec::new_with_keyring("", [("", KEY_A)]); + assert!(matches!( + empty_kid, + Err(RequestStateError::InvalidKeyring(_)) + )); + + let oversized = "x".repeat(MAX_KID_LEN + 1); + let oversized_kid = + RequestStateCodec::new_with_keyring(oversized.clone(), [(oversized, KEY_A)]); + assert!(matches!( + oversized_kid, + Err(RequestStateError::InvalidKeyring(_)) + )); + + let absent_active = RequestStateCodec::new_with_keyring("missing", [("a", KEY_A)]); + assert!(matches!( + absent_active, + Err(RequestStateError::InvalidKeyring(_)) + )); + + let short_key = RequestStateCodec::new_with_keyring( + "a", + [("a", vec![0; RequestStateCodec::MIN_KEY_LENGTH - 1])], + ); + assert!(matches!( + short_key, + Err(RequestStateError::KeyTooShort { .. }) + )); + + assert!(matches!( + RequestStateCodec::try_new(KEY_A) + .unwrap() + .with_rs1_signing("a"), + Err(RequestStateError::InvalidKeyring(_)) + )); + assert!(matches!( + two_key_ring("a").with_rs1_signing("missing"), + Err(RequestStateError::InvalidKeyring(_)) + )); + assert!(matches!( + RequestStateCodec::try_new(KEY_A) + .unwrap() + .with_rs1_fallback("a"), + Err(RequestStateError::InvalidKeyring(_)) + )); + assert!(matches!( + two_key_ring("a").with_rs1_fallback("missing"), + Err(RequestStateError::InvalidKeyring(_)) + )); + } + + #[test] + fn maximum_length_kid_roundtrips_and_oversized_wire_kid_is_rejected() { + let max_kid = "x".repeat(MAX_KID_LEN); + let codec = + RequestStateCodec::new_with_keyring(max_kid.clone(), [(max_kid, KEY_A)]).unwrap(); + let token = codec.seal(b"state"); + assert_eq!(codec.open(&token).unwrap(), b"state"); + + let encoded_oversized = URL_SAFE_NO_PAD.encode("x".repeat(MAX_KID_LEN + 1)); + assert!(encoded_oversized.len() > MAX_ENCODED_KID_LEN); + let oversized = format!("rs2.{encoded_oversized}.!!!!.!!!!"); + assert!(matches!( + codec.open(&oversized), + Err(RequestStateError::InvalidKeyId) + )); + + let overlong_invalid_base64 = + format!("rs2.{}.!!!!.!!!!", "!".repeat(MAX_ENCODED_KID_LEN + 1)); + assert!(matches!( + codec.open(&overlong_invalid_base64), + Err(RequestStateError::InvalidKeyId) + )); + } + + #[test] + fn parser_rejects_invalid_tokens_and_unavailable_keys() { + let codec = two_key_ring("a"); + for malformed in [ + "rs2", + "rs2.YQ", + "rs2.YQ.body", + "rs2.YQ.body.tag.extra", + "rs3.YQ.body.tag", + ] { + assert!(matches!( + codec.open(malformed), + Err(RequestStateError::MalformedFormat) + )); + } + + let valid_rs2 = codec.seal(b"state"); + assert!(matches!( + codec.open(&replace_segment(&valid_rs2, 1, "!!!!")), + Err(RequestStateError::InvalidEncoding) + )); + assert!(matches!( + codec.open(&replace_segment(&valid_rs2, 1, "")), + Err(RequestStateError::InvalidKeyId) + )); + let non_utf8 = URL_SAFE_NO_PAD.encode([0xff]); + assert!(matches!( + codec.open(&replace_segment(&valid_rs2, 1, &non_utf8)), + Err(RequestStateError::InvalidKeyId) + )); + + assert!(matches!( + RequestStateCodec::try_new(KEY_A).unwrap().open(&valid_rs2), + Err(RequestStateError::UnknownKeyId) + )); + + let valid_rs1 = RequestStateCodec::try_new(KEY_A).unwrap().seal(b"state"); + assert!(matches!( + codec.open(&valid_rs1), + Err(RequestStateError::UnknownKeyId) + )); + } + + #[test] + fn authenticated_short_body_is_malformed() { + let short_body = b"short"; + let tag = RequestStateCodec::mac_v2(KEY_A, b"a", b"", short_body) + .finalize() + .into_bytes(); + let token = raw_rs2(b"a", short_body, tag.as_slice()); + assert!(matches!( + two_key_ring("a").open(&token), + Err(RequestStateError::MalformedFormat) + )); + } + + #[test] + fn serde_is_reached_only_after_integrity_verification() { + let codec = two_key_ring("a"); + let invalid_json = codec.seal(b"{not-json"); + let opened: Result = codec.open_json(&invalid_json); + assert!(matches!(opened, Err(RequestStateError::Deserialization(_)))); + + let valid_json = codec.seal(b"{}"); + let tampered_body = body(0, b"{not-json"); + let tampered = replace_segment(&valid_json, 2, &URL_SAFE_NO_PAD.encode(tampered_body)); + let opened: Result = codec.open_json(&tampered); + assert!(matches!( + opened, + Err(RequestStateError::IntegrityCheckFailed) + )); + } + + #[test] + fn ring_debug_redacts_all_key_material() { + let codec = two_key_ring("b").with_rs1_fallback("a").unwrap(); + let rendered = format!("{codec:?}"); + assert!(!rendered.contains(std::str::from_utf8(KEY_A).unwrap())); + assert!(!rendered.contains(std::str::from_utf8(KEY_B).unwrap())); + assert!(rendered.contains("redacted")); + } + + #[test] + fn mutated_tokens_do_not_panic() { + let codec = two_key_ring("a").with_rs1_fallback("a").unwrap(); + let valid = [ + RequestStateCodec::try_new(KEY_A).unwrap().seal(b"state"), + codec.seal(b"state"), + ]; + let mut corpus = vec![ + String::new(), + ".".to_owned(), + "...".to_owned(), + "rs1".to_owned(), + "rs2".to_owned(), + "💥".to_owned(), + "rs2.💥...".to_owned(), + ]; + + for token in valid { + for end in 0..=token.len() { + corpus.push(token[..end].to_owned()); + } + for index in 0..token.len() { + let mut bytes = token.as_bytes().to_vec(); + bytes[index] = b'!'; + corpus.push(String::from_utf8(bytes).expect("token is ASCII")); + } + } + + for candidate in corpus { + let _ = codec.open(&candidate); + let _ = codec.open_with(&candidate, b"context"); + } + } + } } From ec041b7430126fab4a0ebf537d0402c7ff33e397 Mon Sep 17 00:00:00 2001 From: Dale Seo <5466341+DaleSeo@users.noreply.github.com> Date: Sat, 29 Aug 2026 07:42:11 +0900 Subject: [PATCH 09/12] ci: reduce dependabot update noise (#1184) --- .github/dependabot.yml | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/.github/dependabot.yml b/.github/dependabot.yml index bdd4faffc..0fa34ce69 100644 --- a/.github/dependabot.yml +++ b/.github/dependabot.yml @@ -16,11 +16,11 @@ updates: - package-ecosystem: "github-actions" directory: "/" schedule: - interval: "daily" + interval: "weekly" labels: # Mark PRs as CI related change. - T-CI open-pull-requests-limit: 3 commit-message: prefix: "chore" - include: "scope" \ No newline at end of file + include: "scope" From 1afde5b4fcee215d2fd6dc437364d1c73ad75b22 Mon Sep 17 00:00:00 2001 From: Dale Seo <5466341+DaleSeo@users.noreply.github.com> Date: Sat, 29 Aug 2026 07:42:20 +0900 Subject: [PATCH 10/12] docs: use auto lifecycle in HTTP example (#1187) --- examples/clients/src/streamable_http.rs | 31 ++++++++++++++----------- 1 file changed, 18 insertions(+), 13 deletions(-) diff --git a/examples/clients/src/streamable_http.rs b/examples/clients/src/streamable_http.rs index 0c27fd358..d2c88dc7e 100644 --- a/examples/clients/src/streamable_http.rs +++ b/examples/clients/src/streamable_http.rs @@ -1,14 +1,15 @@ use anyhow::Result; use rmcp::{ - ServiceExt, - model::{CallToolRequestParams, ClientCapabilities, ClientInfo, Implementation}, + ClientLifecycleMode, ClientServiceExt, + model::{ + CallToolRequestParams, ClientCapabilities, ClientInfo, Implementation, ProtocolVersion, + }, transport::StreamableHttpClientTransport, }; use tracing_subscriber::{layer::SubscriberExt, util::SubscriberInitExt}; #[tokio::main] async fn main() -> Result<()> { - // Initialize logging tracing_subscriber::registry() .with( tracing_subscriber::EnvFilter::try_from_default_env() @@ -19,25 +20,29 @@ async fn main() -> Result<()> { let transport = StreamableHttpClientTransport::from_uri("http://localhost:8000/mcp"); let client_info = ClientInfo::new( ClientCapabilities::default(), - Implementation::new("test sse client", "0.0.1"), + Implementation::new("streamable-http-client", "0.0.1"), ); - let client = client_info.serve(transport).await.inspect_err(|e| { - tracing::error!("client error: {:?}", e); - })?; + let client = client_info + .serve_with_lifecycle( + transport, + ClientLifecycleMode::Auto { + preferred_versions: vec![ProtocolVersion::V_2026_07_28], + legacy_version: Some(ProtocolVersion::V_2025_11_25), + }, + ) + .await + .inspect_err(|e| { + tracing::error!("client error: {:?}", e); + })?; - // Initialize let server_info = client.peer_info(); tracing::info!("Connected to server: {server_info:#?}"); - // List tools let tools = client.list_tools(Default::default()).await?; tracing::info!("Available tools: {tools:#?}"); let tool_result = client - .call_tool( - CallToolRequestParams::new("increment") - .with_arguments(serde_json::json!({}).as_object().cloned().unwrap()), - ) + .call_tool(CallToolRequestParams::new("increment")) .await?; tracing::info!("Tool result: {tool_result:#?}"); client.cancel().await?; From 3a2ebbc92034f6711e4eac9e93f5de0423be8dfe Mon Sep 17 00:00:00 2001 From: "dependabot[bot]" <49699333+dependabot[bot]@users.noreply.github.com> Date: Sun, 30 Aug 2026 08:26:41 +0900 Subject: [PATCH 11/12] chore(deps): bump taiki-e/install-action from 2.85.13 to 2.86.7 (#1230) Bumps [taiki-e/install-action](https://github.com/taiki-e/install-action) from 2.85.13 to 2.86.7. - [Release notes](https://github.com/taiki-e/install-action/releases) - [Changelog](https://github.com/taiki-e/install-action/blob/main/CHANGELOG.md) - [Commits](https://github.com/taiki-e/install-action/compare/82cd3e7658a6f96c86c0234aeeda1748937cb0a1...b6ff580856c41316412a0b9b60540fbc6f8c82cc) --- updated-dependencies: - dependency-name: taiki-e/install-action dependency-version: 2.86.7 dependency-type: direct:production update-type: version-update:semver-minor ... Signed-off-by: dependabot[bot] Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com> --- .github/workflows/ci.yml | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index c72680b33..a43adf760 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -79,7 +79,7 @@ jobs: - uses: Swatinem/rust-cache@6323deb102c322ba6fcbdcafc7e3dddab59af2b6 # v2 - name: Install cargo-semver-checks - uses: taiki-e/install-action@82cd3e7658a6f96c86c0234aeeda1748937cb0a1 # v2.85.13 + uses: taiki-e/install-action@b6ff580856c41316412a0b9b60540fbc6f8c82cc # v2.86.7 with: tool: cargo-semver-checks @@ -133,7 +133,7 @@ jobs: - uses: Swatinem/rust-cache@6323deb102c322ba6fcbdcafc7e3dddab59af2b6 # v2 - name: Install cargo-public-api - uses: taiki-e/install-action@82cd3e7658a6f96c86c0234aeeda1748937cb0a1 # v2.85.13 + uses: taiki-e/install-action@b6ff580856c41316412a0b9b60540fbc6f8c82cc # v2.86.7 with: tool: cargo-public-api From e3bc6c71f5ee6d708fa79f860280a96788ebdf27 Mon Sep 17 00:00:00 2001 From: Jacob Magar Date: Mon, 31 Aug 2026 02:40:15 -0400 Subject: [PATCH 12/12] Add typed custom request responses --- crates/rmcp/CHANGELOG.md | 5 + crates/rmcp/Cargo.toml | 2 +- crates/rmcp/src/service.rs | 166 +++++++++++++---- crates/rmcp/src/service/client.rs | 10 +- crates/rmcp/src/service/server.rs | 10 +- crates/rmcp/src/transport.rs | 43 ++++- crates/rmcp/src/transport/async_rw.rs | 22 ++- crates/rmcp/src/transport/child_process.rs | 14 +- .../src/transport/common/client_side_sse.rs | 7 +- .../common/reqwest/streamable_http_client.rs | 11 +- .../rmcp/src/transport/common/unix_socket.rs | 12 +- .../src/transport/streamable_http_client.rs | 113 ++++++++---- crates/rmcp/src/transport/worker.rs | 96 +++++++++- crates/rmcp/tests/test_custom_request.rs | 172 +++++++++++++++++- .../test_streamable_http_json_response.rs | 78 +++++++- 15 files changed, 665 insertions(+), 96 deletions(-) diff --git a/crates/rmcp/CHANGELOG.md b/crates/rmcp/CHANGELOG.md index b7b7d542a..7bd421081 100644 --- a/crates/rmcp/CHANGELOG.md +++ b/crates/rmcp/CHANGELOG.md @@ -7,6 +7,11 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] +### Added + +- add `Peer::send_request_as` and option-aware typed request handles for + method-specific extension responses, plus an additive raw-response transport hook + ## [3.1.4](https://github.com/modelcontextprotocol/rust-sdk/compare/rmcp-v3.1.3...rmcp-v3.1.4) - 2026-08-18 ### Fixed diff --git a/crates/rmcp/Cargo.toml b/crates/rmcp/Cargo.toml index e554e8606..a4e2a42b4 100644 --- a/crates/rmcp/Cargo.toml +++ b/crates/rmcp/Cargo.toml @@ -293,7 +293,7 @@ path = "tests/test_streamable_http_priming.rs" [[test]] name = "test_streamable_http_json_response" -required-features = ["server", "client", "transport-streamable-http-server", "reqwest"] +required-features = ["server", "client", "transport-streamable-http-server", "transport-streamable-http-client-reqwest"] path = "tests/test_streamable_http_json_response.rs" [[test]] diff --git a/crates/rmcp/src/service.rs b/crates/rmcp/src/service.rs index 20fd2e981..c87612c20 100644 --- a/crates/rmcp/src/service.rs +++ b/crates/rmcp/src/service.rs @@ -94,6 +94,12 @@ pub enum ServiceError { /// The peer kept returning `input_required` beyond the configured round cap. #[error("input_required did not complete within {max_rounds} MRTR rounds")] InputRequiredRoundsExceeded { max_rounds: usize }, + /// A response could not be decoded as the type requested by the caller. + #[error("failed to deserialize response: {0}")] + ResponseDeserialization(#[source] serde_json::Error), + /// The selected transport cannot retain an extension response's raw result. + #[error("transport does not preserve raw response results required by typed requests")] + RawResponseUnavailable, } trait TransferObject: @@ -270,6 +276,23 @@ pub type RxJsonRpcMessage = JsonRpcMessage< ::PeerResp, ::PeerNot, >; +pub type RawRxJsonRpcMessage = + JsonRpcMessage<::PeerReq, serde_json::Value, ::PeerNot>; + +pub(crate) fn decode_peer_response( + message: RawRxJsonRpcMessage, +) -> Result, serde_json::Error> { + Ok(match message { + JsonRpcMessage::Request(request) => JsonRpcMessage::Request(request), + JsonRpcMessage::Response(response) => JsonRpcMessage::Response(JsonRpcResponse { + jsonrpc: response.jsonrpc, + id: response.id, + result: serde_json::from_value(response.result)?, + }), + JsonRpcMessage::Notification(notification) => JsonRpcMessage::Notification(notification), + JsonRpcMessage::Error(error) => JsonRpcMessage::Error(error), + }) +} #[cfg(not(feature = "local"))] pub trait Service: Send + Sync + 'static { @@ -528,8 +551,8 @@ type SubscriptionChannelMap = HashMap>; /// or wait for response by call [`RequestHandle::await_response`] #[derive(Debug)] #[non_exhaustive] -pub struct RequestHandle { - pub rx: tokio::sync::oneshot::Receiver>, +pub struct RequestHandle::PeerResp> { + pub rx: tokio::sync::oneshot::Receiver>, pub options: PeerRequestOptions, pub peer: Peer, pub id: RequestId, @@ -537,11 +560,11 @@ pub struct RequestHandle { progress_reset_rx: Option>, } -impl RequestHandle { +impl RequestHandle { pub const REQUEST_TIMEOUT_REASON: &str = "request timeout"; pub const REQUEST_MAX_TOTAL_TIMEOUT_REASON: &str = "maximum total timeout exceeded"; - pub async fn await_response(mut self) -> Result { + pub async fn await_response(mut self) -> Result { let timeout = self.options.timeout; let max_total_timeout = self.options.max_total_timeout; let reset_timeout_on_progress = self.options.reset_timeout_on_progress; @@ -603,7 +626,7 @@ impl RequestHandle { timeout: Option, max_total_timeout: Option, reset_timeout_on_progress: bool, - ) -> Result { + ) -> Result { let mut idle_sleep = timeout.map(|timeout| (timeout, Box::pin(tokio::time::sleep(timeout)))); let mut max_total_sleep = @@ -687,12 +710,52 @@ impl RequestHandle { } } +pub(crate) enum PendingResponder { + Standard(Responder>), + Typed(Box) + Send>), +} + +impl std::fmt::Debug for PendingResponder { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter.write_str(match self { + Self::Standard(_) => "PendingResponder::Standard", + Self::Typed(_) => "PendingResponder::Typed", + }) + } +} + +impl PendingResponder { + fn requires_raw_response(&self) -> bool { + matches!(self, Self::Typed(_)) + } + + fn send_error(self, error: ServiceError) { + match self { + Self::Standard(responder) => { + let _ = responder.send(Err(error)); + } + Self::Typed(responder) => responder(Err(error)), + } + } + + fn send_result(self, value: serde_json::Value) { + match self { + Self::Standard(responder) => { + let result = + serde_json::from_value(value).map_err(ServiceError::ResponseDeserialization); + let _ = responder.send(result); + } + Self::Typed(responder) => responder(Ok(value)), + } + } +} + #[derive(Debug)] pub(crate) enum PeerSinkMessage { Request { request: R::Req, id: RequestId, - responder: Responder>, + responder: PendingResponder, }, Notification { notification: R::Not, @@ -839,6 +902,44 @@ impl Peer { .await } + /// Send a request and deserialize its result directly as `T`. + /// + /// This is intended for extension methods whose result type is not part of + /// the role's core response union. Decoding the result as the caller's + /// concrete type avoids ambiguous `#[serde(untagged)]` union matching. + /// The caller is responsible for pairing the request method with its + /// correct result type. + pub async fn send_request_as(&self, request: R::Req) -> Result + where + T: serde::de::DeserializeOwned + Send + 'static, + { + self.send_request_as_with_option(request, PeerRequestOptions::no_options()) + .await? + .await_response() + .await + } + + /// Send a typed request with the same timeout, metadata, progress, and + /// cancellation lifecycle available to core requests. + pub async fn send_request_as_with_option( + &self, + request: R::Req, + options: PeerRequestOptions, + ) -> Result, ServiceError> + where + T: serde::de::DeserializeOwned + Send + 'static, + { + self.send_request_with_option_and_subscription(request, options, None, |responder| { + PendingResponder::Typed(Box::new(move |result| { + let result = result.and_then(|value| { + serde_json::from_value(value).map_err(ServiceError::ResponseDeserialization) + }); + let _ = responder.send(result); + })) + }) + .await + } + pub async fn send_cancellable_request( &self, request: R::Req, @@ -852,16 +953,22 @@ impl Peer { request: R::Req, options: PeerRequestOptions, ) -> Result, ServiceError> { - self.send_request_with_option_and_subscription(request, options, None) - .await + self.send_request_with_option_and_subscription( + request, + options, + None, + PendingResponder::Standard, + ) + .await } - async fn send_request_with_option_and_subscription( + async fn send_request_with_option_and_subscription( &self, mut request: R::Req, options: PeerRequestOptions, subscription_sender: Option>, - ) -> Result, ServiceError> { + wrap_responder: impl FnOnce(Responder>) -> PendingResponder, + ) -> Result, ServiceError> { R::enforce_request_association( &request, self.peer_info().as_deref(), @@ -906,7 +1013,7 @@ impl Peer { .send(PeerSinkMessage::Request { request, id: id.clone(), - responder, + responder: wrap_responder(responder), }) .await .is_err() @@ -943,6 +1050,7 @@ impl Peer { request, options, Some((sender, channel_capacity)), + PendingResponder::Standard, ) .await?; Ok((handle, receiver)) @@ -1342,8 +1450,7 @@ where tracing::info!(?peer_info, "Service initialized as server"); } - let mut local_responder_pool = - HashMap::>>::new(); + let mut local_responder_pool = HashMap::>::new(); let mut local_ct_pool = HashMap::::new(); let shared_service = Arc::new(service); // for return @@ -1356,7 +1463,7 @@ where let current_span = tracing::Span::current(); let handle = spawn_service_task(async move { let mut transport = transport.into_transport(); - let mut batch_messages = VecDeque::>::new(); + let mut batch_messages = VecDeque::>::new(); let mut send_task_set = tokio::task::JoinSet::::new(); let mut response_send_tasks = tokio::task::JoinSet::<()>::new(); #[derive(Debug)] @@ -1374,7 +1481,7 @@ where #[derive(Debug)] enum Event { ProxyMessage(PeerSinkMessage), - PeerMessage(RxJsonRpcMessage), + PeerMessage(RawRxJsonRpcMessage), ToSink(TxJsonRpcMessage), SendTaskResult(SendTaskResult), ResponseSendTaskResult(Result<(), tokio::task::JoinError>), @@ -1392,7 +1499,7 @@ where continue } } - m = transport.receive() => { + m = transport.receive_raw() => { if let Some(m) = m { Event::PeerMessage(m) } else { @@ -1440,7 +1547,7 @@ where Event::SendTaskResult(SendTaskResult::Request { id, result }) => { if let Err(e) = result && let Some(responder) = local_responder_pool.remove(&id) { - let _ = responder.send(Err(ServiceError::TransportSend(e))); + responder.send_error(ServiceError::TransportSend(e)); } } Event::SendTaskResult(SendTaskResult::Notification { @@ -1458,9 +1565,9 @@ where && let Some(request_id) = ¶m.request_id && let Some(responder) = local_responder_pool.remove(request_id) { tracing::info!(id = %request_id, reason = param.reason, "cancelled"); - let _response_result = responder.send(Err(ServiceError::Cancelled { + responder.send_error(ServiceError::Cancelled { reason: param.reason.clone(), - })); + }); } } Event::ResponseSendTaskResult(result) => { @@ -1495,6 +1602,10 @@ where id, responder, }) => { + if responder.requires_raw_response() && !T::preserves_raw_responses() { + responder.send_error(ServiceError::RawResponseUnavailable); + continue; + } local_responder_pool.insert(id.clone(), responder); let send = transport.send(JsonRpcMessage::request(request, id.clone())); { @@ -1604,9 +1715,9 @@ where if let Some(responder) = local_responder_pool.remove(request_id) { - let _ = responder.send(Err(ServiceError::Cancelled { + responder.send_error(ServiceError::Cancelled { reason: cancelled.reason.clone(), - })); + }); } } else if let Some(ct) = local_ct_pool.remove(request_id) { tracing::info!(id = %request_id, reason = cancelled.reason, "cancelled"); @@ -1637,8 +1748,7 @@ where && let Some(responder) = local_responder_pool.remove(&subscription_id) { - let _ = responder - .send(Err(ServiceError::SubscriptionLagged { capacity })); + responder.send_error(ServiceError::SubscriptionLagged { capacity }); } peer.unregister_subscription(&subscription_id); peer.try_cancel_request( @@ -1689,10 +1799,7 @@ where if let Some(responder) = remove_pending_request(&mut local_responder_pool, &id) { - let response_result = responder.send(Ok(result)); - if let Err(_error) = response_result { - tracing::warn!(%id, "Error sending response"); - } + responder.send_result(result); } } Event::PeerMessage(JsonRpcMessage::Error(JsonRpcError { error, id, .. })) => { @@ -1710,10 +1817,7 @@ where } else { ServiceError::McpError(error) }; - let _response_result = responder.send(Err(service_error)); - if let Err(_error) = _response_result { - tracing::warn!(%id, "Error sending response"); - } + responder.send_error(service_error); } } } diff --git a/crates/rmcp/src/service/client.rs b/crates/rmcp/src/service/client.rs index 520410fb1..6a54aa7a0 100644 --- a/crates/rmcp/src/service/client.rs +++ b/crates/rmcp/src/service/client.rs @@ -40,6 +40,8 @@ use crate::{ #[derive(Error, Debug)] #[non_exhaustive] pub enum ClientInitializeError { + #[error("failed to deserialize initialization response: {0}")] + ResponseDeserialization(#[source] serde_json::Error), #[error("expect initialized response, but received: {0:?}")] ExpectedInitResponse(Option), @@ -152,10 +154,12 @@ async fn expect_next_message( where T: Transport, { - transport - .receive() + let message = transport + .receive_raw() .await - .ok_or_else(|| ClientInitializeError::ConnectionClosed(context.to_string())) + .ok_or_else(|| ClientInitializeError::ConnectionClosed(context.to_string()))?; + super::decode_peer_response::(message) + .map_err(ClientInitializeError::ResponseDeserialization) } /// Helper function to expect a response from the stream, correlated to diff --git a/crates/rmcp/src/service/server.rs b/crates/rmcp/src/service/server.rs index c2f57a2a6..19719f593 100644 --- a/crates/rmcp/src/service/server.rs +++ b/crates/rmcp/src/service/server.rs @@ -82,6 +82,8 @@ impl ServiceRole for RoleServer { #[derive(Error, Debug)] #[non_exhaustive] pub enum ServerInitializeError { + #[error("failed to deserialize initialization response: {0}")] + ResponseDeserialization(#[source] serde_json::Error), #[error("expect initialized request, but received: {0:?}")] ExpectedInitializeRequest(Option), @@ -436,10 +438,12 @@ async fn expect_next_message( where T: Transport, { - transport - .receive() + let message = transport + .receive_raw() .await - .ok_or_else(|| ServerInitializeError::ConnectionClosed(context.to_string())) + .ok_or_else(|| ServerInitializeError::ConnectionClosed(context.to_string()))?; + super::decode_peer_response::(message) + .map_err(ServerInitializeError::ResponseDeserialization) } pub async fn serve_server_with_ct( diff --git a/crates/rmcp/src/transport.rs b/crates/rmcp/src/transport.rs index 06fde8e51..332f683f8 100644 --- a/crates/rmcp/src/transport.rs +++ b/crates/rmcp/src/transport.rs @@ -69,7 +69,7 @@ use std::{borrow::Cow, sync::Arc}; -use crate::service::{RxJsonRpcMessage, ServiceRole, TxJsonRpcMessage}; +use crate::service::{RawRxJsonRpcMessage, RxJsonRpcMessage, ServiceRole, TxJsonRpcMessage}; pub mod sink_stream; @@ -145,6 +145,47 @@ where /// Receive a message from the transport, this operation is sequential. fn receive(&mut self) -> impl Future>> + Send; + /// Whether this transport preserves response result bodies as raw JSON. + /// + /// Typed extension requests require this capability. Existing transports + /// default to `false` because decoding into the role response union can + /// irreversibly discard extension fields. + fn preserves_raw_responses() -> bool { + false + } + + /// Receive a message while preserving the raw JSON-RPC result value. + /// + /// Transports should override this method when they can retain the raw + /// response body. The default preserves compatibility for existing custom + /// transports, but an extension result already decoded through the role's + /// response union cannot recover information discarded by that union. + fn receive_raw(&mut self) -> impl Future>> + Send { + async move { + self.receive().await.map(|message| match message { + crate::model::JsonRpcMessage::Request(request) => { + crate::model::JsonRpcMessage::Request(request) + } + crate::model::JsonRpcMessage::Response(response) => { + crate::model::JsonRpcMessage::Response(crate::model::JsonRpcResponse { + jsonrpc: response.jsonrpc, + id: response.id, + result: serde_json::to_value(response.result).unwrap_or_else(|error| { + tracing::error!(%error, "failed to serialize peer response"); + serde_json::Value::Null + }), + }) + } + crate::model::JsonRpcMessage::Notification(notification) => { + crate::model::JsonRpcMessage::Notification(notification) + } + crate::model::JsonRpcMessage::Error(error) => { + crate::model::JsonRpcMessage::Error(error) + } + }) + } + } + /// Close the transport fn close(&mut self) -> impl Future> + Send; } diff --git a/crates/rmcp/src/transport/async_rw.rs b/crates/rmcp/src/transport/async_rw.rs index e18fd4fce..dbfa6318b 100644 --- a/crates/rmcp/src/transport/async_rw.rs +++ b/crates/rmcp/src/transport/async_rw.rs @@ -15,7 +15,9 @@ use tokio_util::{ use super::{IntoTransport, Transport}; use crate::{ model::ErrorData, - service::{RxJsonRpcMessage, ServiceRole, TxJsonRpcMessage}, + service::{ + RawRxJsonRpcMessage, RxJsonRpcMessage, ServiceRole, TxJsonRpcMessage, decode_peer_response, + }, }; #[non_exhaustive] @@ -104,6 +106,10 @@ where { type Error = std::io::Error; + fn preserves_raw_responses() -> bool { + true + } + fn send( &mut self, item: TxJsonRpcMessage, @@ -123,6 +129,18 @@ where } async fn receive(&mut self) -> Option> { + loop { + let message = self.receive_raw().await?; + match decode_peer_response::(message) { + Ok(message) => return Some(message), + Err(error) => { + tracing::debug!(%error, "Ignoring response with invalid result shape") + } + } + } + } + + async fn receive_raw(&mut self) -> Option> { loop { // `read_until` is not cancellation-safe on its own, and `receive` is // polled inside a `select!` in the service loop: an in-progress line @@ -155,7 +173,7 @@ where self.line_buf.clear(); continue; } - try_parse_with_compatibility::>(line, "receive") + try_parse_with_compatibility::>(line, "receive") }; self.line_buf.clear(); match parsed { diff --git a/crates/rmcp/src/transport/child_process.rs b/crates/rmcp/src/transport/child_process.rs index 6e19a0c3b..6b220e618 100644 --- a/crates/rmcp/src/transport/child_process.rs +++ b/crates/rmcp/src/transport/child_process.rs @@ -4,7 +4,9 @@ use futures::future::Future; use process_wrap::tokio::{ChildWrapper, CommandWrap}; use tokio::process::{ChildStderr, ChildStdin, ChildStdout}; -use super::{RxJsonRpcMessage, Transport, TxJsonRpcMessage, async_rw::AsyncRwTransport}; +use super::{ + RawRxJsonRpcMessage, RxJsonRpcMessage, Transport, TxJsonRpcMessage, async_rw::AsyncRwTransport, +}; use crate::RoleClient; const MAX_WAIT_ON_DROP_SECS: u64 = 3; @@ -168,6 +170,10 @@ impl TokioChildProcessBuilder { impl Transport for TokioChildProcess { type Error = std::io::Error; + fn preserves_raw_responses() -> bool { + true + } + fn send( &mut self, item: TxJsonRpcMessage, @@ -179,6 +185,12 @@ impl Transport for TokioChildProcess { self.transport.receive() } + fn receive_raw( + &mut self, + ) -> impl Future>> + Send { + self.transport.receive_raw() + } + fn close(&mut self) -> impl Future> + Send { self.graceful_shutdown() } diff --git a/crates/rmcp/src/transport/common/client_side_sse.rs b/crates/rmcp/src/transport/common/client_side_sse.rs index e668d63df..e0748ddbc 100644 --- a/crates/rmcp/src/transport/common/client_side_sse.rs +++ b/crates/rmcp/src/transport/common/client_side_sse.rs @@ -10,7 +10,7 @@ use futures::{Stream, StreamExt, stream::BoxStream}; use sse_stream::{Error as SseError, Sse, SseStream}; use thiserror::Error; -use crate::model::ServerJsonRpcMessage; +use crate::{RoleClient, service::RawRxJsonRpcMessage}; pub type BoxedSseResponse = BoxStream<'static, Result>; @@ -371,7 +371,7 @@ impl Stream for SseAutoReconnectStream where R: SseStreamReconnect, { - type Item = Result; + type Item = Result, R::Error>; fn poll_next( mut self: Pin<&mut Self>, cx: &mut std::task::Context<'_>, @@ -406,7 +406,8 @@ where } } if let Some(data) = sse.data { - match serde_json::from_str::(&data) { + match serde_json::from_str::>(&data) + { Err(e) => { // Downgrade to debug to avoid noisy logs when servers emit // non-JSON payloads as message frames. Include last_event_id diff --git a/crates/rmcp/src/transport/common/reqwest/streamable_http_client.rs b/crates/rmcp/src/transport/common/reqwest/streamable_http_client.rs index b0b5090b5..56e9b5e0b 100644 --- a/crates/rmcp/src/transport/common/reqwest/streamable_http_client.rs +++ b/crates/rmcp/src/transport/common/reqwest/streamable_http_client.rs @@ -49,6 +49,10 @@ fn parse_json_rpc_error(body: &str) -> Option { impl StreamableHttpClient for reqwest::Client { type Error = reqwest::Error; + fn preserves_raw_responses() -> bool { + true + } + async fn get_stream( &self, uri: Arc, @@ -311,8 +315,11 @@ impl StreamableHttpClient for reqwest::Client { // Try to parse as a valid JSON-RPC message. If the body is // malformed (e.g. a 200 response to a notification that lacks // an `id` field), treat it as accepted rather than failing. - match response.json::().await { - Ok(message) => Ok(StreamableHttpPostResponse::Json(message, session_id)), + match response + .json::>() + .await + { + Ok(message) => Ok(StreamableHttpPostResponse::RawJson(message, session_id)), Err(e) => { tracing::warn!( "could not parse JSON response as ServerJsonRpcMessage, treating as accepted: {e}" diff --git a/crates/rmcp/src/transport/common/unix_socket.rs b/crates/rmcp/src/transport/common/unix_socket.rs index e3b34898f..7dc29f44e 100644 --- a/crates/rmcp/src/transport/common/unix_socket.rs +++ b/crates/rmcp/src/transport/common/unix_socket.rs @@ -10,7 +10,7 @@ use sse_stream::Sse; use tokio::net::UnixStream; use crate::{ - model::{ClientJsonRpcMessage, ServerJsonRpcMessage}, + model::ClientJsonRpcMessage, transport::{ common::{ client_side_sse::{DEFAULT_MAX_SSE_EVENT_SIZE, bounded_sse_stream}, @@ -164,6 +164,10 @@ fn apply_custom_headers( impl StreamableHttpClient for UnixSocketHttpClient { type Error = UnixSocketError; + fn preserves_raw_responses() -> bool { + true + } + async fn post_message( &self, uri: Arc, @@ -321,8 +325,10 @@ impl StreamableHttpClient for UnixSocketHttpClient { .await .map_err(|e| StreamableHttpError::Client(UnixSocketError::Hyper(e)))? .to_bytes(); - match serde_json::from_slice::(&body) { - Ok(message) => Ok(StreamableHttpPostResponse::Json(message, session_id)), + match serde_json::from_slice::>( + &body, + ) { + Ok(message) => Ok(StreamableHttpPostResponse::RawJson(message, session_id)), Err(e) => { tracing::warn!( "could not parse JSON response as ServerJsonRpcMessage, treating as accepted: {e}" diff --git a/crates/rmcp/src/transport/streamable_http_client.rs b/crates/rmcp/src/transport/streamable_http_client.rs index d702bc1ca..a8e59a244 100644 --- a/crates/rmcp/src/transport/streamable_http_client.rs +++ b/crates/rmcp/src/transport/streamable_http_client.rs @@ -27,7 +27,7 @@ use crate::{ InitializedNotification, JsonObject, ProtocolVersion, RequestId, ServerJsonRpcMessage, ServerResult, }, - service::InboundStreamOrigin, + service::{InboundStreamOrigin, RawRxJsonRpcMessage, decode_peer_response}, transport::{ common::{client_side_sse::SseAutoReconnectStream, mcp_headers}, worker::{ @@ -253,6 +253,8 @@ pub enum StreamableHttpProtocolError { pub enum StreamableHttpPostResponse { Accepted, Json(ServerJsonRpcMessage, Option), + /// A JSON response whose result body has not been decoded through the core union. + RawJson(RawRxJsonRpcMessage, Option), Sse(BoxedSseStream, Option), } @@ -261,6 +263,7 @@ impl std::fmt::Debug for StreamableHttpPostResponse { match self { Self::Accepted => write!(f, "Accepted"), Self::Json(arg0, arg1) => f.debug_tuple("Json").field(arg0).field(arg1).finish(), + Self::RawJson(arg0, arg1) => f.debug_tuple("RawJson").field(arg0).field(arg1).finish(), Self::Sse(_, arg1) => f.debug_tuple("Sse").field(arg1).finish(), } } @@ -275,6 +278,9 @@ impl StreamableHttpPostResponse { { match self { Self::Json(message, session_id) => Ok((message, session_id)), + Self::RawJson(message, session_id) => { + Ok((decode_peer_response::(message)?, session_id)) + } Self::Sse(mut stream, session_id) => { while let Some(event) = stream.next().await { let event = event?; @@ -283,7 +289,8 @@ impl StreamableHttpPostResponse { continue; } - let message: ServerJsonRpcMessage = serde_json::from_str(&payload)?; + let raw: RawRxJsonRpcMessage = serde_json::from_str(&payload)?; + let message = decode_peer_response::(raw)?; if matches!(message, ServerJsonRpcMessage::Response(_)) { return Ok((message, session_id)); @@ -311,6 +318,7 @@ impl StreamableHttpPostResponse { { match self { Self::Json(message, ..) => Ok(message), + Self::RawJson(message, ..) => Ok(decode_peer_response::(message)?), got => Err(StreamableHttpError::UnexpectedServerResponse( format!("expect json, got {got:?}").into(), )), @@ -325,6 +333,7 @@ impl StreamableHttpPostResponse { Self::Accepted => Ok(()), // Tolerate servers that return 200 with JSON for notifications Self::Json(..) => Ok(()), + Self::RawJson(..) => Ok(()), got => Err(StreamableHttpError::UnexpectedServerResponse( format!("expect accepted or json, got {got:?}").into(), )), @@ -381,6 +390,10 @@ pub(super) fn legacy_discover_response( /// handle the cancellation using state owned by that stream. pub trait StreamableHttpClient: Clone + Send + 'static { type Error: std::error::Error + Send + Sync + 'static; + /// Whether JSON responses are returned as [`StreamableHttpPostResponse::RawJson`]. + fn preserves_raw_responses() -> bool { + false + } fn post_message( &self, uri: Arc, @@ -649,17 +662,27 @@ impl StreamableHttpClientWorker { } } - fn server_response_id(message: &ServerJsonRpcMessage) -> Option<&RequestId> { + fn server_response_id( + message: &crate::model::JsonRpcMessage< + crate::model::ServerRequest, + Resp, + crate::model::ServerNotification, + >, + ) -> Option<&RequestId> { match message { - ServerJsonRpcMessage::Response(response) => Some(&response.id), - ServerJsonRpcMessage::Error(error) => error.id.as_ref(), + crate::model::JsonRpcMessage::Response(response) => Some(&response.id), + crate::model::JsonRpcMessage::Error(error) => error.id.as_ref(), _ => None, } } - fn clear_stream_response_pending( + fn clear_stream_response_pending( pending_stream_response_ids: &mut HashSet, - message: &ServerJsonRpcMessage, + message: &crate::model::JsonRpcMessage< + crate::model::ServerRequest, + Resp, + crate::model::ServerNotification, + >, ) -> Option { let response_id = Self::server_response_id(message)?; if let Some(id) = pending_stream_response_ids.take(response_id) { @@ -670,7 +693,7 @@ impl StreamableHttpClientWorker { } async fn drain_queued_stream_messages( - sse_worker_rx: &mut tokio::sync::mpsc::Receiver, + sse_worker_rx: &mut tokio::sync::mpsc::Receiver>, context: &mut super::worker::WorkerContext, pending_stream_response_ids: &mut HashSet, ) -> Result<(), WorkerQuitReason>> { @@ -763,8 +786,9 @@ impl StreamableHttpClientWorker { custom_headers: HashMap, max_sse_event_size: usize, retry_config: Arc, - ) -> impl Stream>> + Send + 'static - { + ) -> impl Stream, StreamableHttpError>> + + Send + + 'static { SseAutoReconnectStream::new_after_event_id( stream, StreamableHttpClientReconnect { @@ -792,7 +816,8 @@ impl StreamableHttpClientWorker { custom_headers: HashMap, max_sse_event_size: usize, retry_config: Arc, - ) -> BoxStream<'static, Result>> { + ) -> BoxStream<'static, Result, StreamableHttpError>> + { Self::reconnecting_sse_to_jsonrpc( stream, client, @@ -809,9 +834,9 @@ impl StreamableHttpClientWorker { async fn run_response_stream( mut sse_stream: BoxStream< 'static, - Result>, + Result, StreamableHttpError>, >, - sse_worker_tx: tokio::sync::mpsc::Sender, + sse_worker_tx: tokio::sync::mpsc::Sender>, origin: InboundStreamOrigin, request_ct: CancellationToken, stream_ct: CancellationToken, @@ -831,10 +856,24 @@ impl StreamableHttpClientWorker { } } - async fn execute_sse_stream( - sse_stream: impl Stream>> - + Send, - sse_worker_tx: tokio::sync::mpsc::Sender, + async fn execute_sse_stream( + sse_stream: impl Stream< + Item = Result< + crate::model::JsonRpcMessage< + crate::model::ServerRequest, + Resp, + crate::model::ServerNotification, + >, + StreamableHttpError, + >, + > + Send, + sse_worker_tx: tokio::sync::mpsc::Sender< + crate::model::JsonRpcMessage< + crate::model::ServerRequest, + Resp, + crate::model::ServerNotification, + >, + >, origin: InboundStreamOrigin, close_on_response: bool, ct: CancellationToken, @@ -855,12 +894,12 @@ impl StreamableHttpClientWorker { }; // SEP-2260: mark inbound requests with the stream they arrived on // for the client receive-side association check. - if let ServerJsonRpcMessage::Request(request) = &mut message { + if let crate::model::JsonRpcMessage::Request(request) = &mut message { request.request.extensions_mut().insert(origin.clone()); } let is_response = matches!( message, - ServerJsonRpcMessage::Response(_) | ServerJsonRpcMessage::Error(_) + crate::model::JsonRpcMessage::Response(_) | crate::model::JsonRpcMessage::Error(_) ); let yield_result = sse_worker_tx.send(message).await; if yield_result.is_err() { @@ -887,7 +926,7 @@ impl StreamableHttpClientWorker { session_id: Arc, config: &StreamableHttpClientTransportConfig, protocol_headers: HashMap, - sse_worker_tx: tokio::sync::mpsc::Sender, + sse_worker_tx: tokio::sync::mpsc::Sender>, transport_task_ct: CancellationToken, ) { let uri = config.uri.clone(); @@ -1018,6 +1057,9 @@ impl StreamableHttpClientWorker { impl Worker for StreamableHttpClientWorker { type Role = RoleClient; type Error = StreamableHttpError; + fn preserves_raw_responses() -> bool { + C::preserves_raw_responses() + } fn is_control_message(message: &ClientJsonRpcMessage) -> bool { match message { ClientJsonRpcMessage::Response(_) | ClientJsonRpcMessage::Error(_) => true, @@ -1049,7 +1091,7 @@ impl Worker for StreamableHttpClientWorker { ) -> Result<(), WorkerQuitReason> { let channel_buffer_capacity = self.config.channel_buffer_capacity; let (sse_worker_tx, mut sse_worker_rx) = - tokio::sync::mpsc::channel::(channel_buffer_capacity); + tokio::sync::mpsc::channel::>(channel_buffer_capacity); let config = self.config.clone(); let transport_task_ct = context.cancellation_token.clone(); let _drop_guard = transport_task_ct.clone().drop_guard(); @@ -1172,7 +1214,7 @@ impl Worker for StreamableHttpClientWorker { StartPost(WorkerSendRequest>), PostResult(PostResult), RecoveryTimeout, - ServerMessage(ServerJsonRpcMessage), + ServerMessage(RawRxJsonRpcMessage), StreamResult { request_id: Option, result: Result<(), StreamableHttpError>, @@ -1639,6 +1681,10 @@ impl Worker for StreamableHttpClientWorker { context.send_to_handler(message).await?; Ok(()) } + Ok(StreamableHttpPostResponse::RawJson(message, ..)) => { + context.send_to_handler(message).await?; + Ok(()) + } Ok(StreamableHttpPostResponse::Sse(stream, ..)) => { let stream_request_id = request_id; let sse_stream = Self::response_sse_to_jsonrpc( @@ -1710,7 +1756,7 @@ impl Worker for StreamableHttpClientWorker { } let _ = responder.send(send_result); } - Event::ServerMessage(mut json_rpc_message) => { + Event::ServerMessage(json_rpc_message) => { // Match against all pending requests, not just open response streams. if let Some(request_id) = Self::clear_stream_response_pending( &mut pending_stream_response_ids, @@ -1718,11 +1764,15 @@ impl Worker for StreamableHttpClientWorker { ) { drop(request_stream_cancellations.remove(&request_id)); } - cache_tools_from_response( - &mut tool_header_cache, - &mut json_rpc_message, - &negotiated_version, - ); + if let Ok(mut typed_message) = + decode_peer_response::(json_rpc_message.clone()) + { + cache_tools_from_response( + &mut tool_header_cache, + &mut typed_message, + &negotiated_version, + ); + } // send the message to the handler if let Err(e) = context.send_to_handler(json_rpc_message).await { break 'main_loop Err(e); @@ -2299,7 +2349,7 @@ mod tests { let message = stream.next().await.expect("replayed response").unwrap(); - assert!(matches!(message, ServerJsonRpcMessage::Response(_))); + assert!(matches!(message, crate::model::JsonRpcMessage::Response(_))); assert_eq!( reconnects.lock().expect("lock reconnects").as_slice(), &[(None, Some("event-0".into()))] @@ -2417,7 +2467,8 @@ mod tests { &[(None, Some("e1".into()))], "the request must arrive on the resumed connection" ); - let ServerJsonRpcMessage::Request(request) = rx.recv().await.expect("request forwarded") + let crate::model::JsonRpcMessage::Request(request) = + rx.recv().await.expect("request forwarded") else { panic!("expected request first"); }; @@ -2428,7 +2479,7 @@ mod tests { ); assert!(matches!( rx.recv().await.expect("response forwarded"), - ServerJsonRpcMessage::Response(_) + crate::model::JsonRpcMessage::Response(_) )); } diff --git a/crates/rmcp/src/transport/worker.rs b/crates/rmcp/src/transport/worker.rs index e32da7d7d..6f2256f07 100644 --- a/crates/rmcp/src/transport/worker.rs +++ b/crates/rmcp/src/transport/worker.rs @@ -14,7 +14,9 @@ use tracing::{Instrument, Level}; use super::{IntoTransport, Transport}; use crate::{ model::{CancelledNotification, JsonRpcMessage, RequestId}, - service::{RxJsonRpcMessage, ServiceRole, TxJsonRpcMessage}, + service::{ + RawRxJsonRpcMessage, RxJsonRpcMessage, ServiceRole, TxJsonRpcMessage, decode_peer_response, + }, }; #[derive(Debug, thiserror::Error)] @@ -35,6 +37,8 @@ pub enum WorkerQuitReason { HandlerTerminated, #[error("Worker idle timeout after {}ms", _0.as_millis())] IdleTimeout(Duration), + #[error("failed to serialize peer response: {0}")] + ResponseSerialization(#[source] serde_json::Error), } impl WorkerQuitReason { @@ -77,6 +81,10 @@ pub trait Worker: Sized + Send + 'static { fn supports_request_cancellation() -> bool { false } + /// Whether this worker sends inbound responses through the lossless raw path. + fn preserves_raw_responses() -> bool { + false + } } type RequestCancellations = Arc>>>; @@ -167,7 +175,7 @@ impl WorkerSendRequest { } pub struct WorkerTransport { - rx: tokio::sync::mpsc::Receiver>, + rx: tokio::sync::mpsc::Receiver>, send_service: tokio::sync::mpsc::Sender>, control_send_service: tokio::sync::mpsc::Sender>, request_cancellations: RequestCancellations, @@ -214,12 +222,41 @@ impl WorkerTransport { 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) = + let (raw_to_handler_tx, from_transport_rx) = tokio::sync::mpsc::channel::< + RawRxJsonRpcMessage, + >(config.channel_buffer_capacity); + let (to_handler_tx, mut legacy_from_transport_rx) = tokio::sync::mpsc::channel::>(config.channel_buffer_capacity); + let legacy_raw_tx = raw_to_handler_tx.clone(); + tokio::spawn(async move { + while let Some(message) = legacy_from_transport_rx.recv().await { + let raw = match message { + JsonRpcMessage::Request(request) => JsonRpcMessage::Request(request), + JsonRpcMessage::Response(response) => JsonRpcMessage::Response( + crate::model::JsonRpcResponse { + jsonrpc: response.jsonrpc, + id: response.id, + result: serde_json::to_value(response.result).unwrap_or_else(|error| { + tracing::error!(%error, "failed to serialize legacy worker response"); + serde_json::Value::Null + }), + }, + ), + JsonRpcMessage::Notification(notification) => { + JsonRpcMessage::Notification(notification) + } + JsonRpcMessage::Error(error) => JsonRpcMessage::Error(error), + }; + if legacy_raw_tx.send(raw).await.is_err() { + break; + } + } + }); let request_cancellations = RequestCancellations::default(); let control_generation = Arc::new(AtomicU64::new(0)); let context = WorkerContext { to_handler_tx, + raw_to_handler_tx, from_handler_rx, control_from_handler_rx, control_generation: control_generation.clone(), @@ -248,6 +285,9 @@ impl WorkerTransport { WorkerQuitReason::Fatal { error, context } => { tracing::error!("worker quit with fatal: {error}, when {context}"); } + WorkerQuitReason::ResponseSerialization(error) => { + tracing::error!(%error, "worker failed to serialize peer response"); + } }) .inspect(|_| { tracing::debug!("worker quit"); @@ -298,6 +338,7 @@ pub struct WorkerContext { pub control_from_handler_rx: tokio::sync::mpsc::Receiver>, pub cancellation_token: CancellationToken, control_generation: Arc, + raw_to_handler_tx: tokio::sync::mpsc::Sender>, } impl WorkerContext { @@ -320,11 +361,37 @@ impl WorkerContext { .wrapping_add(1) } - pub async fn send_to_handler( + pub async fn send_to_handler( &mut self, - item: RxJsonRpcMessage, - ) -> Result<(), WorkerQuitReason> { - self.to_handler_tx + item: crate::model::JsonRpcMessage< + ::PeerReq, + Resp, + ::PeerNot, + >, + ) -> Result<(), WorkerQuitReason> + where + Resp: serde::Serialize, + { + let item = match item { + crate::model::JsonRpcMessage::Request(request) => { + crate::model::JsonRpcMessage::Request(request) + } + crate::model::JsonRpcMessage::Response(response) => { + crate::model::JsonRpcMessage::Response(crate::model::JsonRpcResponse { + jsonrpc: response.jsonrpc, + id: response.id, + result: serde_json::to_value(response.result) + .map_err(WorkerQuitReason::ResponseSerialization)?, + }) + } + crate::model::JsonRpcMessage::Notification(notification) => { + crate::model::JsonRpcMessage::Notification(notification) + } + crate::model::JsonRpcMessage::Error(error) => { + crate::model::JsonRpcMessage::Error(error) + } + }; + self.raw_to_handler_tx .send(item) .await .map_err(|_| WorkerQuitReason::HandlerTerminated) @@ -343,6 +410,10 @@ impl WorkerContext { impl Transport for WorkerTransport { type Error = W::Error; + fn preserves_raw_responses() -> bool { + W::preserves_raw_responses() + } + fn send( &mut self, item: TxJsonRpcMessage, @@ -397,6 +468,17 @@ impl Transport for WorkerTransport { } } async fn receive(&mut self) -> Option> { + loop { + let message = self.rx.recv().await?; + match decode_peer_response::(message) { + Ok(message) => return Some(message), + Err(error) => { + tracing::debug!(%error, "Ignoring response with invalid result shape") + } + } + } + } + async fn receive_raw(&mut self) -> Option> { self.rx.recv().await } async fn close(&mut self) -> Result<(), Self::Error> { diff --git a/crates/rmcp/tests/test_custom_request.rs b/crates/rmcp/tests/test_custom_request.rs index 66ee1ff99..6ccc5992f 100644 --- a/crates/rmcp/tests/test_custom_request.rs +++ b/crates/rmcp/tests/test_custom_request.rs @@ -1,18 +1,49 @@ #![cfg(not(feature = "local"))] -use std::sync::Arc; +use std::{future::Future, sync::Arc, time::Duration}; use rmcp::{ - ClientHandler, ServerHandler, ServiceExt, + ClientHandler, RoleClient, ServerHandler, ServiceExt, model::{ - ClientRequest, ClientResult, CustomRequest, CustomResult, ServerRequest, ServerResult, + ClientRequest, ClientResult, CustomRequest, CustomResult, ErrorCode, ErrorData, + PingRequest, ServerRequest, ServerResult, }, + service::{PeerRequestOptions, ServiceError}, + transport::Transport, }; +use serde::Deserialize; use serde_json::json; use tokio::sync::{Mutex, Notify}; use tracing_subscriber::{layer::SubscriberExt, util::SubscriberInitExt}; type CustomRequestPayload = (String, Option); +struct LegacyTypedTransport; + +impl Transport for LegacyTypedTransport { + type Error = std::convert::Infallible; + + fn send( + &mut self, + _item: rmcp::service::TxJsonRpcMessage, + ) -> impl Future> + Send + 'static { + std::future::ready(Ok(())) + } + + async fn receive(&mut self) -> Option> { + None + } + + fn close(&mut self) -> impl Future> + Send { + std::future::ready(Ok(())) + } +} + +#[tokio::test] +async fn existing_transport_implementations_get_the_raw_receive_compatibility_default() { + let mut transport = LegacyTypedTransport; + assert!(transport.receive_raw().await.is_none()); +} + struct CustomRequestServer { receive_signal: Arc, payload: Arc>>, @@ -86,6 +117,141 @@ async fn test_custom_client_request_reaches_server() -> anyhow::Result<()> { Ok(()) } +#[derive(Debug, Deserialize, PartialEq)] +#[serde(rename_all = "camelCase")] +struct SkillsListResult { + result_type: String, + skills: Vec, + #[serde(rename = "_meta")] + meta: serde_json::Value, +} + +struct TypedCustomRequestServer; + +impl ServerHandler for TypedCustomRequestServer { + async fn on_custom_request( + &self, + request: CustomRequest, + _context: rmcp::service::RequestContext, + ) -> Result { + if request.method == "skills/malformed" { + return Ok(CustomResult::new(json!({ + "resultType": "complete", + "skills": "not-an-array", + "_meta": {} + }))); + } + if request.method == "skills/error" { + return Err(ErrorData::new( + ErrorCode::INVALID_PARAMS, + "invalid skills request", + None, + )); + } + if request.method == "skills/slow" { + tokio::time::sleep(Duration::from_millis(100)).await; + } + Ok(CustomResult::new(json!({ + "resultType": "complete", + "skills": ["example"], + "_meta": {"io.modelcontextprotocol/serverInfo": {"name": "skills"}} + }))) + } +} + +async fn typed_test_client() -> anyhow::Result> +{ + let (server_transport, client_transport) = tokio::io::duplex(4096); + tokio::spawn(async move { + let server = TypedCustomRequestServer.serve(server_transport).await?; + server.waiting().await?; + anyhow::Ok(()) + }); + Ok(().serve(client_transport).await?) +} + +#[tokio::test] +async fn typed_custom_request_bypasses_the_core_response_union() -> anyhow::Result<()> { + let client = typed_test_client().await?; + let response: SkillsListResult = client + .send_request_as(ClientRequest::CustomRequest(CustomRequest::new( + "skills/list", + Some(json!({})), + ))) + .await?; + + assert_eq!(response.result_type, "complete"); + assert_eq!(response.skills, ["example"]); + assert_eq!( + response.meta["io.modelcontextprotocol/serverInfo"]["name"], + "skills" + ); + + client.cancel().await?; + Ok(()) +} + +#[tokio::test] +async fn typed_custom_requests_preserve_errors_correlation_and_recovery() -> anyhow::Result<()> { + let client = typed_test_client().await?; + + let malformed = client + .send_request_as::(ClientRequest::CustomRequest(CustomRequest::new( + "skills/malformed", + None, + ))) + .await; + assert!(matches!( + malformed, + Err(ServiceError::ResponseDeserialization(_)) + )); + + let protocol_error = client + .send_request_as::(ClientRequest::CustomRequest(CustomRequest::new( + "skills/error", + None, + ))) + .await; + assert!(matches!(protocol_error, Err(ServiceError::McpError(_)))); + + let typed = client.send_request_as::(ClientRequest::CustomRequest( + CustomRequest::new("skills/list", None), + )); + let standard = client.send_request(ClientRequest::PingRequest(PingRequest::default())); + let (typed, standard) = tokio::join!(typed, standard); + assert_eq!(typed?.skills, ["example"]); + assert!(matches!(standard?, ServerResult::EmptyResult(_))); + + client.cancel().await?; + Ok(()) +} + +#[tokio::test] +async fn typed_custom_requests_use_standard_timeout_and_cancellation() -> anyhow::Result<()> { + let client = typed_test_client().await?; + let handle = client + .send_request_as_with_option::( + ClientRequest::CustomRequest(CustomRequest::new("skills/slow", None)), + PeerRequestOptions::with_timeout(Duration::from_millis(10)), + ) + .await?; + assert!(matches!( + handle.await_response().await, + Err(ServiceError::Timeout { .. }) + )); + + let recovered: SkillsListResult = client + .send_request_as(ClientRequest::CustomRequest(CustomRequest::new( + "skills/list", + None, + ))) + .await?; + assert_eq!(recovered.skills, ["example"]); + + client.cancel().await?; + Ok(()) +} + struct CustomRequestClient { receive_signal: Arc, payload: Arc>>, diff --git a/crates/rmcp/tests/test_streamable_http_json_response.rs b/crates/rmcp/tests/test_streamable_http_json_response.rs index b1c09f512..5deb20b8d 100644 --- a/crates/rmcp/tests/test_streamable_http_json_response.rs +++ b/crates/rmcp/tests/test_streamable_http_json_response.rs @@ -1,15 +1,20 @@ #![cfg(not(feature = "local"))] use rmcp::{ - ErrorData, ServerHandler, + ErrorData, ServerHandler, ServiceExt, model::{ - CallToolRequestParams, CallToolResponse, CallToolResult, ContentBlock, - ProgressNotificationParam, ServerCapabilities, ServerInfo, + CallToolRequestParams, CallToolResponse, CallToolResult, ContentBlock, CustomRequest, + CustomResult, ProgressNotificationParam, ServerCapabilities, ServerInfo, }, service::RequestContext, - transport::streamable_http_server::{ - StreamableHttpServerConfig, StreamableHttpService, session::local::LocalSessionManager, + transport::{ + StreamableHttpClientTransport, + streamable_http_server::{ + StreamableHttpServerConfig, StreamableHttpService, session::local::LocalSessionManager, + }, }, }; +use serde::Deserialize; +use serde_json::json; use tokio_util::sync::CancellationToken; mod common; @@ -56,6 +61,32 @@ const NEGOTIATED_CALL_WITH_PROGRESS_BODY: &str = r#"{ #[derive(Clone)] struct ProgressServer; +#[derive(Clone)] +struct TypedExtensionServer; + +impl ServerHandler for TypedExtensionServer { + async fn on_custom_request( + &self, + _request: CustomRequest, + _context: RequestContext, + ) -> Result { + Ok(CustomResult::new(json!({ + "resultType": "complete", + "skills": ["http-skill"], + "_meta": {"source": "streamable-http"} + }))) + } +} + +#[derive(Debug, Deserialize)] +#[serde(rename_all = "camelCase")] +struct TypedExtensionResult { + result_type: String, + skills: Vec, + #[serde(rename = "_meta")] + meta: serde_json::Value, +} + impl ServerHandler for ProgressServer { fn get_info(&self) -> ServerInfo { ServerInfo::new(ServerCapabilities::builder().enable_tools().build()) @@ -135,6 +166,43 @@ async fn spawn_progress_server( (client, base_url, ct) } +#[tokio::test] +async fn streamable_http_preserves_typed_extension_results() -> anyhow::Result<()> { + let ct = CancellationToken::new(); + let service: StreamableHttpService = + StreamableHttpService::new( + || Ok(TypedExtensionServer), + Default::default(), + StreamableHttpServerConfig::default() + .with_legacy_session_mode(false) + .with_json_response(true) + .with_sse_keep_alive(None) + .with_cancellation_token(ct.child_token()), + ); + let router = axum::Router::new().nest_service("/mcp", service); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await?; + let url = format!("http://{}/mcp", listener.local_addr()?); + let server_ct = ct.clone(); + tokio::spawn(async move { + let _ = axum::serve(listener, router) + .with_graceful_shutdown(async move { server_ct.cancelled().await }) + .await; + }); + + let client = ().serve(StreamableHttpClientTransport::from_uri(url)).await?; + let result: TypedExtensionResult = client + .send_request_as(rmcp::model::ClientRequest::CustomRequest( + CustomRequest::new("skills/list", None), + )) + .await?; + assert_eq!(result.result_type, "complete"); + assert_eq!(result.skills, ["http-skill"]); + assert_eq!(result.meta["source"], "streamable-http"); + client.cancel().await?; + ct.cancel(); + Ok(()) +} + #[tokio::test] async fn stateless_json_response_returns_application_json() -> anyhow::Result<()> { let ct = CancellationToken::new();