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/29] 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/29] 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/29] 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/29] 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/29] 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/29] 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/29] 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/29] 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/29] 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/29] 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/29] 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 4e744997831819fd7662aa5ac8b27196780d10b9 Mon Sep 17 00:00:00 2001 From: Matthew Zeng Date: Mon, 31 Aug 2026 13:59:37 -0700 Subject: [PATCH 12/29] feat(auth): coordinate OAuth refreshes through credential stores (#1232) * feat(auth): coordinate refreshes through credential stores Add an optional owned guard covering credential reload, refresh, and save. Keep credential-store failures distinct from reauthorization, including reactive refresh after an HTTP 401. Existing stores retain their default uncoordinated behavior. Cover guard ordering, concurrent rotation, storage failures, and client identity checks. * docs(auth): clarify refresh guard coordination --- crates/rmcp/src/transport.rs | 2 +- crates/rmcp/src/transport/auth.rs | 332 +++++++++++++++++- .../common/auth/streamable_http_client.rs | 62 ++++ 3 files changed, 387 insertions(+), 9 deletions(-) diff --git a/crates/rmcp/src/transport.rs b/crates/rmcp/src/transport.rs index 06fde8e51..fdedc9c3a 100644 --- a/crates/rmcp/src/transport.rs +++ b/crates/rmcp/src/transport.rs @@ -100,7 +100,7 @@ pub use auth::JwtSigningAlgorithm; #[cfg(feature = "auth")] pub use auth::{ AuthClient, AuthError, AuthorizationManager, AuthorizationRequest, AuthorizationSession, - AuthorizedHttpClient, ClientCredentialsConfig, CredentialStore, + AuthorizedHttpClient, ClientCredentialsConfig, CredentialRefreshGuard, CredentialStore, EXTENSION_OAUTH_CLIENT_CREDENTIALS, InMemoryCredentialStore, InMemoryStateStore, OAuthHttpClient, OAuthHttpClientError, OAuthHttpClientFuture, OAuthHttpRedirectPolicy, OAuthHttpRequest, ScopeUpgradeConfig, StateStore, StoredAuthorizationState, StoredCredentials, diff --git a/crates/rmcp/src/transport/auth.rs b/crates/rmcp/src/transport/auth.rs index 2eae2b220..6762286d9 100644 --- a/crates/rmcp/src/transport/auth.rs +++ b/crates/rmcp/src/transport/auth.rs @@ -255,11 +255,31 @@ impl StoredCredentials { } } +/// An owned guard held across a credential refresh and its save. +/// +/// Stores can wrap a file lock, an owned mutex guard, or another coordination +/// primitive. Dropping this value releases the guard. +#[must_use = "dropping the guard releases refresh coordination"] +pub struct CredentialRefreshGuard { + _guard: Box, +} + +impl CredentialRefreshGuard { + /// Wrap a guard whose lifetime coordinates access to the stored credentials. + pub fn new(guard: impl Send + 'static) -> Self { + Self { + _guard: Box::new(guard), + } + } +} + /// Trait for storing and retrieving OAuth2 credentials /// /// Implementations of this trait can provide custom storage backends /// for OAuth2 credentials, such as file-based storage, keychain integration, /// or database storage. +/// Return [`AuthError::CredentialStoreError`] for backend or locking failures +/// so they remain distinct from errors requiring reauthorization. #[async_trait] pub trait CredentialStore: Send + Sync { async fn load(&self) -> Result, AuthError>; @@ -267,6 +287,16 @@ pub trait CredentialStore: Send + Sync { async fn save(&self, credentials: StoredCredentials) -> Result<(), AuthError>; async fn clear(&self) -> Result<(), AuthError>; + + /// Optionally coordinate refreshes that share these credentials. + /// + /// The manager acquires this guard before loading credentials and retains it + /// through the token request and save. `load` and `save` must not reacquire + /// the same lock. Writers that bypass the guard are not coordinated with it. + /// The default does not coordinate refreshes. + async fn acquire_refresh_guard(&self) -> Result, AuthError> { + Ok(None) + } } /// In-memory credential store (default implementation) @@ -515,6 +545,9 @@ pub enum AuthError { #[error("OAuth refresh token was rejected: {0}")] TokenRefreshRejected(String), + #[error("OAuth credential store failed: {0}")] + CredentialStoreError(String), + #[error("HTTP error: {0}")] HttpError(#[from] reqwest::Error), @@ -2203,8 +2236,14 @@ impl AuthorizationManager { .as_ref() .ok_or_else(|| AuthError::InternalError("OAuth client not configured".to_string()))?; + let refresh_guard = self.credential_store.acquire_refresh_guard().await?; let stored = self.credential_store.load().await?; let stored_credentials = stored.ok_or(AuthError::AuthorizationRequired)?; + if refresh_guard.is_some() + && stored_credentials.client_id != oauth_client.client_id().as_str() + { + return Err(AuthError::AuthorizationRequired); + } let current_credentials = stored_credentials .token_response .ok_or(AuthError::AuthorizationRequired)?; @@ -2220,6 +2259,7 @@ impl AuthorizationManager { // RFC 8707: the resource indicator is required on token requests, including refreshes .add_extra_param("resource", self.oauth_resource().await); let mut refresh_scopes = stored_credentials.granted_scopes; + let authoritative_scopes = refresh_guard.is_some().then(|| refresh_scopes.clone()); self.add_offline_access_if_supported(&mut refresh_scopes); for scope in refresh_scopes { refresh_request = refresh_request.add_scope(Scope::new(scope)); @@ -2246,9 +2286,10 @@ impl AuthorizationManager { token_result.set_refresh_token(Some(refresh_token_value)); } - let granted_scopes: Vec = match token_result.scopes() { - Some(scopes) => scopes.iter().map(|s| s.to_string()).collect(), - None => self.current_scopes.read().await.clone(), + let granted_scopes: Vec = match (token_result.scopes(), authoritative_scopes) { + (Some(scopes), _) => scopes.iter().map(|s| s.to_string()).collect(), + (None, Some(scopes)) => scopes, + (None, None) => self.current_scopes.read().await.clone(), }; *self.current_scopes.write().await = granted_scopes.clone(); @@ -3904,17 +3945,19 @@ mod tests { sync::{Arc, Mutex as StdMutex}, }; - use oauth2::{AuthType, CsrfToken, HttpResponse, PkceCodeVerifier}; + use oauth2::{AuthType, CsrfToken, HttpResponse, PkceCodeVerifier, TokenResponse}; use reqwest::StatusCode; use rstest::rstest; + use tokio::sync::{Mutex, OwnedMutexGuard, Semaphore}; use url::Url; use super::{ AuthError, AuthorizationCallback, AuthorizationManager, AuthorizationMetadata, - AuthorizationMetadataSource, AuthorizationRequest, AuthorizationSession, CredentialStore, - InMemoryCredentialStore, InMemoryStateStore, OAuthClientConfig, OAuthHttpClient, - OAuthHttpClientError, OAuthHttpClientFuture, OAuthHttpRedirectPolicy, OAuthHttpRequest, - ScopeUpgradeConfig, StateStore, StoredAuthorizationState, is_https_url, + AuthorizationMetadataSource, AuthorizationRequest, AuthorizationSession, + CredentialRefreshGuard, CredentialStore, InMemoryCredentialStore, InMemoryStateStore, + OAuthClientConfig, OAuthHttpClient, OAuthHttpClientError, OAuthHttpClientFuture, + OAuthHttpRedirectPolicy, OAuthHttpRequest, ScopeUpgradeConfig, StateStore, + StoredAuthorizationState, is_https_url, }; use crate::transport::auth::VendorExtraTokenFields; @@ -8342,4 +8385,277 @@ mod tests { "a rotated refresh token from the response should replace the old one" ); } + + #[derive(Clone)] + struct RefreshStore { + credentials: InMemoryCredentialStore, + lock: Arc>, + events: Arc>>, + guard_requested: Arc, + save_started: Arc, + save_gate: Option>, + fail_at: Option<&'static str>, + } + + struct ObservedRefreshGuard { + _lock: OwnedMutexGuard<()>, + events: Arc>>, + } + + impl Drop for ObservedRefreshGuard { + fn drop(&mut self) { + self.events.lock().unwrap().push("release"); + } + } + + #[async_trait::async_trait] + impl CredentialStore for RefreshStore { + async fn load(&self) -> Result, AuthError> { + self.events.lock().unwrap().push("load"); + if self.fail_at == Some("load") { + return Err(AuthError::CredentialStoreError("load failed".into())); + } + self.credentials.load().await + } + + async fn save(&self, credentials: StoredCredentials) -> Result<(), AuthError> { + self.events.lock().unwrap().push("save"); + self.save_started.add_permits(1); + if let Some(gate) = &self.save_gate { + gate.acquire().await.unwrap().forget(); + } + if self.fail_at == Some("save") { + return Err(AuthError::CredentialStoreError("save failed".into())); + } + self.credentials.save(credentials).await?; + self.events.lock().unwrap().push("saved"); + Ok(()) + } + + async fn clear(&self) -> Result<(), AuthError> { + self.credentials.clear().await + } + + async fn acquire_refresh_guard(&self) -> Result, AuthError> { + self.events.lock().unwrap().push("acquire"); + self.guard_requested.add_permits(1); + if self.fail_at == Some("guard") { + return Err(AuthError::CredentialStoreError("guard failed".into())); + } + let lock = self.lock.clone().lock_owned().await; + self.events.lock().unwrap().push("acquired"); + Ok(Some(CredentialRefreshGuard::new(ObservedRefreshGuard { + _lock: lock, + events: self.events.clone(), + }))) + } + } + + fn refresh_store() -> RefreshStore { + let credentials = StoredCredentials::new( + "my-client".into(), + Some(make_token_response_with_refresh("old-token", "old-refresh")), + vec!["read".into()], + Some(AuthorizationManager::now_epoch_secs()), + ); + RefreshStore { + credentials: InMemoryCredentialStore { + credentials: Arc::new(tokio::sync::RwLock::new(Some(credentials))), + }, + lock: Arc::new(Mutex::new(())), + events: Arc::new(StdMutex::new(Vec::new())), + guard_requested: Arc::new(Semaphore::new(0)), + save_started: Arc::new(Semaphore::new(0)), + save_gate: None, + fail_at: None, + } + } + + struct RefreshHttpClient { + recording: RecordingOAuthHttpClient, + events: Arc>>, + } + + impl OAuthHttpClient for RefreshHttpClient { + fn execute(&self, request: OAuthHttpRequest) -> OAuthHttpClientFuture<'_> { + self.events.lock().unwrap().push("provider"); + self.recording.execute(request) + } + } + + fn refresh_http_client(store: &RefreshStore) -> Arc { + Arc::new(RefreshHttpClient { + recording: RecordingOAuthHttpClient::with_responses( + [ + ("new-token", "new-refresh"), + ("latest-token", "latest-refresh"), + ] + .into_iter() + .map(|(access, refresh)| { + http_response( + 200, + serde_json::json!({ + "access_token": access, "token_type": "Bearer", + "expires_in": 3600, "refresh_token": refresh + }), + ) + }) + .collect(), + ), + events: store.events.clone(), + }) + } + + async fn refresh_manager( + store: RefreshStore, + http_client: Arc, + ) -> AuthorizationManager { + let mut manager = AuthorizationManager::new_with_oauth_http_client( + "https://mcp.example.com/mcp", + http_client, + ) + .await + .unwrap(); + manager.set_metadata(AuthorizationMetadata { + authorization_endpoint: "https://auth.example.com/authorize".into(), + token_endpoint: "https://auth.example.com/token".into(), + ..Default::default() + }); + manager.configure_client(test_client_config()).unwrap(); + manager.set_credential_store(store); + *manager.current_scopes.write().await = vec!["cached".into()]; + manager + } + + async fn wait_for_permits(semaphore: &Semaphore, count: u32) { + tokio::time::timeout( + std::time::Duration::from_secs(5), + semaphore.acquire_many(count), + ) + .await + .unwrap() + .unwrap() + .forget(); + } + + #[tokio::test] + async fn refresh_guard_spans_load_exchange_and_completed_save() { + let store = refresh_store(); + let manager = refresh_manager(store.clone(), refresh_http_client(&store)).await; + + manager.refresh_token().await.unwrap(); + + assert_eq!( + *store.events.lock().unwrap(), + [ + "acquire", "acquired", "load", "provider", "save", "saved", "release" + ] + ); + let saved = store.credentials.load().await.unwrap().unwrap(); + assert_eq!( + saved + .token_response + .unwrap() + .refresh_token() + .unwrap() + .secret(), + "new-refresh" + ); + assert_eq!(saved.granted_scopes, ["read"]); + assert!(store.lock.try_lock().is_ok()); + } + + #[tokio::test] + async fn concurrent_refreshes_wait_for_save_and_use_the_latest_token() { + let mut store = refresh_store(); + let save_gate = Arc::new(Semaphore::new(0)); + store.save_gate = Some(save_gate.clone()); + let http_client = refresh_http_client(&store); + let first_manager = refresh_manager(store.clone(), http_client.clone()).await; + let second_manager = refresh_manager(store.clone(), http_client.clone()).await; + + let first = tokio::spawn(async move { first_manager.refresh_token().await }); + wait_for_permits(&store.save_started, 1).await; + let second = tokio::spawn(async move { second_manager.refresh_token().await }); + wait_for_permits(&store.guard_requested, 2).await; + assert_eq!(http_client.recording.requests().len(), 1); + assert!(store.lock.try_lock().is_err()); + + save_gate.add_permits(2); + let (first, second) = tokio::join!(first, second); + assert_eq!(first.unwrap().unwrap().access_token().secret(), "new-token"); + assert_eq!( + second.unwrap().unwrap().access_token().secret(), + "latest-token" + ); + let refresh_tokens: Vec = http_client + .recording + .requests() + .iter() + .map(|request| { + url::form_urlencoded::parse(&request.body) + .find_map(|(key, value)| (key == "refresh_token").then(|| value.into_owned())) + .unwrap() + }) + .collect(); + assert_eq!(refresh_tokens, ["old-refresh", "new-refresh"]); + let saved = store.credentials.load().await.unwrap().unwrap(); + assert_eq!( + saved + .token_response + .unwrap() + .refresh_token() + .unwrap() + .secret(), + "latest-refresh" + ); + assert!(store.lock.try_lock().is_ok()); + } + + #[tokio::test] + async fn guarded_refresh_rejects_credentials_for_another_client() { + let store = refresh_store(); + let mut credentials = store.credentials.load().await.unwrap().unwrap(); + credentials.client_id = "other-client".into(); + store.credentials.save(credentials).await.unwrap(); + let http_client = refresh_http_client(&store); + let manager = refresh_manager(store.clone(), http_client.clone()).await; + + assert!(matches!( + manager.refresh_token().await, + Err(AuthError::AuthorizationRequired) + )); + assert!(http_client.recording.requests().is_empty()); + assert!(store.lock.try_lock().is_ok()); + } + + #[rstest] + #[case("guard", 0)] + #[case("load", 0)] + #[case("save", 1)] + #[tokio::test] + async fn refresh_store_failures_release_the_guard( + #[case] phase: &'static str, + #[case] provider_requests: usize, + ) { + let mut store = refresh_store(); + store.fail_at = Some(phase); + let http_client = refresh_http_client(&store); + let manager = refresh_manager(store.clone(), http_client.clone()).await; + + assert!(matches!(manager.refresh_token().await, + Err(AuthError::CredentialStoreError(message)) if message == format!("{phase} failed"))); + assert_eq!(http_client.recording.requests().len(), provider_requests); + let saved = store.credentials.load().await.unwrap().unwrap(); + assert_eq!( + saved + .token_response + .unwrap() + .refresh_token() + .unwrap() + .secret(), + "old-refresh" + ); + assert!(store.lock.try_lock().is_ok()); + } } diff --git a/crates/rmcp/src/transport/common/auth/streamable_http_client.rs b/crates/rmcp/src/transport/common/auth/streamable_http_client.rs index 2053920d2..97432f90e 100644 --- a/crates/rmcp/src/transport/common/auth/streamable_http_client.rs +++ b/crates/rmcp/src/transport/common/auth/streamable_http_client.rs @@ -54,6 +54,7 @@ where match refreshed { Ok(fresh_token) if fresh_token != sent_token => call(Some(fresh_token)).await, Ok(_) => Err(StreamableHttpError::AuthRequired(challenge)), + Err(error @ AuthError::CredentialStoreError(_)) => Err(error.into()), Err(error) => { debug!("token refresh after server rejection failed: {error}"); Err(StreamableHttpError::AuthRequired(challenge)) @@ -208,3 +209,64 @@ where .await } } + +#[cfg(all(test, feature = "transport-streamable-http-client-reqwest"))] +mod tests { + use super::*; + use crate::transport::{ + auth::{ + AuthorizationManager, AuthorizationMetadata, CredentialRefreshGuard, CredentialStore, + StoredCredentials, + }, + streamable_http_client::AuthRequiredError, + }; + + struct UnavailableStore; + + #[async_trait::async_trait] + impl CredentialStore for UnavailableStore { + async fn load(&self) -> Result, AuthError> { + unreachable!("guard failure must stop the credential load") + } + + async fn save(&self, _: StoredCredentials) -> Result<(), AuthError> { + unreachable!("guard failure must stop the credential save") + } + + async fn clear(&self) -> Result<(), AuthError> { + unreachable!("refresh must not clear credentials") + } + + async fn acquire_refresh_guard(&self) -> Result, AuthError> { + Err(AuthError::CredentialStoreError("guard unavailable".into())) + } + } + + #[tokio::test] + async fn reactive_refresh_preserves_credential_store_failure() { + let mut manager = AuthorizationManager::new("https://mcp.example.com/mcp") + .await + .unwrap(); + manager.set_metadata(AuthorizationMetadata { + authorization_endpoint: "https://auth.example.com/authorize".into(), + token_endpoint: "https://auth.example.com/token".into(), + ..Default::default() + }); + manager.configure_client_id("client").unwrap(); + manager.set_credential_store(UnavailableStore); + let client = AuthClient::new(reqwest::Client::new(), manager); + + let error = client + .call_reacting_to_challenges(Some("old-token".into()), |_| async { + Err::<(), _>(StreamableHttpError::AuthRequired(AuthRequiredError::new( + "Bearer".into(), + ))) + }) + .await + .unwrap_err(); + + assert!(matches!(error, + StreamableHttpError::Auth(AuthError::CredentialStoreError(message)) + if message == "guard unavailable")); + } +} From ad9832ec212baf526e1a69d73ee04cd8305ae331 Mon Sep 17 00:00:00 2001 From: Dale Seo <5466341+DaleSeo@users.noreply.github.com> Date: Tue, 1 Sep 2026 07:38:57 +0900 Subject: [PATCH 13/29] fix: keep initialize on legacy protocol versions (#1228) Claude-Session: https://claude.ai/code/session_01EPcBfiUWJJueJn5cve1sjp --- crates/rmcp/src/handler/server.rs | 9 +- crates/rmcp/src/service.rs | 9 +- crates/rmcp/src/service/server.rs | 94 ++++++++++++---- .../transport/streamable_http_server/tower.rs | 30 +++-- crates/rmcp/tests/test_handler_cache_hints.rs | 39 +++++-- .../test_protocol_version_negotiation.rs | 105 ++++++++++++++---- .../tests/test_resource_not_found_version.rs | 37 +++--- crates/rmcp/tests/test_result_type_version.rs | 33 ++++-- .../test_sep_2260_request_association.rs | 42 +++++-- .../tests/test_stateless_protocol_version.rs | 28 ++--- .../test_streamable_http_protocol_version.rs | 63 +++++++++++ 11 files changed, 369 insertions(+), 120 deletions(-) diff --git a/crates/rmcp/src/handler/server.rs b/crates/rmcp/src/handler/server.rs index 6dc7883ed..b5f70be3f 100644 --- a/crates/rmcp/src/handler/server.rs +++ b/crates/rmcp/src/handler/server.rs @@ -322,12 +322,15 @@ macro_rules! server_handler_methods { ) -> impl Future> + MaybeSendFuture + '_ { context.peer.set_peer_info(request.clone()); let mut info = self.get_info(); - info.protocol_version = negotiate_protocol_version( + let negotiated = negotiate_protocol_version( &request.protocol_version, - info.protocol_version, + std::mem::take(&mut info.protocol_version), &self.supported_protocol_versions(), ); - std::future::ready(Ok(info)) + std::future::ready(negotiated.map(|version| { + info.protocol_version = version; + info + })) } /// Return the protocol versions supported by this server. /// diff --git a/crates/rmcp/src/service.rs b/crates/rmcp/src/service.rs index 20fd2e981..b6cc5e538 100644 --- a/crates/rmcp/src/service.rs +++ b/crates/rmcp/src/service.rs @@ -193,12 +193,17 @@ pub enum PeerRequestAssociation { Unknown { has_pending_outbound_request: bool }, } +/// Whether `version` predates `2026-07-28`, the revision that replaced the +/// `initialize` handshake with per-request metadata. +pub(crate) fn is_legacy_version(version: &ProtocolVersion) -> bool { + version.as_str() < ProtocolVersion::V_2026_07_28.as_str() +} + pub(crate) fn uses_legacy_lifecycle( protocol_version: Option<&ProtocolVersion>, uses_discover_lifecycle: bool, ) -> bool { - !uses_discover_lifecycle - && protocol_version.is_none_or(|version| version < &ProtocolVersion::V_2026_07_28) + !uses_discover_lifecycle && protocol_version.is_none_or(is_legacy_version) } pub(crate) fn peer_request_association( diff --git a/crates/rmcp/src/service/server.rs b/crates/rmcp/src/service/server.rs index c2f57a2a6..70a148641 100644 --- a/crates/rmcp/src/service/server.rs +++ b/crates/rmcp/src/service/server.rs @@ -460,27 +460,60 @@ where } } -/// Echoes the client-requested version if the server supports it; otherwise -/// returns `server_fallback`. +/// Echoes the client-requested version if the server can serve it over the +/// `initialize` handshake; otherwise returns a legacy version the server does +/// support. /// /// `server_supported` comes from [`Service::supported_protocol_versions`], so a /// server that narrows that list is never made to answer `initialize` with a -/// version it cannot serve. +/// version it cannot serve. `2026-07-28` replaced the handshake with +/// per-request metadata, so a client naming that revision or later is answered +/// with the server's newest legacy version instead. +/// +/// # Errors +/// +/// Returns [`ErrorCode::UNSUPPORTED_PROTOCOL_VERSION`] when the server supports +/// no version that still has an `initialize` handshake. +/// +/// [`ErrorCode::UNSUPPORTED_PROTOCOL_VERSION`]: crate::model::ErrorCode::UNSUPPORTED_PROTOCOL_VERSION pub(crate) fn negotiate_protocol_version( client_requested: &ProtocolVersion, server_fallback: ProtocolVersion, server_supported: &[ProtocolVersion], -) -> ProtocolVersion { - if server_supported.contains(client_requested) { - client_requested.clone() +) -> Result { + if is_legacy_version(client_requested) && server_supported.contains(client_requested) { + return Ok(client_requested.clone()); + } + let legacy_fallback = if is_legacy_version(&server_fallback) { + Some(server_fallback) } else { + newest_legacy_version(server_supported) + }; + let Some(legacy_fallback) = legacy_fallback else { tracing::warn!( client_requested = %client_requested, - server_fallback = %server_fallback, - "client requested unsupported protocol version; falling back to server default" + "server supports no protocol version with an initialize handshake; rejecting" ); - server_fallback - } + return Err(ErrorData::unsupported_protocol_version( + client_requested.clone(), + server_supported, + )); + }; + tracing::warn!( + client_requested = %client_requested, + server_fallback = %legacy_fallback, + "client requested a protocol version unavailable over initialize; falling back to server default" + ); + Ok(legacy_fallback) +} + +/// The newest of `versions` that still has an `initialize` handshake. +fn newest_legacy_version(versions: &[ProtocolVersion]) -> Option { + versions + .iter() + .filter(|version| is_legacy_version(version)) + .max_by(|left, right| left.as_str().cmp(right.as_str())) + .cloned() } fn missing_request_metadata_error(missing: &[&str]) -> ErrorData { @@ -493,6 +526,27 @@ fn missing_request_metadata_error(missing: &[&str]) -> ErrorData { ) } +/// Sends `error` as the response to the `initialize` request and reports it as +/// the reason the handshake failed. +async fn report_initialize_failure( + transport: &mut T, + error: ErrorData, + id: RequestId, +) -> ServerInitializeError +where + T: Transport + 'static, +{ + match transport + .send(ServerJsonRpcMessage::error(error.clone(), Some(id))) + .await + { + Ok(()) => ServerInitializeError::InitializeFailed(error), + Err(send_error) => { + ServerInitializeError::transport::(send_error, "sending error response") + } + } +} + async fn serve_server_with_ct_inner( service: S, transport: T, @@ -589,27 +643,21 @@ where peer: peer.clone(), }; // Send initialize response - let init_response = service.handle_request(request, context).await; - let mut init_response = match init_response { + let mut init_response = match service.handle_request(request, context).await { Ok(ServerResult::InitializeResult(init_response)) => init_response, Ok(result) => { return Err(ServerInitializeError::UnexpectedInitializeResponse(result)); } - Err(e) => { - transport - .send(ServerJsonRpcMessage::error(e.clone(), Some(id))) - .await - .map_err(|error| { - ServerInitializeError::transport::(error, "sending error response") - })?; - return Err(ServerInitializeError::InitializeFailed(e)); - } + Err(e) => return Err(report_initialize_failure(&mut transport, e, id).await), }; - init_response.protocol_version = negotiate_protocol_version( + init_response.protocol_version = match negotiate_protocol_version( &requested_protocol_version, init_response.protocol_version, &service.supported_protocol_versions(), - ); + ) { + Ok(version) => version, + Err(e) => return Err(report_initialize_failure(&mut transport, e, id).await), + }; // Update peer_info so context.protocol_version() reflects the negotiated // version in all subsequent request handlers. negotiated_peer_info.protocol_version = init_response.protocol_version.clone(); diff --git a/crates/rmcp/src/transport/streamable_http_server/tower.rs b/crates/rmcp/src/transport/streamable_http_server/tower.rs index f1fef585f..e15dfd434 100644 --- a/crates/rmcp/src/transport/streamable_http_server/tower.rs +++ b/crates/rmcp/src/transport/streamable_http_server/tower.rs @@ -322,7 +322,7 @@ impl> Service for NegotiatingStatelessHttpSer &requested, result.protocol_version.clone(), &self.0.supported_protocol_versions(), - ); + )?; if let Some(peer_info) = peer.peer_info() { let mut peer_info = (*peer_info).clone(); peer_info.protocol_version = result.protocol_version.clone(); @@ -375,22 +375,30 @@ fn is_legacy_request( validate_request_protocol_version_meta(headers, message)?; } + // An `initialize` request selects legacy semantics whatever version it names: + // the handshake exists only in the revisions before 2026-07-28, so the + // version in its params never routes it to the stateless path. The + // handshake itself answers with a legacy version the server supports. + if matches!( + message, + Some(ClientJsonRpcMessage::Request(req)) + if matches!(&req.request, ClientRequest::InitializeRequest(_)) + ) { + return Ok(true); + } + let uses_discover_lifecycle = matches!( message, Some(ClientJsonRpcMessage::Request(req)) - if !matches!(&req.request, ClientRequest::InitializeRequest(_)) - && req - .request - .get_meta() - .missing_required_keys(&ProtocolVersion::V_2026_07_28) - .is_empty() + if req + .request + .get_meta() + .missing_required_keys(&ProtocolVersion::V_2026_07_28) + .is_empty() ); let from_body = match message { - Some(ClientJsonRpcMessage::Request(req)) => match &req.request { - ClientRequest::InitializeRequest(init) => Some(init.params.protocol_version.clone()), - _ => req.request.get_meta().protocol_version(), - }, + Some(ClientJsonRpcMessage::Request(req)) => req.request.get_meta().protocol_version(), _ => None, }; let version = from_body diff --git a/crates/rmcp/tests/test_handler_cache_hints.rs b/crates/rmcp/tests/test_handler_cache_hints.rs index d6b626d2a..7f44e8280 100644 --- a/crates/rmcp/tests/test_handler_cache_hints.rs +++ b/crates/rmcp/tests/test_handler_cache_hints.rs @@ -2,10 +2,14 @@ #![cfg(feature = "client")] use rmcp::{ - ClientHandler, ServerHandler, ServiceExt, + ClientHandler, RoleClient, RoleServer, ServerHandler, handler::server::router::{prompt::PromptRouter, tool::ToolRouter}, - model::{CacheScope, ClientInfo, ListPromptsResult, ListToolsResult, ProtocolVersion}, - prompt_handler, tool_handler, + model::{ + CacheScope, ClientInfo, ListPromptsResult, ListToolsResult, ProtocolVersion, ServerInfo, + }, + prompt_handler, + service::serve_directly, + tool_handler, }; #[derive(Debug, Clone)] @@ -40,22 +44,33 @@ impl ClientHandler for VersionedClient { } } +/// Wires the pair up directly on `protocol_version`. `2026-07-28` removed the +/// `initialize` handshake, so a peer on that revision is reached the way the +/// discover lifecycle leaves one: with the version already agreed. async fn list_results(protocol_version: ProtocolVersion) -> (ListToolsResult, ListPromptsResult) { let (server_transport, client_transport) = tokio::io::duplex(4096); + let client_handler = VersionedClient { + protocol_version: protocol_version.clone(), + }; + let mut server_peer_info = ServerInfo::default(); + server_peer_info.protocol_version = protocol_version; + + let server = serve_directly::( + CacheHintServer::new(), + server_transport, + Some(client_handler.get_info()), + ); let server_handle = tokio::spawn(async move { - CacheHintServer::new() - .serve(server_transport) - .await? - .waiting() - .await?; + server.waiting().await?; anyhow::Ok(()) }); - let client = VersionedClient { protocol_version } - .serve(client_transport) - .await - .expect("client should connect"); + let client = serve_directly::( + client_handler, + client_transport, + Some(server_peer_info.into()), + ); let tools = client .list_tools(None) .await diff --git a/crates/rmcp/tests/test_protocol_version_negotiation.rs b/crates/rmcp/tests/test_protocol_version_negotiation.rs index e91ecf97a..717a63f18 100644 --- a/crates/rmcp/tests/test_protocol_version_negotiation.rs +++ b/crates/rmcp/tests/test_protocol_version_negotiation.rs @@ -1,6 +1,7 @@ //! Tests for protocol version negotiation in the default ServerHandler::initialize impl. //! -//! Known versions are echoed back; unknown versions fall back to LATEST. +//! Handshake versions are echoed back; every other version falls back to one +//! the server can serve over `initialize`. #![cfg(not(feature = "local"))] #![cfg(feature = "client")] @@ -8,8 +9,11 @@ use std::borrow::Cow; use rmcp::{ ClientHandler, ErrorData, RoleServer, ServerHandler, ServiceExt, - model::{ClientInfo, InitializeRequestParams, InitializeResult, ProtocolVersion, ServerInfo}, - service::RequestContext, + model::{ + ClientInfo, ErrorCode, InitializeRequestParams, InitializeResult, ProtocolVersion, + ServerInfo, + }, + service::{ClientInitializeError, RequestContext}, }; #[derive(Debug, Clone, Default)] @@ -21,9 +25,10 @@ impl ServerHandler for EchoServer { } } -/// Every known version except `2026-07-28`, standing in for a server that has -/// not implemented that revision. -const NARROWED_VERSIONS: &[ProtocolVersion] = &[ +/// Every known version whose lifecycle still runs the `initialize` handshake. +/// `2026-07-28` replaced the handshake with per-request metadata, so this is +/// also the list a server that has not implemented that revision supports. +const HANDSHAKE_VERSIONS: &[ProtocolVersion] = &[ ProtocolVersion::V_2024_11_05, ProtocolVersion::V_2025_03_26, ProtocolVersion::V_2025_06_18, @@ -39,7 +44,25 @@ impl ServerHandler for NarrowedServer { } fn supported_protocol_versions(&self) -> Cow<'static, [ProtocolVersion]> { - Cow::Borrowed(NARROWED_VERSIONS) + Cow::Borrowed(HANDSHAKE_VERSIONS) + } +} + +/// Supports only revisions that have no `initialize` handshake at all. +#[derive(Debug, Clone, Default)] +struct ModernOnlyServer; + +const MODERN_ONLY_VERSIONS: &[ProtocolVersion] = &[ProtocolVersion::V_2026_07_28]; + +impl ServerHandler for ModernOnlyServer { + fn get_info(&self) -> ServerInfo { + let mut info = ServerInfo::default(); + info.protocol_version = ProtocolVersion::V_2026_07_28; + info + } + + fn supported_protocol_versions(&self) -> Cow<'static, [ProtocolVersion]> { + Cow::Borrowed(MODERN_ONLY_VERSIONS) } } @@ -55,7 +78,7 @@ impl ServerHandler for NarrowedOverridingServer { } fn supported_protocol_versions(&self) -> Cow<'static, [ProtocolVersion]> { - Cow::Borrowed(NARROWED_VERSIONS) + Cow::Borrowed(HANDSHAKE_VERSIONS) } async fn initialize( @@ -88,23 +111,28 @@ async fn negotiated_version_with( server: S, client_version: ProtocolVersion, ) -> ProtocolVersion { + negotiate_with(server, client_version) + .await + .expect("client should connect") +} + +async fn negotiate_with( + server: S, + client_version: ProtocolVersion, +) -> Result { let (server_transport, client_transport) = tokio::io::duplex(4096); tokio::spawn(async move { - let _ = server - .serve(server_transport) - .await - .expect("server should start") - .waiting() - .await; + if let Ok(running) = server.serve(server_transport).await { + let _ = running.waiting().await; + } }); let client = VersionedClient { protocol_version: client_version, } .serve(client_transport) - .await - .expect("client should connect"); + .await?; let version = client .peer_info() @@ -113,20 +141,55 @@ async fn negotiated_version_with( .clone(); client.cancel().await.expect("client should cancel"); - version + Ok(version) } #[tokio::test] -async fn known_version_echoed_back() { - for version in ProtocolVersion::KNOWN_VERSIONS { +async fn handshake_version_echoed_back() { + for version in HANDSHAKE_VERSIONS { let negotiated = negotiated_version(version.clone()).await; assert_eq!( negotiated, *version, - "known version {version} should be echoed back" + "handshake version {version} should be echoed back" ); } } +/// `initialize` disappeared in `2026-07-28`, so agreeing to it here would leave +/// the peers speaking a revision that has no handshake at all. +#[tokio::test] +async fn handshake_never_agrees_to_a_version_that_dropped_it() { + let negotiated = negotiated_version(ProtocolVersion::V_2026_07_28).await; + assert_eq!( + negotiated, + ProtocolVersion::LATEST, + "a version that dropped the handshake should fall back to the server's own" + ); +} + +#[tokio::test] +async fn modern_only_server_rejects_the_handshake() { + let error = negotiate_with(ModernOnlyServer, ProtocolVersion::V_2026_07_28) + .await + .expect_err("a server with no handshake version cannot answer initialize"); + let ClientInitializeError::JsonRpcError(error) = error else { + panic!("expected a JSON-RPC error, got {error:?}"); + }; + assert_eq!( + error.code, + ErrorCode::UNSUPPORTED_PROTOCOL_VERSION, + "a server with no handshake version should reject initialize" + ); + assert_eq!( + error.data, + Some(serde_json::json!({ + "requested": "2026-07-28", + "supported": ["2026-07-28"], + })), + "the rejection should name the versions the server does support" + ); +} + #[tokio::test] async fn unknown_version_falls_back_to_latest() { let unknown: ProtocolVersion = serde_json::from_str(r#""1999-01-01""#).unwrap(); @@ -140,7 +203,7 @@ async fn unknown_version_falls_back_to_latest() { #[tokio::test] async fn narrowed_server_still_echoes_versions_it_supports() { - for version in NARROWED_VERSIONS { + for version in HANDSHAKE_VERSIONS { let negotiated = negotiated_version_with(NarrowedServer, version.clone()).await; assert_eq!( negotiated, *version, diff --git a/crates/rmcp/tests/test_resource_not_found_version.rs b/crates/rmcp/tests/test_resource_not_found_version.rs index 44eb3631f..1d86ca8c3 100644 --- a/crates/rmcp/tests/test_resource_not_found_version.rs +++ b/crates/rmcp/tests/test_resource_not_found_version.rs @@ -6,12 +6,12 @@ #![cfg(feature = "client")] use rmcp::{ - ClientHandler, RoleServer, ServerHandler, ServiceError, ServiceExt, + ClientHandler, RoleClient, RoleServer, ServerHandler, ServiceError, model::{ ClientInfo, ErrorCode, ErrorData, ProtocolVersion, ReadResourceRequestParams, - ReadResourceResponse, + ReadResourceResponse, ServerInfo, }, - service::RequestContext, + service::{RequestContext, serve_directly}, }; #[derive(Debug, Clone, Default)] @@ -40,24 +40,33 @@ impl ClientHandler for VersionedClient { } } +/// Wires the pair up directly on `client_version`. `2026-07-28` removed the +/// `initialize` handshake, so a peer on that revision is reached the way the +/// discover lifecycle leaves one: with the version already agreed. async fn not_found_code(client_version: ProtocolVersion) -> ErrorCode { let (server_transport, client_transport) = tokio::io::duplex(4096); + let client_handler = VersionedClient { + protocol_version: client_version.clone(), + }; + let mut server_peer_info = ServerInfo::default(); + server_peer_info.protocol_version = client_version; + + let server = serve_directly::( + ResourceServer, + server_transport, + Some(client_handler.get_info()), + ); let server_handle = tokio::spawn(async move { - ResourceServer - .serve(server_transport) - .await? - .waiting() - .await?; + server.waiting().await?; anyhow::Ok(()) }); - let client = VersionedClient { - protocol_version: client_version, - } - .serve(client_transport) - .await - .expect("client should connect"); + let client = serve_directly::( + client_handler, + client_transport, + Some(server_peer_info.into()), + ); let error = client .read_resource(ReadResourceRequestParams::new("missing://resource")) diff --git a/crates/rmcp/tests/test_result_type_version.rs b/crates/rmcp/tests/test_result_type_version.rs index 849a6610e..ab779d283 100644 --- a/crates/rmcp/tests/test_result_type_version.rs +++ b/crates/rmcp/tests/test_result_type_version.rs @@ -6,12 +6,12 @@ #![cfg(feature = "client")] use rmcp::{ - ClientHandler, RoleServer, ServerHandler, ServiceExt, + ClientHandler, RoleClient, RoleServer, ServerHandler, model::{ CallToolRequestParams, CallToolResponse, CallToolResult, ClientInfo, ContentBlock, - ErrorData, ProtocolVersion, ResultType, + ErrorData, ProtocolVersion, ResultType, ServerInfo, }, - service::RequestContext, + service::{RequestContext, serve_directly}, }; #[derive(Debug, Clone, Default)] @@ -40,20 +40,33 @@ impl ClientHandler for VersionedClient { } } +/// Wires the pair up directly on `client_version`. `2026-07-28` removed the +/// `initialize` handshake, so a peer on that revision is reached the way the +/// discover lifecycle leaves one: with the version already agreed. async fn call_tool_result_type(client_version: ProtocolVersion) -> Option { let (server_transport, client_transport) = tokio::io::duplex(4096); + let client_handler = VersionedClient { + protocol_version: client_version.clone(), + }; + let mut server_peer_info = ServerInfo::default(); + server_peer_info.protocol_version = client_version; + + let server = serve_directly::( + ToolServer, + server_transport, + Some(client_handler.get_info()), + ); let server_handle = tokio::spawn(async move { - ToolServer.serve(server_transport).await?.waiting().await?; + server.waiting().await?; anyhow::Ok(()) }); - let client = VersionedClient { - protocol_version: client_version, - } - .serve(client_transport) - .await - .expect("client should connect"); + let client = serve_directly::( + client_handler, + client_transport, + Some(server_peer_info.into()), + ); let result = client .call_tool(CallToolRequestParams::new("echo")) diff --git a/crates/rmcp/tests/test_sep_2260_request_association.rs b/crates/rmcp/tests/test_sep_2260_request_association.rs index d4e20e1e9..b0e7af905 100644 --- a/crates/rmcp/tests/test_sep_2260_request_association.rs +++ b/crates/rmcp/tests/test_sep_2260_request_association.rs @@ -10,7 +10,7 @@ use rmcp::{ CreateMessageRequest, CreateMessageRequestParams, CreateMessageResult, ProtocolVersion, SamplingMessage, ServerCapabilities, ServerInfo, ServerRequest, }, - service::RequestContext, + service::{RequestContext, RunningService, serve_directly}, }; use serde_json::{Value, json}; use tokio::{ @@ -95,21 +95,44 @@ impl ClientHandler for SamplingClient { } } +/// Connects the pair on `2026-07-28`. That revision dropped the `initialize` +/// handshake, so the version is agreed up front the way a discover-lifecycle +/// startup leaves it. +fn serve_modern_pair( + server: SamplingServer, +) -> ( + RunningService, + RunningService, +) { + let (server_transport, client_transport) = tokio::io::duplex(4096); + let mut server_peer_info = server.get_info(); + server_peer_info.protocol_version = ProtocolVersion::V_2026_07_28; + + let running_server = serve_directly::( + server, + server_transport, + Some(SamplingClient.get_info()), + ); + let client = serve_directly::( + SamplingClient, + client_transport, + Some(server_peer_info.into()), + ); + (running_server, client) +} + #[tokio::test] async fn nested_sampling_allowed_standalone_rejected() -> anyhow::Result<()> { - let (server_transport, client_transport) = tokio::io::duplex(4096); let (tx, rx) = oneshot::channel(); let server = SamplingServer { outside: Arc::new(Mutex::new(Some(tx))), }; + let (running_server, client) = serve_modern_pair(server); let server_handle = tokio::spawn(async move { - let running = server.serve(server_transport).await?; - running.waiting().await?; + running_server.waiting().await?; anyhow::Ok(()) }); - let client = SamplingClient.serve(client_transport).await?; - let result = client .peer() .call_tool(CallToolRequestParams::new("sample")) @@ -129,19 +152,16 @@ async fn nested_sampling_allowed_standalone_rejected() -> anyhow::Result<()> { #[tokio::test] async fn generic_send_request_bypass_rejected() -> anyhow::Result<()> { - let (server_transport, client_transport) = tokio::io::duplex(4096); let (tx, rx) = oneshot::channel(); let server = SamplingServer { outside: Arc::new(Mutex::new(Some(tx))), }; + let (running_server, client) = serve_modern_pair(server); let server_handle = tokio::spawn(async move { - let running = server.serve(server_transport).await?; - running.waiting().await?; + running_server.waiting().await?; anyhow::Ok(()) }); - let client = SamplingClient.serve(client_transport).await?; - let result = client .peer() .call_tool(CallToolRequestParams::new("sample_generic")) diff --git a/crates/rmcp/tests/test_stateless_protocol_version.rs b/crates/rmcp/tests/test_stateless_protocol_version.rs index 02222ec2c..dbfcaf27c 100644 --- a/crates/rmcp/tests/test_stateless_protocol_version.rs +++ b/crates/rmcp/tests/test_stateless_protocol_version.rs @@ -1,7 +1,8 @@ //! Tests for protocol version negotiation in stateless HTTP mode. //! -//! Supported versions are echoed back; unknown versions, and versions outside -//! the server's `supported_protocol_versions`, fall back to the handler's own +//! Supported handshake versions are echoed back; unknown versions, versions +//! outside the server's `supported_protocol_versions`, and versions that no +//! longer have an `initialize` handshake fall back to the handler's own //! version. #![cfg(not(feature = "local"))] @@ -36,9 +37,10 @@ impl ServerHandler for OverridingInitialize { } } -/// Every known version except `2026-07-28`, standing in for a server that has -/// not implemented that revision. -const NARROWED_VERSIONS: &[ProtocolVersion] = &[ +/// Every known version whose lifecycle still runs the `initialize` handshake. +/// `2026-07-28` replaced the handshake with per-request metadata, so this is +/// also the list a server that has not implemented that revision supports. +const HANDSHAKE_VERSIONS: &[ProtocolVersion] = &[ ProtocolVersion::V_2024_11_05, ProtocolVersion::V_2025_03_26, ProtocolVersion::V_2025_06_18, @@ -56,7 +58,7 @@ impl ServerHandler for NarrowedOverridingInitialize { } fn supported_protocol_versions(&self) -> Cow<'static, [ProtocolVersion]> { - Cow::Borrowed(NARROWED_VERSIONS) + Cow::Borrowed(HANDSHAKE_VERSIONS) } async fn initialize( @@ -148,15 +150,15 @@ async fn post_init(client: &reqwest::Client, url: &str, body_version: &str) -> s } #[tokio::test] -async fn stateless_json_init_echoes_known_versions_when_handler_overrides_initialize() { +async fn stateless_json_init_echoes_handshake_versions_when_handler_overrides_initialize() { let (client, url, ct) = spawn_server(stateless_json_config()).await; - for version in ProtocolVersion::KNOWN_VERSIONS { + for version in HANDSHAKE_VERSIONS { let resp = post_init(&client, &url, version.as_str()).await; assert_eq!( resp["result"]["protocolVersion"], version.as_str(), - "known version {version} should be echoed back" + "handshake version {version} should be echoed back" ); } @@ -164,15 +166,15 @@ async fn stateless_json_init_echoes_known_versions_when_handler_overrides_initia } #[tokio::test] -async fn stateless_sse_init_echoes_known_versions_when_handler_overrides_initialize() { +async fn stateless_sse_init_echoes_handshake_versions_when_handler_overrides_initialize() { let (client, url, ct) = spawn_server(stateless_sse_config()).await; - for version in ProtocolVersion::KNOWN_VERSIONS { + for version in HANDSHAKE_VERSIONS { let resp = post_init(&client, &url, version.as_str()).await; assert_eq!( resp["result"]["protocolVersion"], version.as_str(), - "known version {version} should be echoed back" + "handshake version {version} should be echoed back" ); } @@ -198,7 +200,7 @@ async fn stateless_json_init_echoes_versions_the_server_narrowed_to() { let (client, url, ct) = spawn_server_of::(stateless_json_config()).await; - for version in NARROWED_VERSIONS { + for version in HANDSHAKE_VERSIONS { let resp = post_init(&client, &url, version.as_str()).await; assert_eq!( resp["result"]["protocolVersion"], diff --git a/crates/rmcp/tests/test_streamable_http_protocol_version.rs b/crates/rmcp/tests/test_streamable_http_protocol_version.rs index fcbeb41a0..373549aa9 100644 --- a/crates/rmcp/tests/test_streamable_http_protocol_version.rs +++ b/crates/rmcp/tests/test_streamable_http_protocol_version.rs @@ -87,6 +87,17 @@ async fn post_init( req.send().await.expect("send initialize request") } +/// First JSON-RPC message of an SSE response, skipping the empty priming event. +async fn sse_payload(response: reqwest::Response) -> anyhow::Result { + let body = response.text().await?; + let data = body + .lines() + .filter_map(|line| line.strip_prefix("data: ")) + .find(|data| !data.is_empty()) + .expect("response must carry a JSON-RPC payload"); + Ok(serde_json::from_str(data)?) +} + async fn post_non_initialize(client: &reqwest::Client, url: &str) -> reqwest::Response { client .post(url) @@ -186,6 +197,27 @@ async fn stateless_init_accepts_when_header_absent() -> anyhow::Result<()> { Ok(()) } +#[tokio::test] +async fn stateless_init_naming_a_post_handshake_version_negotiates_down() -> anyhow::Result<()> { + let (client, url, ct) = spawn_server(stateless_json_config()).await; + + let response = post_init(&client, &url, None, "2026-07-28").await; + assert_eq!( + response.status(), + 200, + "initialize selects legacy semantics whatever version it names" + ); + + let body: serde_json::Value = response.json().await?; + assert_eq!( + body["result"]["protocolVersion"], "2025-11-25", + "initialize must settle on a version that still has the handshake" + ); + + ct.cancel(); + Ok(()) +} + #[tokio::test] async fn stateful_init_rejects_when_header_mismatches_body() -> anyhow::Result<()> { let (client, url, ct) = spawn_server(stateful_config()).await; @@ -218,6 +250,37 @@ async fn stateful_rejected_initial_posts_do_not_create_sessions() -> anyhow::Res Ok(()) } +/// The `initialize` handshake exists only in the revisions before `2026-07-28`, +/// so naming a later version in it does not make the request a modern one: the +/// server keeps legacy semantics, opens a session, and answers with a version it +/// can actually serve over the handshake. +#[tokio::test] +async fn stateful_init_naming_a_post_handshake_version_opens_a_legacy_session() -> anyhow::Result<()> +{ + let (client, url, ct) = spawn_server(stateful_config()).await; + + let response = post_init(&client, &url, None, "2026-07-28").await; + assert_eq!( + response.status(), + 200, + "initialize selects legacy semantics whatever version it names" + ); + assert!( + response.headers().contains_key("Mcp-Session-Id"), + "the handshake must open a session, got headers: {:?}", + response.headers() + ); + + let payload = sse_payload(response).await?; + assert_eq!( + payload["result"]["protocolVersion"], "2025-11-25", + "initialize must settle on a version that still has the handshake" + ); + + ct.cancel(); + Ok(()) +} + #[tokio::test] async fn stateless_missing_protocol_header_returns_header_mismatch() -> anyhow::Result<()> { let (client, url, ct) = spawn_server(stateless_json_config()).await; From 51ccb42993d6eb5075399672ce7a0c21a0e55eea Mon Sep 17 00:00:00 2001 From: "github-actions[bot]" <41898282+github-actions[bot]@users.noreply.github.com> Date: Mon, 31 Aug 2026 19:08:18 -0400 Subject: [PATCH 14/29] chore: release v3.2.0 (#1227) Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com> --- Cargo.toml | 6 +++--- crates/rmcp-macros/CHANGELOG.md | 10 ++++++++++ crates/rmcp/CHANGELOG.md | 13 +++++++++++++ 3 files changed, 26 insertions(+), 3 deletions(-) diff --git a/Cargo.toml b/Cargo.toml index 001096245..b3a0da365 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -4,13 +4,13 @@ default-members = ["crates/rmcp", "crates/rmcp-macros"] resolver = "2" [workspace.dependencies] -rmcp = { version = "3.1.4", path = "./crates/rmcp" } -rmcp-macros = { version = "3.1.4", path = "./crates/rmcp-macros" } +rmcp = { version = "3.2.0", path = "./crates/rmcp" } +rmcp-macros = { version = "3.2.0", path = "./crates/rmcp-macros" } [workspace.package] edition = "2024" rust-version = "1.88" -version = "3.1.4" +version = "3.2.0" authors = ["4t145 "] license = "Apache-2.0" repository = "https://github.com/modelcontextprotocol/rust-sdk/" diff --git a/crates/rmcp-macros/CHANGELOG.md b/crates/rmcp-macros/CHANGELOG.md index 910257a28..4f59b7513 100644 --- a/crates/rmcp-macros/CHANGELOG.md +++ b/crates/rmcp-macros/CHANGELOG.md @@ -7,6 +7,16 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] +## [3.2.0](https://github.com/modelcontextprotocol/rust-sdk/compare/rmcp-macros-v3.1.4...rmcp-macros-v3.2.0) - 2026-08-31 + +### Added + +- add request-state key rotation ([#1128](https://github.com/modelcontextprotocol/rust-sdk/pull/1128)) + +### Fixed + +- allow concurrent streamable http requests ([#1186](https://github.com/modelcontextprotocol/rust-sdk/pull/1186)) + ## [3.1.3](https://github.com/modelcontextprotocol/rust-sdk/compare/rmcp-macros-v3.1.2...rmcp-macros-v3.1.3) - 2026-08-17 ### Fixed diff --git a/crates/rmcp/CHANGELOG.md b/crates/rmcp/CHANGELOG.md index b7b7d542a..86ab5b146 100644 --- a/crates/rmcp/CHANGELOG.md +++ b/crates/rmcp/CHANGELOG.md @@ -7,6 +7,19 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] +## [3.2.0](https://github.com/modelcontextprotocol/rust-sdk/compare/rmcp-v3.1.4...rmcp-v3.2.0) - 2026-08-31 + +### Added + +- *(auth)* coordinate OAuth refreshes through credential stores ([#1232](https://github.com/modelcontextprotocol/rust-sdk/pull/1232)) +- add request-state key rotation ([#1128](https://github.com/modelcontextprotocol/rust-sdk/pull/1128)) + +### Fixed + +- keep initialize on legacy protocol versions ([#1228](https://github.com/modelcontextprotocol/rust-sdk/pull/1228)) +- *(transport)* fall back after sessionless HTTP discover rejections ([#1211](https://github.com/modelcontextprotocol/rust-sdk/pull/1211)) +- allow concurrent streamable http requests ([#1186](https://github.com/modelcontextprotocol/rust-sdk/pull/1186)) + ## [3.1.4](https://github.com/modelcontextprotocol/rust-sdk/compare/rmcp-v3.1.3...rmcp-v3.1.4) - 2026-08-18 ### Fixed From 42cf8d9fa3e8578e33dd33ba6aee446cb574c493 Mon Sep 17 00:00:00 2001 From: "dependabot[bot]" <49699333+dependabot[bot]@users.noreply.github.com> Date: Tue, 1 Sep 2026 23:13:12 +0900 Subject: [PATCH 15/29] chore(deps): update process-wrap requirement from 9.0 to 10.0 (#1229) Updates the requirements on [process-wrap](https://github.com/watchexec/process-wrap) to permit the latest version. - [Changelog](https://github.com/watchexec/process-wrap/blob/main/CHANGELOG.md) - [Commits](https://github.com/watchexec/process-wrap/compare/v9.0.0...v10.0.0) --- updated-dependencies: - dependency-name: process-wrap dependency-version: 10.0.0 dependency-type: direct:production ... Signed-off-by: dependabot[bot] Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com> --- crates/rmcp/Cargo.toml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/crates/rmcp/Cargo.toml b/crates/rmcp/Cargo.toml index e554e8606..66cb02966 100644 --- a/crates/rmcp/Cargo.toml +++ b/crates/rmcp/Cargo.toml @@ -87,7 +87,7 @@ url = { version = "2.4", optional = true } tower-service = { version = "0.3", optional = true } # for child process transport -process-wrap = { version = "9.0", features = ["tokio1"], optional = true } +process-wrap = { version = "10.0", features = ["tokio1"], optional = true } # for cross-platform executable path resolution which = { version = "8", optional = true } From 3b5ca4dd3f2a7f34404ed3d97ea6880b0ffa90bd Mon Sep 17 00:00:00 2001 From: "dependabot[bot]" <49699333+dependabot[bot]@users.noreply.github.com> Date: Wed, 2 Sep 2026 07:01:50 +0900 Subject: [PATCH 16/29] chore(deps): bump astral-sh/setup-uv from 7.6.0 to 10.0.1 (#1238) Bumps [astral-sh/setup-uv](https://github.com/astral-sh/setup-uv) from 7.6.0 to 10.0.1. - [Release notes](https://github.com/astral-sh/setup-uv/releases) - [Commits](https://github.com/astral-sh/setup-uv/compare/37802adc94f370d6bfd71619e3f0bf239e1f3b78...20cfd1bf945f4377ade1205e4dbc17946fc9a30d) --- updated-dependencies: - dependency-name: astral-sh/setup-uv dependency-version: 10.0.1 dependency-type: direct:production update-type: version-update:semver-major ... Signed-off-by: dependabot[bot] Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com> --- .github/workflows/ci.yml | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index a43adf760..c71d9d502 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -245,7 +245,7 @@ jobs: node-version: '22' - name: Install uv - uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7 + uses: astral-sh/setup-uv@20cfd1bf945f4377ade1205e4dbc17946fc9a30d # v10.0.1 - name: Install Rust uses: dtolnay/rust-toolchain@4360b52568e2003a75bf9bc1d59f33a8e3fc893c # stable @@ -274,7 +274,7 @@ jobs: node-version: '22' - name: Install uv - uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7 + uses: astral-sh/setup-uv@20cfd1bf945f4377ade1205e4dbc17946fc9a30d # v10.0.1 - name: Install Rust uses: dtolnay/rust-toolchain@4360b52568e2003a75bf9bc1d59f33a8e3fc893c # stable @@ -310,7 +310,7 @@ jobs: node-version: '22' - name: Install uv - uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7 + uses: astral-sh/setup-uv@20cfd1bf945f4377ade1205e4dbc17946fc9a30d # v10.0.1 - name: Install Rust uses: dtolnay/rust-toolchain@4360b52568e2003a75bf9bc1d59f33a8e3fc893c # stable @@ -345,7 +345,7 @@ jobs: node-version: '22' - name: Install uv - uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7 + uses: astral-sh/setup-uv@20cfd1bf945f4377ade1205e4dbc17946fc9a30d # v10.0.1 - name: Install Rust uses: dtolnay/rust-toolchain@4360b52568e2003a75bf9bc1d59f33a8e3fc893c # stable From 302319861a4b5ab538f6aebf25befdc3c7dfe039 Mon Sep 17 00:00:00 2001 From: Nick Steele Date: Fri, 4 Sep 2026 15:04:08 -0400 Subject: [PATCH 17/29] feat(auth): add enterprise refresh-token and ID-JAG exchanges (#1234) * feat(auth): request ID-JAGs with enterprise refresh tokens * feat(auth): exchange ID-JAGs for MCP access tokens Redeem ID-JAGs at the approved resource authorization server and return its bearer token, lifetime, and effective granted scopes. Preserve scope narrowing and redacted diagnostics, and reuse the default HTTP adapter. Support independently configured client authentication at both servers. Document the exchange profile and test redirects and staged failures. Partially addresses modelcontextprotocol/rust-sdk#531. * test(conformance): cover EMA refresh-token exchange --- conformance/Cargo.toml | 2 + conformance/src/bin/client.rs | 81 +- crates/rmcp/Cargo.toml | 4 +- crates/rmcp/README.md | 5 + crates/rmcp/src/transport/auth.rs | 74 +- crates/rmcp/src/transport/auth/enterprise.rs | 767 +++++++++++ .../src/transport/auth/enterprise_tests.rs | 1145 +++++++++++++++++ docs/OAUTH_SUPPORT.md | 125 ++ 8 files changed, 2196 insertions(+), 7 deletions(-) create mode 100644 crates/rmcp/src/transport/auth/enterprise.rs create mode 100644 crates/rmcp/src/transport/auth/enterprise_tests.rs diff --git a/conformance/Cargo.toml b/conformance/Cargo.toml index fc4be6d90..e8de6f072 100644 --- a/conformance/Cargo.toml +++ b/conformance/Cargo.toml @@ -19,6 +19,7 @@ rmcp = { path = "../crates/rmcp", features = [ "elicitation", "auth", "auth-client-credentials-jwt", + "auth-enterprise-managed", "request-state", "transport-streamable-http-server", "transport-streamable-http-client-reqwest", @@ -31,6 +32,7 @@ tracing = "0.1" tracing-subscriber = { version = "0.3", features = ["env-filter"] } axum = { version = "0.8", features = ["macros"] } anyhow = "1" +oauth2 = { version = "5.0", default-features = false } reqwest = { version = "0.13", features = ["json"] } urlencoding = "2" url = "2" diff --git a/conformance/src/bin/client.rs b/conformance/src/bin/client.rs index 5d654105a..ea98a8de6 100644 --- a/conformance/src/bin/client.rs +++ b/conformance/src/bin/client.rs @@ -1,3 +1,5 @@ +use anyhow::Context; +use oauth2::{ClientSecret, RefreshToken}; use rmcp::{ ClientHandler, ClientLifecycleMode, ClientServiceExt, ErrorData, RoleClient, ServiceExt, model::*, @@ -6,7 +8,8 @@ use rmcp::{ AuthClient, AuthorizationManager, StreamableHttpClientTransport, auth::{ AuthorizationCallback, AuthorizationRequest, ClientCredentialsConfig, - InMemoryCredentialStore, JwtSigningAlgorithm, OAuthState, + InMemoryCredentialStore, JwtSigningAlgorithm, OAuthState, default_oauth_http_client, + enterprise::{EmaAuthorizationServer, EmaClientAuthentication, EmaExchangeRequest}, }, streamable_http_client::StreamableHttpClientTransportConfig, }, @@ -36,6 +39,17 @@ struct ConformanceContext { private_key_pem: Option, #[serde(default)] signing_algorithm: Option, + // enterprise-managed-authorization-refresh-token + #[serde(default)] + idp_client_id: Option, + #[serde(default)] + idp_client_secret: Option, + #[serde(default)] + idp_refresh_token: Option, + #[serde(default)] + idp_issuer: Option, + #[serde(default)] + idp_token_endpoint: Option, } fn load_context() -> ConformanceContext { @@ -760,6 +774,66 @@ async fn run_client_credentials_jwt( Ok(()) } +/// Exchange the fixture's IdP refresh token, then exercise authenticated MCP access. +async fn run_ema_refresh_token_client( + server_url: &str, + ctx: &ConformanceContext, +) -> anyhow::Result<()> { + let manager = AuthorizationManager::new(server_url).await?; + let metadata = manager.resolve_metadata().await?.metadata; + let idp = EmaAuthorizationServer::new( + ctx.idp_issuer.as_deref().context("Missing idp_issuer")?, + ctx.idp_token_endpoint + .as_deref() + .context("Missing idp_token_endpoint")?, + ctx.idp_client_id + .as_deref() + .context("Missing idp_client_id")?, + ) + .with_client_authentication(EmaClientAuthentication::ClientSecretBasic( + ClientSecret::new( + ctx.idp_client_secret + .clone() + .context("Missing idp_client_secret")?, + ), + )); + let resource_as = EmaAuthorizationServer::new( + metadata + .issuer + .context("Missing authorization server issuer")?, + metadata.token_endpoint, + ctx.client_id.as_deref().context("Missing client_id")?, + ) + .with_client_authentication(EmaClientAuthentication::ClientSecretBasic( + ClientSecret::new(ctx.client_secret.clone().context("Missing client_secret")?), + )); + let refresh_token = RefreshToken::new( + ctx.idp_refresh_token + .clone() + .context("Missing idp_refresh_token")?, + ); + let http = default_oauth_http_client()?; + let token = EmaExchangeRequest::new(idp, resource_as, server_url, &refresh_token) + .with_scopes(manager.select_scopes(None, &[])) + .exchange(&http, &http) + .await?; + + let transport = StreamableHttpClientTransport::from_config( + StreamableHttpClientTransportConfig::with_uri(server_url) + .auth_header(token.access_token.secret()), + ); + let client = BasicClientHandler + .serve_with_lifecycle(transport, conformance_lifecycle()) + .await?; + let tools = client.list_tools(Default::default()).await?; + for tool in tools.tools { + let args = build_tool_arguments(&tool); + client.call_tool(call_tool_params(tool.name, args)).await?; + } + client.cancel().await?; + Ok(()) +} + /// Cross-app access flow (SEP-1046 extension). async fn run_cross_app_access_client( server_url: &str, @@ -1110,6 +1184,11 @@ async fn run_scenario( "auth/client-credentials-basic" => run_client_credentials_basic(server_url, ctx).await?, "auth/client-credentials-jwt" => run_client_credentials_jwt(server_url, ctx).await?, + // Auth - enterprise-managed authorization with a refresh-token subject + "auth/enterprise-managed-authorization-refresh-token" => { + run_ema_refresh_token_client(server_url, ctx).await? + } + // Auth - cross-app access "auth/cross-app-access-complete-flow" => { run_cross_app_access_client(server_url, ctx).await? diff --git a/crates/rmcp/Cargo.toml b/crates/rmcp/Cargo.toml index 66cb02966..ac51ebfcd 100644 --- a/crates/rmcp/Cargo.toml +++ b/crates/rmcp/Cargo.toml @@ -18,6 +18,7 @@ exhaustive_enums = "warn" features = [ "auth", "auth-client-credentials-jwt", + "auth-enterprise-managed", "base64", "client", "client-side-sse", @@ -196,10 +197,11 @@ transport-streamable-http-server-session = [ tower = ["dep:tower-service"] auth = ["dep:async-trait", "dep:oauth2", "__reqwest", "dep:url"] auth-client-credentials-jwt = ["auth", "dep:jsonwebtoken", "uuid"] +auth-enterprise-managed = ["auth", "base64"] schemars = ["dep:schemars"] [dev-dependencies] -tokio = { version = "1", features = ["full"] } +tokio = { version = "1", features = ["full", "test-util"] } schemars = { version = "1.1.0", features = ["chrono04"] } axum = { version = "0.8", default-features = false, features = ["http1", "tokio"] } hyper = { version = "1", features = ["server", "http1"] } diff --git a/crates/rmcp/README.md b/crates/rmcp/README.md index bb7837e84..f34469ef4 100644 --- a/crates/rmcp/README.md +++ b/crates/rmcp/README.md @@ -24,6 +24,7 @@ For **getting started**, **usage guides**, and **full MCP feature documentation* | `macros` | `#[tool]` / `#[prompt]` macros (re-exports [`rmcp-macros`](../rmcp-macros)) | ✅ | | `schemars` | JSON Schema generation for tool definitions | | | `auth` | OAuth 2.0 authentication support | | +| `auth-enterprise-managed` | EMA/XAA refresh-token and ID-JAG exchanges for registered public and confidential clients (includes `auth`) | | | `elicitation` | Elicitation support | | ### Transport features @@ -45,6 +46,10 @@ For **getting started**, **usage guides**, and **full MCP feature documentation* | `reqwest-native-tls` | Uses platform-native TLS (OpenSSL / Secure Transport / SChannel) | | `reqwest-tls-no-provider` | Uses rustls without a default crypto provider (bring your own) | +For enterprise-managed authorization, enable `auth-enterprise-managed` and a TLS +backend such as `reqwest`. See the [EMA/XAA guide](../../docs/OAUTH_SUPPORT.md#enterprise-managed-authorization-emaxaa) +for client authentication and an MCP connection example. + ## Transports The transport layer is pluggable. Two built-in pairs cover the most common cases: diff --git a/crates/rmcp/src/transport/auth.rs b/crates/rmcp/src/transport/auth.rs index 6762286d9..8c9e41216 100644 --- a/crates/rmcp/src/transport/auth.rs +++ b/crates/rmcp/src/transport/auth.rs @@ -28,6 +28,9 @@ use tracing::{debug, warn}; use crate::transport::common::http_header::HEADER_MCP_PROTOCOL_VERSION; +#[cfg(feature = "auth-enterprise-managed")] +pub mod enterprise; + const DEFAULT_HTTP_TIMEOUT: Duration = Duration::from_secs(30); const MAX_OAUTH_HTTP_RESPONSE_BODY_BYTES: usize = 1024 * 1024; const MAX_OAUTH_DISCOVERY_REDIRECTS: usize = 10; @@ -99,6 +102,19 @@ pub trait OAuthHttpClient: Send + Sync { fn execute(&self, request: OAuthHttpRequest) -> OAuthHttpClientFuture<'_>; } +/// Create an OAuth HTTP client with the SDK's default reqwest configuration. +/// +/// Honors each request's redirect policy, with a 30-second timeout and bounded +/// response bodies. Enable a TLS feature such as `reqwest` for HTTPS requests. +/// Implement [`OAuthHttpClient`] instead when custom network policy is required. +pub fn default_oauth_http_client() -> Result { + let client = ReqwestClient::builder() + .timeout(DEFAULT_HTTP_TIMEOUT) + .build() + .map_err(|error| AuthError::InternalError(error.to_string()))?; + ReqwestOAuthHttpClient::new(client) +} + struct ReqwestOAuthHttpClient { follow_redirects: ReqwestClient, stop_redirects: ReqwestClient, @@ -1307,13 +1323,9 @@ impl AuthorizationManager { /// create new auth manager with base url pub async fn new(base_url: U) -> Result { - let http_client = ReqwestClient::builder() - .timeout(DEFAULT_HTTP_TIMEOUT) - .build() - .map_err(|e| AuthError::InternalError(e.to_string()))?; Self::new_inner( base_url, - Arc::new(ReqwestOAuthHttpClient::new(http_client)?), + Arc::new(default_oauth_http_client()?), OAuthHttpRedirectPolicy::Stop, ) .await @@ -4043,6 +4055,58 @@ mod tests { ); } + #[tokio::test] + async fn default_oauth_http_client_honors_redirect_policy() { + use axum::{Router, routing::post}; + + let received = Arc::new(StdMutex::new(Vec::new())); + let capture = Arc::clone(&received); + let app = Router::new() + .route( + "/redirect", + post(|| async { (StatusCode::TEMPORARY_REDIRECT, [("location", "/token")]) }), + ) + .route( + "/token", + post(move |body: String| { + let capture = Arc::clone(&capture); + async move { + capture.lock().unwrap().push(body); + StatusCode::OK + } + }), + ); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let endpoint = format!("http://{}/redirect", listener.local_addr().unwrap()); + tokio::spawn(async move { axum::serve(listener, app).await.unwrap() }); + let client = super::default_oauth_http_client().unwrap(); + + for (policy, expected) in [ + ( + OAuthHttpRedirectPolicy::Stop, + StatusCode::TEMPORARY_REDIRECT, + ), + (OAuthHttpRedirectPolicy::Follow, StatusCode::OK), + ] { + let request = oauth2::http::Request::builder() + .method("POST") + .uri(&endpoint) + .body(b"credential-sentinel".to_vec()) + .unwrap(); + let response = client + .execute(OAuthHttpRequest::new(request, policy)) + .await + .unwrap(); + assert_eq!(response.status(), expected); + let expected_bodies = if policy == OAuthHttpRedirectPolicy::Stop { + vec![] + } else { + vec!["credential-sentinel".to_owned()] + }; + assert_eq!(*received.lock().unwrap(), expected_bodies); + } + } + #[tokio::test] async fn default_http_client_preserves_connection_failure_cause() { let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap(); diff --git a/crates/rmcp/src/transport/auth/enterprise.rs b/crates/rmcp/src/transport/auth/enterprise.rs new file mode 100644 index 000000000..226c8d531 --- /dev/null +++ b/crates/rmcp/src/transport/auth/enterprise.rs @@ -0,0 +1,767 @@ +//! Non-interactive enterprise-managed authorization (EMA/XAA) token exchanges. +//! +//! Exchanges an enterprise refresh token for an ID-JAG, then for an MCP access +//! token. Callers must discover and approve both servers and their registrations +//! before supplying a credential. This module does not discover servers, log in, +//! persist credentials, or decide when to reauthenticate. +//! +//! Configure each pre-registered client's approved authentication method separately: +//! HTTP Basic, a client secret in the request body, or a freshly signed JWT client +//! assertion. Public-client authentication requires the server's explicit approval. +//! This helper requires one MCP resource and does not implement Rich +//! Authorization Requests or DPoP. Redemption consumes the SDK's ID-JAG handle, +//! without automatic retries; server-side replay policy remains the server's responsibility. +//! +//! ID-JAG checks below enforce structure and claim bindings, not cryptographic +//! signature verification. Assertions come directly from the trusted IdP token +//! endpoint; the resource authorization server must verify their signatures. +//! +//! ```no_run +//! use oauth2::{ClientSecret, RefreshToken}; +//! use rmcp::transport::auth::{default_oauth_http_client, enterprise::*}; +//! +//! # async fn authorize(refresh: &RefreshToken, idp_secret: ClientSecret, resource_secret: ClientSecret) -> Result<(), Box> { +//! // Enable `auth-enterprise-managed` and a TLS feature such as `reqwest`. +//! let http = default_oauth_http_client()?; +//! let token = EmaExchangeRequest::new( +//! EmaAuthorizationServer::new("https://idp.example", "https://idp.example/token", "idp-client") +//! .with_client_authentication(EmaClientAuthentication::ClientSecretBasic(idp_secret)), +//! EmaAuthorizationServer::new("https://as.example", "https://as.example/token", "mcp-client") +//! .with_client_authentication(EmaClientAuthentication::ClientSecretBasic(resource_secret)), +//! "https://mcp.example", refresh, +//! ).with_scopes(["files.read"]).exchange(&http, &http).await?; +//! // Use token.access_token only for the approved MCP resource; never log it. +//! // token.scopes contains the final granted scopes, which may be narrower. +//! # Ok(()) } +//! ``` + +use std::{ + collections::HashSet, + sync::Arc, + time::{Duration, SystemTime, UNIX_EPOCH}, +}; + +use base64::{ + Engine, + engine::general_purpose::{STANDARD, URL_SAFE_NO_PAD}, +}; +use oauth2::{AccessToken, ClientSecret, RefreshToken}; +use serde::{Deserialize, de::DeserializeOwned}; +use thiserror::Error; +use url::{Host, Url}; + +use super::{ + DEFAULT_HTTP_TIMEOUT, MAX_OAUTH_HTTP_RESPONSE_BODY_BYTES, OAuthHttpClient, + OAuthHttpRedirectPolicy, OAuthHttpRequest, +}; + +const ID_JAG_TOKEN_TYPE: &str = "urn:ietf:params:oauth:token-type:id-jag"; + +/// Authentication approved for a pre-registered client at one authorization server. +/// +/// The selected method is used as configured, without negotiation or fallback. +#[derive(Clone)] +#[non_exhaustive] +pub enum EmaClientAuthentication { + /// Public client (`token_endpoint_auth_method=none`), only if the server permits it. + None, + /// `client_secret_basic`, with OAuth form encoding before HTTP Basic encoding. + ClientSecretBasic(ClientSecret), + /// `client_secret_post`, for servers requiring credentials in the request body. + ClientSecretPost(ClientSecret), + /// Fresh JWT client assertions, such as `private_key_jwt` or `client_secret_jwt`. + /// Signing, key custody, claims, and the registered algorithm belong to the provider. + JwtAssertion(Arc), +} + +impl std::fmt::Debug for EmaClientAuthentication { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.write_str(match self { + Self::None => "None", + Self::ClientSecretBasic(_) => "ClientSecretBasic { .. }", + Self::ClientSecretPost(_) => "ClientSecretPost { .. }", + Self::JwtAssertion(_) => "JwtAssertion { .. }", + }) + } +} + +impl EmaClientAuthentication { + fn validate(&self) -> Result<(), EmaError> { + if let Self::ClientSecretBasic(secret) | Self::ClientSecretPost(secret) = self + && secret.secret().trim().is_empty() + { + return Err(EmaError::InvalidRequest("client secret must not be empty")); + } + Ok(()) + } +} + +/// A signed JWT used to authenticate the client, distinct from the ID-JAG grant. +pub struct EmaClientAssertion(String); + +impl EmaClientAssertion { + /// Wrap a fresh assertion without exposing it through `Debug`. + pub fn new(assertion: String) -> Self { + Self(assertion) + } +} + +impl std::fmt::Debug for EmaClientAssertion { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.write_str("EmaClientAssertion { .. }") + } +} + +/// Creates client assertions on demand, allowing keys to remain in an external signer. +/// +/// Called once immediately before each token request, including delayed ID-JAG +/// redemption. Set `iss` and `sub` to the registered client identifier and `aud` +/// to the server's approved audience, with a short expiration and a fresh `jti`. +/// The SDK does not sign or validate these assertions. Provider failures are +/// sanitized and the call shares the token request's timeout. Cancellation may +/// occur when that deadline expires. +#[async_trait::async_trait] +pub trait EmaClientAssertionProvider: Send + Sync { + async fn create_assertion( + &self, + server: &EmaAuthorizationServer, + ) -> Result>; +} + +/// A trusted server with a pre-registered client and its approved authentication method. +#[derive(Clone)] +#[non_exhaustive] +pub struct EmaAuthorizationServer { + /// Exact issuer identifier from approved metadata. + pub issuer: String, + /// Token endpoint from that metadata. + pub token_endpoint: String, + /// Pre-registered client identifier. + pub client_id: String, + client_authentication: EmaClientAuthentication, +} + +impl EmaAuthorizationServer { + /// Use approved metadata and a public-client registration (`token_endpoint_auth_method=none`). + /// Set [`Self::with_client_authentication`] for a confidential client. + pub fn new( + issuer: impl Into, + token_endpoint: impl Into, + client_id: impl Into, + ) -> Self { + Self { + issuer: issuer.into(), + token_endpoint: token_endpoint.into(), + client_id: client_id.into(), + client_authentication: EmaClientAuthentication::None, + } + } + + /// Select the authentication method approved for this server's client registration. + pub fn with_client_authentication(mut self, authentication: EmaClientAuthentication) -> Self { + self.client_authentication = authentication; + self + } +} + +/// A refresh-token exchange bound to one MCP resource and two registered clients. +pub struct EmaExchangeRequest<'a> { + idp: EmaAuthorizationServer, + resource_as: EmaAuthorizationServer, + resource: &'a str, + refresh_token: &'a RefreshToken, + scopes: Vec, +} + +impl<'a> EmaExchangeRequest<'a> { + /// No scope parameter is sent until [`Self::with_scopes`] is used. + pub fn new( + idp: EmaAuthorizationServer, + resource_as: EmaAuthorizationServer, + resource: &'a str, + refresh_token: &'a RefreshToken, + ) -> Self { + Self { + idp, + resource_as, + resource, + refresh_token, + scopes: Vec::new(), + } + } + + /// Request distinct non-empty scope tokens; the IdP may narrow them. + /// An empty iterator omits `scope`, rather than requesting an empty grant. + pub fn with_scopes(mut self, scopes: I) -> Self + where + I: IntoIterator, + S: Into, + { + self.scopes = scopes.into_iter().map(Into::into).collect(); + self + } + + /// Obtain an ID-JAG without redirects or retries, checking its resource/client bindings. + /// The returned assertion does not retain the enterprise refresh token. + pub async fn exchange_id_jag( + self, + idp_http: &dyn OAuthHttpClient, + ) -> Result { + for endpoint in [ + self.resource, + &self.idp.issuer, + &self.idp.token_endpoint, + &self.resource_as.issuer, + &self.resource_as.token_endpoint, + ] { + validate_endpoint(endpoint)?; + } + if self.idp.issuer == self.resource_as.issuer { + return Err(EmaError::InvalidRequest( + "IdP and resource AS issuers must differ", + )); + } + if self.idp.client_id.trim().is_empty() + || self.resource_as.client_id.trim().is_empty() + || self.refresh_token.secret().trim().is_empty() + { + return Err(EmaError::InvalidRequest( + "client IDs and refresh token must not be empty", + )); + } + self.idp.client_authentication.validate()?; + self.resource_as.client_authentication.validate()?; + let requested: HashSet<&str> = self.scopes.iter().map(String::as_str).collect(); + if requested.len() != self.scopes.len() || self.scopes.iter().any(|s| !is_scope_token(s)) { + return Err(EmaError::InvalidRequest( + "scopes must be distinct non-empty tokens", + )); + } + let mut params = vec![ + ( + "grant_type", + "urn:ietf:params:oauth:grant-type:token-exchange", + ), + ("requested_token_type", ID_JAG_TOKEN_TYPE), + ( + "subject_token_type", + "urn:ietf:params:oauth:token-type:refresh_token", + ), + ("subject_token", self.refresh_token.secret()), + ("audience", self.resource_as.issuer.as_str()), + ("resource", self.resource), + ]; + let scope = self.scopes.join(" "); + if !scope.is_empty() { + params.push(("scope", &scope)); + } + let jag: IdJagResponse = post_form( + idp_http, + &self.idp, + ¶ms, + EmaExchangeStage::IdentityProvider, + None, + unix_time, + ) + .await?; + let (granted, expires_at) = jag.validate(&self, &requested)?; + Ok(EmaIdJag { + assertion: AccessToken::new(jag.access_token), + scopes: granted, + resource_as: self.resource_as, + resource: self.resource.to_owned(), + expires_at, + }) + } + + /// Perform both exchanges with independently routed HTTP clients and no automatic retries. + pub async fn exchange( + self, + idp_http: &dyn OAuthHttpClient, + resource_http: &dyn OAuthHttpClient, + ) -> Result { + self.exchange_id_jag(idp_http) + .await? + .exchange(resource_http) + .await + } +} + +/// An IdP-issued ID-JAG whose structure and bindings have been checked, not its signature. +pub struct EmaIdJag { + assertion: AccessToken, + scopes: HashSet, + resource_as: EmaAuthorizationServer, + resource: String, + expires_at: u64, +} + +impl std::fmt::Debug for EmaAuthorizationServer { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.write_str("EmaAuthorizationServer { .. }") + } +} + +impl std::fmt::Debug for EmaExchangeRequest<'_> { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.write_str("EmaExchangeRequest { .. }") + } +} + +impl std::fmt::Debug for EmaIdJag { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.write_str("EmaIdJag { .. }") + } +} + +impl EmaIdJag { + /// The assertion to present to the approved resource authorization server. + pub fn assertion(&self) -> &AccessToken { + &self.assertion + } + + /// The scope tokens carried by the assertion; an empty set means scope was omitted. + pub fn scopes(&self) -> &HashSet { + &self.scopes + } + + /// Redeem this assertion once, at its approved resource AS, without redirects or retries. + /// Consume the grant so this helper cannot accidentally replay it after a failed exchange. + pub async fn exchange(self, http: &dyn OAuthHttpClient) -> Result { + self.exchange_with_clock(http, unix_time).await + } + + async fn exchange_with_clock( + self, + http: &dyn OAuthHttpClient, + now: impl Fn() -> Result + Sync, + ) -> Result { + if self.expires_at <= now()? { + return Err(EmaError::InvalidRequest("ID-JAG expired before redemption")); + } + // Only the assertion carries authority: repeating resource/scope could undo narrowing. + let token: ResourceTokenResponse = post_form( + http, + &self.resource_as, + &[ + ("grant_type", "urn:ietf:params:oauth:grant-type:jwt-bearer"), + ("assertion", self.assertion.secret()), + ], + EmaExchangeStage::ResourceAuthorizationServer, + Some(self.expires_at), + now, + ) + .await?; + token.validate(&self.resource, self.scopes) + } +} + +/// A resource-bound bearer with secret-safe diagnostics and its optional lifetime. +#[derive(Clone)] +#[non_exhaustive] +pub struct EmaAccessToken { + pub access_token: AccessToken, + pub expires_in: Option, + /// Resource-AS scopes, or the ID-JAG scopes when the response omits `scope`. + /// Empty when both the ID-JAG and response omit scopes. + pub scopes: HashSet, +} + +impl std::fmt::Debug for EmaAccessToken { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.write_str("EmaAccessToken { .. }") + } +} + +/// The endpoint that failed, allowing the caller to apply its own credential lifecycle policy. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +#[non_exhaustive] +pub enum EmaExchangeStage { + IdentityProvider, + ResourceAuthorizationServer, +} + +/// Sanitized failures. Raw HTTP adapter errors and provider response bodies are never retained. +#[derive(Debug, Error, PartialEq, Eq)] +#[non_exhaustive] +pub enum EmaError { + #[error("invalid EMA exchange request: {0}")] + InvalidRequest(&'static str), + #[error("invalid EMA response from {stage:?}: {message}")] + InvalidResponse { + stage: EmaExchangeStage, + message: &'static str, + }, + #[error("EMA request to {0:?} failed")] + RequestFailed(EmaExchangeStage), + #[error("{0:?}: invalid_grant")] + InvalidGrant(EmaExchangeStage), + #[error("{0:?}: insufficient_user_authentication")] + InsufficientUserAuthentication(EmaExchangeStage), + #[error("{stage:?} returned HTTP {status}: {code}")] + OAuthRejected { + stage: EmaExchangeStage, + status: u16, + code: &'static str, + }, +} + +async fn post_form( + http: &dyn OAuthHttpClient, + server: &EmaAuthorizationServer, + params: &[(&str, &str)], + stage: EmaExchangeStage, + grant_expires_at: Option, + now: impl Fn() -> Result + Sync, +) -> Result { + let response = tokio::time::timeout(DEFAULT_HTTP_TIMEOUT, async { + // Generate assertions at the request boundary, not when server configuration is built. + let client_assertion = match &server.client_authentication { + EmaClientAuthentication::JwtAssertion(provider) => { + let assertion = provider + .create_assertion(server) + .await + .map_err(|_| EmaError::RequestFailed(stage))?; + if assertion.0.trim().is_empty() { + return Err(EmaError::InvalidRequest( + "client assertion must not be empty", + )); + } + Some(assertion) + } + _ => None, + }; + let request = { + let mut form = url::form_urlencoded::Serializer::new(String::new()); + form.extend_pairs(params.iter().copied()); + let mut request = oauth2::http::Request::builder() + .method("POST") + .uri(&server.token_endpoint) + .header("content-type", "application/x-www-form-urlencoded") + .header("accept", "application/json"); + match &server.client_authentication { + EmaClientAuthentication::ClientSecretBasic(secret) => { + let client_id: String = + url::form_urlencoded::byte_serialize(server.client_id.as_bytes()).collect(); + let secret: String = + url::form_urlencoded::byte_serialize(secret.secret().as_bytes()).collect(); + let encoded = STANDARD.encode(format!("{client_id}:{secret}")); + let mut header = + oauth2::http::HeaderValue::from_str(&format!("Basic {encoded}")).map_err( + |_| EmaError::InvalidRequest("invalid client authentication header"), + )?; + header.set_sensitive(true); + request = request.header(oauth2::http::header::AUTHORIZATION, header); + } + EmaClientAuthentication::ClientSecretPost(secret) => { + form.append_pair("client_id", &server.client_id) + .append_pair("client_secret", secret.secret()); + } + EmaClientAuthentication::None | EmaClientAuthentication::JwtAssertion(_) => { + form.append_pair("client_id", &server.client_id); + } + } + if let Some(assertion) = client_assertion { + form.append_pair( + "client_assertion_type", + "urn:ietf:params:oauth:client-assertion-type:jwt-bearer", + ) + .append_pair("client_assertion", &assertion.0); + } + request + .body(form.finish().into_bytes()) + .map_err(|_| EmaError::InvalidRequest("invalid token endpoint URI"))? + }; + // An external signer may outlive the grant even when it meets the request deadline. + if let Some(expires_at) = grant_expires_at + && expires_at <= now()? + { + return Err(EmaError::InvalidRequest("ID-JAG expired before redemption")); + } + http.execute(OAuthHttpRequest::new( + request, + OAuthHttpRedirectPolicy::Stop, + )) + .await + .map_err(|_| EmaError::RequestFailed(stage)) + }) + .await + .map_err(|_| EmaError::RequestFailed(stage))??; + let invalid = |message| EmaError::InvalidResponse { stage, message }; + if response.body().len() > MAX_OAUTH_HTTP_RESPONSE_BODY_BYTES { + return Err(invalid("response body too large")); + } + if !response.status().is_success() { + #[derive(Deserialize)] + struct OAuthError { + error: Option, + } + let error = serde_json::from_slice::(response.body()).ok(); + let code = match error.as_ref().and_then(|e| e.error.as_deref()) { + Some("invalid_grant") => return Err(EmaError::InvalidGrant(stage)), + Some("insufficient_user_authentication") => { + return Err(EmaError::InsufficientUserAuthentication(stage)); + } + Some("invalid_request") => "invalid_request", + Some("invalid_client") => "invalid_client", + Some("invalid_scope") => "invalid_scope", + Some("invalid_target") => "invalid_target", + Some("unauthorized_client") => "unauthorized_client", + Some("unsupported_grant_type") => "unsupported_grant_type", + Some("access_denied") => "access_denied", + Some("temporarily_unavailable") => "temporarily_unavailable", + Some("server_error") => "server_error", + _ => "OAuth token request rejected", + }; + return Err(EmaError::OAuthRejected { + stage, + status: response.status().as_u16(), + code, + }); + } + serde_json::from_slice(response.body()).map_err(|_| invalid("malformed token response")) +} + +fn validate_endpoint(value: &str) -> Result<(), EmaError> { + let url = Url::parse(value).map_err(|_| EmaError::InvalidRequest("invalid endpoint URL"))?; + let loopback = match url.host() { + Some(Host::Domain(host)) => host.eq_ignore_ascii_case("localhost"), + Some(Host::Ipv4(ip)) => ip.is_loopback(), + Some(Host::Ipv6(ip)) => ip.is_loopback(), + None => false, + }; + if (url.scheme() != "https" && !(url.scheme() == "http" && loopback)) + || !url.username().is_empty() + || url.password().is_some() + || url.fragment().is_some() + { + return Err(EmaError::InvalidRequest( + "endpoint must use HTTPS or HTTP loopback without userinfo or fragments", + )); + } + Ok(()) +} + +#[derive(Deserialize)] +#[serde(untagged)] +enum Resource { + Single(String), + Multiple(Vec), +} + +impl Resource { + fn is_exact(&self, expected: &str) -> bool { + match self { + Self::Single(value) => value == expected, + Self::Multiple(values) => values.as_slice() == [expected], + } + } +} + +#[derive(Deserialize)] +struct JwtHeader { + alg: String, + typ: Option, +} + +#[derive(Deserialize)] +struct IdJagClaims { + iss: String, + sub: String, + aud: Resource, + client_id: String, + jti: String, + exp: u64, + iat: u64, + resource: Resource, + scope: Option, + authorization_details: Option>, +} + +#[derive(Deserialize)] +struct IdJagResponse { + access_token: String, + issued_token_type: String, + token_type: String, + resource: Option, + scope: Option, + refresh_token: Option, + authorization_details: Option>, +} + +impl IdJagResponse { + fn validate( + &self, + request: &EmaExchangeRequest<'_>, + requested: &HashSet<&str>, + ) -> Result<(HashSet, u64), EmaError> { + let invalid = |message| EmaError::InvalidResponse { + stage: EmaExchangeStage::IdentityProvider, + message, + }; + if self.issued_token_type != ID_JAG_TOKEN_TYPE + || self.token_type != "N_A" + || self.refresh_token.is_some() + { + return Err(invalid("unsupported ID-JAG token type or refresh token")); + } + let mut parts = self.access_token.split('.'); + let (Some(header), Some(payload), Some(signature), None) = + (parts.next(), parts.next(), parts.next(), parts.next()) + else { + return Err(invalid("ID-JAG must be a compact signed JWT")); + }; + if header.is_empty() || payload.is_empty() || signature.is_empty() { + return Err(invalid("ID-JAG contains an empty JWT segment")); + } + let decode = |value| { + URL_SAFE_NO_PAD + .decode(value) + .map_err(|_| invalid("malformed ID-JAG encoding")) + }; + decode(signature)?; + let header: JwtHeader = serde_json::from_slice(&decode(header)?) + .map_err(|_| invalid("malformed ID-JAG header"))?; + let claims: IdJagClaims = serde_json::from_slice(&decode(payload)?) + .map_err(|_| invalid("malformed ID-JAG claims"))?; + if [&self.authorization_details, &claims.authorization_details] + .into_iter() + .any(|details| details.as_ref().is_some_and(|details| !details.is_empty())) + { + return Err(invalid("authorization_details is not supported")); + } + if header.alg.trim().is_empty() + || header.alg.eq_ignore_ascii_case("none") + || header.typ.as_deref() != Some("oauth-id-jag+jwt") + || claims.iss != request.idp.issuer + || !claims.aud.is_exact(&request.resource_as.issuer) + || claims.client_id != request.resource_as.client_id + || claims.sub.trim().is_empty() + || claims.jti.trim().is_empty() + { + return Err(invalid( + "ID-JAG type, issuer, audience, client, subject, or JWT ID mismatch", + )); + } + let now = unix_time()?; + if claims.exp <= now || claims.iat > now.saturating_add(60) { + return Err(invalid("expired or future-issued ID-JAG")); + } + if !claims.resource.is_exact(request.resource) + || self + .resource + .as_ref() + .is_some_and(|r| !r.is_exact(request.resource)) + { + return Err(invalid("ID-JAG resource mismatch")); + } + let parse = + |scope| parse_scope(scope).ok_or_else(|| invalid("malformed or duplicate scopes")); + let granted = match claims.scope.as_deref() { + Some(scope) => parse(scope)?, + None if requested.is_empty() => HashSet::new(), + None => return Err(invalid("ID-JAG is missing requested scope authorization")), + }; + if !requested.is_empty() && !granted.is_subset(requested) { + return Err(invalid("ID-JAG scope exceeds the request")); + } + match self.scope.as_deref() { + Some(scope) if parse(scope)? != granted => { + return Err(invalid("response scope differs from ID-JAG scope")); + } + None if !requested.is_empty() && granted != *requested => { + return Err(invalid("response omitted narrowed scope")); + } + _ => {} + } + Ok((granted.into_iter().map(str::to_owned).collect(), claims.exp)) + } +} + +fn parse_scope(scope: &str) -> Option> { + let scopes: HashSet<_> = scope.split(' ').collect(); + (scopes.iter().all(|s| is_scope_token(s)) && scopes.len() == scope.split(' ').count()) + .then_some(scopes) +} + +fn is_scope_token(scope: &str) -> bool { + !scope.is_empty() + && scope + .bytes() + .all(|b| matches!(b, b'!' | b'#'..=b'[' | b']'..=b'~')) +} + +fn unix_time() -> Result { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .map(|duration| duration.as_secs()) + .map_err(|_| EmaError::InvalidRequest("system clock precedes the Unix epoch")) +} + +#[derive(Deserialize)] +struct ResourceTokenResponse { + access_token: String, + token_type: String, + expires_in: Option, + resource: Option, + scope: Option, + refresh_token: Option, + authorization_details: Option>, +} + +impl ResourceTokenResponse { + fn validate( + self, + resource: &str, + granted: HashSet, + ) -> Result { + let invalid = |message| EmaError::InvalidResponse { + stage: EmaExchangeStage::ResourceAuthorizationServer, + message, + }; + if self + .authorization_details + .as_ref() + .is_some_and(|details| !details.is_empty()) + { + return Err(invalid("authorization_details is not supported")); + } + if !self.token_type.eq_ignore_ascii_case("bearer") + || self.access_token.trim().is_empty() + || self.refresh_token.is_some() + || self.expires_in == Some(0) + { + return Err(invalid( + "invalid bearer token, lifetime, or unexpected refresh token", + )); + } + // The resource need not be echoed, but must agree with the ID-JAG if present. + if self + .resource + .as_ref() + .is_some_and(|r| !r.is_exact(resource)) + { + return Err(invalid("access token resource mismatch")); + } + // An omitted scope retains the authority carried by the assertion. + let scopes = if let Some(scope) = self.scope.as_deref() { + let scopes = + parse_scope(scope).ok_or_else(|| invalid("malformed or duplicate scopes"))?; + if !scopes.iter().all(|s| granted.contains(*s)) { + return Err(invalid("access token scope exceeds ID-JAG scope")); + } + scopes.into_iter().map(str::to_owned).collect() + } else { + granted + }; + Ok(EmaAccessToken { + access_token: AccessToken::new(self.access_token), + expires_in: self.expires_in.map(Duration::from_secs), + scopes, + }) + } +} + +#[cfg(test)] +#[path = "enterprise_tests.rs"] +mod tests; diff --git a/crates/rmcp/src/transport/auth/enterprise_tests.rs b/crates/rmcp/src/transport/auth/enterprise_tests.rs new file mode 100644 index 000000000..31a63c1b2 --- /dev/null +++ b/crates/rmcp/src/transport/auth/enterprise_tests.rs @@ -0,0 +1,1145 @@ +use std::{ + collections::BTreeMap, + sync::{Arc, Mutex}, + time::{Duration, SystemTime}, +}; + +use base64::{ + Engine, + engine::general_purpose::{STANDARD, URL_SAFE_NO_PAD}, +}; +use oauth2::{ClientSecret, HttpResponse, RefreshToken}; +use serde_json::{Value, json}; + +use super::{ + EmaExchangeStage::{IdentityProvider as Idp, ResourceAuthorizationServer as ResourceServer}, + *, +}; +use crate::transport::auth::{OAuthHttpClientError, OAuthHttpClientFuture, OAuthHttpRequest}; + +const IDP: &str = "https://idp.example?private-query"; +const AS: &str = "https://as.example?private-query"; +const RESOURCE: &str = "https://mcp.example?private-query"; +const IDP_TOKEN: &str = "https://idp.example/token?private-query"; +const AS_TOKEN: &str = "https://as.example/token?private-query"; +const BAD_SCOPES: &[&str] = &["", "files\tread", "\"", "\\", "\0", "读"]; + +#[derive(Default)] +struct MockHttp { + requests: Mutex>, + response: Mutex>>, +} + +impl MockHttp { + fn new(status: u16, body: Value) -> Self { + Self { + response: Mutex::new(Some(Ok(oauth2::http::Response::builder() + .status(status) + .body(serde_json::to_vec(&body).unwrap()) + .unwrap()))), + ..Self::default() + } + } +} + +impl OAuthHttpClient for MockHttp { + fn execute(&self, request: OAuthHttpRequest) -> OAuthHttpClientFuture<'_> { + self.requests.lock().unwrap().push(request); + let response = self.response.lock().unwrap().take(); + let response = response.expect("unexpected HTTP"); + Box::pin(async move { response }) + } +} + +fn jwt(header: Value, claims: &Value) -> String { + format!( + "{}.{}.{}", + URL_SAFE_NO_PAD.encode(header.to_string()), + URL_SAFE_NO_PAD.encode(claims.to_string()), + URL_SAFE_NO_PAD.encode(b"synthetic-signature") + ) +} + +fn claims() -> Value { + let now = SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap() + .as_secs(); + json!({"iss":IDP,"aud":AS,"sub":"user","client_id":"mcp", + "jti":"jag-id","iat":now,"exp":now + 3600,"resource":RESOURCE,"scope":"files.read"}) +} + +fn jag(claims: &Value) -> Value { + let mut body = json!({"access_token":jwt(json!({"alg":"ES256","typ":"oauth-id-jag+jwt"}), claims), + "issued_token_type":ID_JAG_TOKEN_TYPE,"token_type":"N_A","resource":claims["resource"]}); + if let Some(scope) = claims.get("scope") { + body["scope"] = scope.clone(); + } + body +} + +async fn exchange(idp: &MockHttp, scopes: &str) -> Result { + let scopes = scopes + .split_ascii_whitespace() + .map(str::to_owned) + .collect::>(); + let refresh = RefreshToken::new("refresh-token".into()); + let resource_as = EmaAuthorizationServer::new(AS, AS_TOKEN, "mcp"); + assert!(!format!("{resource_as:?}").contains("private-query")); + let request = EmaExchangeRequest::new( + EmaAuthorizationServer::new(IDP, IDP_TOKEN, "idp"), + resource_as, + RESOURCE, + &refresh, + ) + .with_scopes(scopes); + assert!(!format!("{request:?}").contains(refresh.secret())); + assert!(!format!("{request:?}").contains("private-query")); + let future = request.exchange_id_jag(idp); + fn is_send(_: &T) {} + is_send(&future); + future.await +} + +fn form(request: &OAuthHttpRequest) -> BTreeMap { + assert!(!request.request.headers().contains_key("authorization")); + authenticated_form(request) +} + +fn authenticated_form(request: &OAuthHttpRequest) -> BTreeMap { + assert_eq!(request.request.method(), "POST"); + let headers = request.request.headers(); + assert_eq!(headers["content-type"], "application/x-www-form-urlencoded"); + assert_eq!(headers["accept"], "application/json"); + assert_eq!(request.redirect_policy, OAuthHttpRedirectPolicy::Stop); + assert_eq!(request.timeout, Some(Duration::from_secs(30))); + let pairs = url::form_urlencoded::parse(request.request.body()) + .into_owned() + .collect::>(); + let fields = pairs.iter().cloned().collect::>(); + assert_eq!(pairs.len(), fields.len(), "duplicate form fields"); + fields +} + +#[tokio::test] +async fn refresh_exchange_preserves_exact_forms_and_signed_narrowing() { + for scopes in ["files.read files.write", ""] { + let response = jag(&claims()); + let idp = MockHttp::new(200, response.clone()); + let result = exchange(&idp, scopes).await.unwrap(); + assert_eq!( + result.assertion().secret(), + response["access_token"].as_str().unwrap() + ); + assert_eq!(result.scopes(), &HashSet::from(["files.read".to_owned()])); + assert!(!format!("{result:?}").contains(result.assertion().secret())); + assert!(!format!("{result:?}").contains("private-query")); + let requests = idp.requests.lock().unwrap(); + assert_eq!(requests.len(), 1); + assert_eq!(requests[0].request.uri(), IDP_TOKEN); + let mut fields = form(&requests[0]); + assert_eq!( + fields.remove("scope"), + (!scopes.is_empty()).then(|| scopes.to_owned()) + ); + assert_eq!( + serde_json::to_value(fields).unwrap(), + json!({ + "grant_type":"urn:ietf:params:oauth:grant-type:token-exchange", + "requested_token_type":ID_JAG_TOKEN_TYPE,"subject_token":"refresh-token", + "subject_token_type":"urn:ietf:params:oauth:token-type:refresh_token", + "audience":AS,"resource":RESOURCE,"client_id":"idp" + }) + ); + } +} + +#[tokio::test] +async fn scopes_may_be_omitted_but_never_widened() { + for (requested, signed, echoed, valid) in [ + ("", Some("files.read"), None, true), + ("", None, None, true), + ("files.read", Some("files.read"), None, true), + ("files.read files.write", Some("files.read"), None, false), + ( + "files.read", + Some("files.read files.write"), + Some("files.read files.write"), + false, + ), + ("files.read", None, None, false), + ("", Some("files.read"), Some("files.write"), false), + ("", Some(" \t"), None, false), + ("", Some("files.read files.read"), None, false), + ("", Some("files.read"), Some(" \t"), false), + ("", Some("files.read"), Some("files.read files.read"), false), + ] { + let mut claims = claims(); + claims.as_object_mut().unwrap().remove("scope"); + if let Some(scope) = signed { + claims["scope"] = json!(scope); + } + let mut response = jag(&claims); + response.as_object_mut().unwrap().remove("scope"); + if let Some(scope) = echoed { + response["scope"] = json!(scope); + } + let result = exchange(&MockHttp::new(200, response), requested).await; + assert_eq!( + result.is_ok(), + valid, + "requested={requested:?}, signed={signed:?}, echoed={echoed:?}" + ); + if let Ok(result) = result { + let expected = signed.into_iter().map(str::to_owned).collect(); + assert_eq!(result.scopes(), &expected); + } + } +} + +#[tokio::test] +async fn unsupported_authorization_details_are_rejected_at_each_stage() { + const SECRET: &str = "authorization-details-secret"; + for (location, stage) in [ + ("claims", Idp), + ("idp response", Idp), + ("resource response", ResourceServer), + ] { + for (case, details, valid) in [ + ("absent", None, true), + ("null", Some(Value::Null), true), + ("empty", Some(json!([])), true), + ( + "additional authority", + Some( + json!([{"type":SECRET,"locations":["https://other.example"],"actions":["write"]}]), + ), + false, + ), + ("invalid member", Some(json!([null])), false), + ("object", Some(json!({"type":SECRET})), false), + ("string", Some(json!(SECRET)), false), + ] { + let mut claims = claims(); + let mut response = bearer(); + // Unrelated extension fields remain compatible at every boundary. + claims["vendor_extension"] = json!(SECRET); + response["vendor_extension"] = json!(SECRET); + if location == "claims" + && let Some(details) = &details + { + claims["authorization_details"] = details.clone(); + } + let mut idp_response = jag(&claims); + idp_response["vendor_extension"] = json!(SECRET); + if let Some(details) = details { + match location { + "idp response" => idp_response["authorization_details"] = details, + "resource response" => response["authorization_details"] = details, + _ => {} + } + } + let idp = MockHttp::new(200, idp_response); + let resource = MockHttp::new(200, response); + let result = EmaExchangeRequest::new( + EmaAuthorizationServer::new(IDP, IDP_TOKEN, "idp"), + EmaAuthorizationServer::new(AS, AS_TOKEN, "mcp"), + RESOURCE, + &RefreshToken::new("refresh-token".into()), + ) + .with_scopes(["files.read"]) + .exchange(&idp, &resource) + .await; + assert_eq!(result.is_ok(), valid, "{location}: {case}"); + if let Err(error) = result { + assert!( + matches!(error, EmaError::InvalidResponse { stage: actual, .. } if actual == stage) + ); + assert!(!format!("{error:?} {error}").contains(SECRET)); + } + assert_eq!(idp.requests.lock().unwrap().len(), 1); + assert_eq!( + resource.requests.lock().unwrap().len(), + usize::from(valid || stage == ResourceServer), + "{location}: {case}" + ); + } + } +} + +#[tokio::test] +async fn invalid_jags_never_escape_validation() { + let original = claims(); + let mut cases = Vec::new(); + let invalid = json!({"iss":"https://other.example","aud":[AS,"other"], + "client_id":"other","sub":"","jti":" \t","exp":0,"iat":u64::MAX,"resource":[RESOURCE,"other"]}); + for (key, value) in invalid.as_object().unwrap() { + let mut changed = original.clone(); + changed[key] = value.clone(); + cases.push(jag(&changed)); + } + for scope in BAD_SCOPES { + let mut changed = original.clone(); + changed["scope"] = json!(scope); + cases.push(jag(&changed)); + } + for header in [ + json!({"alg":"ES256","typ":"JWT"}), + json!({"alg":"ES256"}), + json!({"alg":"none","typ":"oauth-id-jag+jwt"}), + ] { + let mut response = jag(&original); + response["access_token"] = json!(jwt(header, &original)); + cases.push(response); + } + let invalid = json!({"issued_token_type":"Bearer","token_type":"Bearer", + "refresh_token":"unsupported","resource":"https://other.example","access_token":"a.b.c.d"}); + for (key, value) in invalid.as_object().unwrap() { + let mut response = jag(&original); + response[key] = value.clone(); + cases.push(response); + } + for signature in ["", "signature", "not+base64url", "c2ln="] { + let mut response = jag(&original); + let assertion = response["access_token"].as_str().unwrap(); + let (signed, _) = assertion.rsplit_once('.').unwrap(); + response["access_token"] = json!(format!("{signed}.{signature}")); + cases.push(response); + } + for response in cases { + let result = exchange(&MockHttp::new(200, response), "files.read").await; + assert!(matches!( + result, + Err(EmaError::InvalidResponse { stage: Idp, .. }) + )); + } +} + +#[tokio::test] +async fn invalid_inputs_fail_before_http() { + // Each server occupies issuer, token endpoint, and client ID slots. + let original = [IDP, IDP_TOKEN, "idp", AS, AS_TOKEN, "mcp", RESOURCE, "rt"]; + let mut cases = Vec::new(); + for index in [0, 1, 3, 4, 6] { + for value in [ + "invalid", + "http://idp.example/token", + "https://user:pass@idp.example/token", + "https://idp.example/token#fragment", + ] { + let mut fields = original; + fields[index] = value; + cases.push((fields, vec![])); + } + } + for (index, value) in [(0, AS), (2, " "), (5, ""), (7, " \t")] { + let mut fields = original; + fields[index] = value; + cases.push((fields, vec![])); + } + for scope in BAD_SCOPES.iter().copied().chain(["files.read files.write"]) { + cases.push((original, vec![scope])); + } + cases.push((original, vec!["files.read", "files.read"])); + for (fields, scopes) in cases { + let http = MockHttp::default(); + let result = EmaExchangeRequest::new( + EmaAuthorizationServer::new(fields[0], fields[1], fields[2]), + EmaAuthorizationServer::new(fields[3], fields[4], fields[5]), + fields[6], + &RefreshToken::new(fields[7].into()), + ) + .with_scopes(scopes) + .exchange_id_jag(&http) + .await; + assert!(matches!(result, Err(EmaError::InvalidRequest(_)))); + assert!(http.requests.lock().unwrap().is_empty()); + } +} + +#[tokio::test] +async fn errors_and_redirects_cannot_reflect_credentials() { + const SECRET: &str = "secret-error-sentinel"; + for (status, code) in [ + (400, "invalid_grant"), + (400, "insufficient_user_authentication"), + (400, "invalid_client"), + (400, SECRET), + (302, SECRET), + ] { + let failure = json!({"error":code,"error_description":SECRET}); + let error = exchange(&MockHttp::new(status, failure), "") + .await + .unwrap_err(); + match code { + "invalid_grant" => assert_eq!(error, EmaError::InvalidGrant(Idp)), + "insufficient_user_authentication" => { + assert_eq!(error, EmaError::InsufficientUserAuthentication(Idp)) + } + _ => assert!( + matches!(error, EmaError::OAuthRejected {stage: Idp, status: actual, ..} if actual == status) + ), + } + assert!(!format!("{error:?} {error}").contains(SECRET)); + } + for adapter_failure in [false, true] { + let failure = MockHttp::new( + 200, + json!({"access_token":SECRET,"issued_token_type":SECRET}), + ); + if adapter_failure { + *failure.response.lock().unwrap() = Some(Err(SECRET.into())); + } + let error = exchange(&failure, "").await.unwrap_err(); + assert!(!format!("{error:?} {error}").contains(SECRET)); + assert!(std::error::Error::source(&error).is_none()); + if adapter_failure { + assert_eq!(error, EmaError::RequestFailed(Idp)); + } + } + let mut oversized = jag(&claims()); + oversized["ignored"] = json!("x".repeat(1024 * 1024)); + let error = exchange(&MockHttp::new(200, oversized), "") + .await + .unwrap_err(); + assert!(matches!( + error, + EmaError::InvalidResponse { stage: Idp, .. } + )); +} + +fn bearer() -> Value { + json!({"access_token":"resource-token","token_type":"Bearer","expires_in":300}) +} + +async fn redeem(http: &MockHttp) -> Result { + exchange( + &MockHttp::new(200, jag(&claims())), + "files.read files.write", + ) + .await? + .exchange(http) + .await +} + +#[tokio::test] +async fn full_exchange_uses_separate_clients_and_only_the_narrowed_assertion() { + for valid in [true, false] { + let mut claims = claims(); + if !valid { + claims["client_id"] = json!("other"); + } + let response = jag(&claims); + let idp = MockHttp::new(200, response.clone()); + let resource_as = MockHttp::new(200, bearer()); + let refresh = RefreshToken::new("refresh-token".into()); + let request = EmaExchangeRequest::new( + EmaAuthorizationServer::new(IDP, IDP_TOKEN, "idp"), + EmaAuthorizationServer::new(AS, AS_TOKEN, "mcp"), + RESOURCE, + &refresh, + ) + .with_scopes(["files.read", "files.write"]); + let future = request.exchange(&idp, &resource_as); + fn is_send(_: &T) {} + is_send(&future); + let result = future.await; + assert_eq!(idp.requests.lock().unwrap().len(), 1); + let requests = resource_as.requests.lock().unwrap(); + assert_eq!(requests.len(), usize::from(valid)); + if !valid { + assert!(matches!( + result, + Err(EmaError::InvalidResponse { stage: Idp, .. }) + )); + continue; + } + let token = result.unwrap(); + assert_eq!(token.access_token.secret(), "resource-token"); + assert_eq!(token.expires_in, Some(Duration::from_secs(300))); + assert_eq!(token.scopes, HashSet::from(["files.read".to_owned()])); + assert!(!format!("{token:?}").contains(token.access_token.secret())); + assert_eq!(requests[0].request.uri(), AS_TOKEN); + assert_eq!( + serde_json::to_value(form(&requests[0])).unwrap(), + json!({ + "grant_type":"urn:ietf:params:oauth:grant-type:jwt-bearer", + "assertion":response["access_token"],"client_id":"mcp" + }) + ); + } +} + +#[tokio::test] +async fn expired_grants_are_not_redeemed() { + let mut grant = exchange(&MockHttp::new(200, jag(&claims())), "") + .await + .unwrap(); + grant.expires_at = 0; + let resource_as = MockHttp::default(); + assert!(matches!( + grant.exchange(&resource_as).await, + Err(EmaError::InvalidRequest(_)) + )); + assert!(resource_as.requests.lock().unwrap().is_empty()); +} + +#[tokio::test] +async fn bearer_responses_cannot_change_resource_or_widen_scope() { + let mut cases = vec![ + ("scope", json!("files.read"), true), + ("scope", json!("files.admin"), false), + ("scope", json!("files.read files.write"), false), + ("scope", json!("files.read files.read"), false), + ("resource", json!(RESOURCE), true), + ("resource", json!([RESOURCE]), true), + ("resource", json!([RESOURCE, "other"]), false), + ("resource", json!("https://mcp.example"), false), + ("expires_in", json!(0), false), + ("expires_in", Value::Null, true), + ("refresh_token", json!("unsupported"), false), + ("token_type", json!("N_A"), false), + ("token_type", json!("bearer"), true), + ("access_token", json!(" \t"), false), + ]; + cases.extend( + BAD_SCOPES + .iter() + .map(|scope| ("scope", json!(scope), false)), + ); + for (key, value, valid) in cases { + let mut token = bearer(); + if value.is_null() { + token.as_object_mut().unwrap().remove(key); + } else { + token[key] = value; + } + let result = redeem(&MockHttp::new(200, token)).await; + assert_eq!(result.is_ok(), valid, "{key}"); + if key == "expires_in" && valid { + assert_eq!(result.unwrap().expires_in, None); + } else if !valid { + assert!(matches!( + result, + Err(EmaError::InvalidResponse { + stage: ResourceServer, + .. + }) + )); + } + } +} + +#[tokio::test] +async fn resource_scope_narrowing_is_reported_without_logging_scope_values() { + for read in ["files.read", "private-scope-sentinel"] { + let requested = format!("{read} files.write"); + let mut claims = claims(); + claims["scope"] = json!(requested); + let grant = exchange(&MockHttp::new(200, jag(&claims)), &requested) + .await + .unwrap(); + let mut response = bearer(); + response["scope"] = json!(read); + let token = grant.exchange(&MockHttp::new(200, response)).await.unwrap(); + assert_eq!(token.scopes, HashSet::from([read.to_owned()])); + assert!(!format!("{token:?}").contains(read)); + } +} + +#[tokio::test] +async fn bearer_scope_may_be_omitted_but_not_added_to_an_unscoped_grant() { + for scope in [None, Some("files.read")] { + let mut claims = claims(); + claims.as_object_mut().unwrap().remove("scope"); + let grant = exchange(&MockHttp::new(200, jag(&claims)), "") + .await + .unwrap(); + let mut token = bearer(); + if let Some(scope) = scope { + token["scope"] = json!(scope); + } + let result = grant.exchange(&MockHttp::new(200, token)).await; + assert_eq!(result.is_ok(), scope.is_none()); + if let Ok(token) = result { + assert!(token.scopes.is_empty()); + } + } +} + +#[tokio::test] +async fn resource_errors_are_staged_and_never_reflect_credentials() { + const SECRET: &str = "secret-resource-error-sentinel"; + for (status, code) in [ + (400, "invalid_grant"), + (400, "insufficient_user_authentication"), + (400, "invalid_client"), + (302, SECRET), + (500, SECRET), + ] { + let http = MockHttp::new(status, json!({"error":code,"error_description":SECRET})); + let error = redeem(&http).await.unwrap_err(); + match code { + "invalid_grant" => assert_eq!(error, EmaError::InvalidGrant(ResourceServer)), + "insufficient_user_authentication" => assert_eq!( + error, + EmaError::InsufficientUserAuthentication(ResourceServer) + ), + _ => assert!( + matches!(error, EmaError::OAuthRejected {stage: ResourceServer, status: actual, ..} if actual == status) + ), + } + assert!(!format!("{error:?} {error}").contains(SECRET)); + assert!(std::error::Error::source(&error).is_none()); + assert_eq!(http.requests.lock().unwrap().len(), 1); + } + let mut oversized = bearer(); + oversized["ignored"] = json!("x".repeat(1024 * 1024)); + for (body, adapter_failure) in [ + (json!({"access_token":SECRET,"expires_in":SECRET}), false), + (oversized, false), + (Value::Null, true), + ] { + let http = MockHttp::new(200, body); + if adapter_failure { + *http.response.lock().unwrap() = Some(Err(SECRET.into())); + } + let error = redeem(&http).await.unwrap_err(); + assert!(!format!("{error:?} {error}").contains(SECRET)); + assert!(std::error::Error::source(&error).is_none()); + if adapter_failure { + assert_eq!(error, EmaError::RequestFailed(ResourceServer)); + } else { + assert!(matches!( + error, + EmaError::InvalidResponse { + stage: ResourceServer, + .. + } + )); + } + } +} + +#[tokio::test(start_paused = true)] +async fn both_exchanges_enforce_the_timeout_when_the_adapter_does_not() { + struct PendingHttp; + impl OAuthHttpClient for PendingHttp { + fn execute(&self, _: OAuthHttpRequest) -> OAuthHttpClientFuture<'_> { + Box::pin(std::future::pending()) + } + } + for stage in [Idp, ResourceServer] { + let ready = MockHttp::new(200, jag(&claims())); + let (idp, resource): (&dyn OAuthHttpClient, &dyn OAuthHttpClient) = match stage { + Idp => (&PendingHttp, &ready), + _ => (&ready, &PendingHttp), + }; + let refresh = RefreshToken::new("refresh-token".into()); + let request = EmaExchangeRequest::new( + EmaAuthorizationServer::new(IDP, IDP_TOKEN, "idp"), + EmaAuthorizationServer::new(AS, AS_TOKEN, "mcp"), + RESOURCE, + &refresh, + ); + let start = tokio::time::Instant::now(); + let result = tokio::time::timeout(Duration::from_secs(31), request.exchange(idp, resource)) + .await + .expect("the SDK must enforce its own deadline"); + assert_eq!(result.unwrap_err(), EmaError::RequestFailed(stage)); + assert_eq!(start.elapsed(), Duration::from_secs(30)); + assert_eq!( + ready.requests.lock().unwrap().len(), + usize::from(stage == ResourceServer) + ); + } +} + +#[derive(Clone, Copy)] +enum AssertionBehavior { + Success, + Error, + Empty(&'static str), + Delay(Duration), + Pending, +} + +struct RecordingAssertionProvider { + calls: Mutex>, + behavior: AssertionBehavior, +} + +impl RecordingAssertionProvider { + fn new(behavior: AssertionBehavior) -> Arc { + Arc::new(Self { + calls: Mutex::new(Vec::new()), + behavior, + }) + } +} + +fn client_assertion(client_id: &str, issuer: &str, sequence: usize) -> String { + jwt( + json!({"alg":"ES256","typ":"client-authentication+jwt"}), + &json!({"iss":client_id,"sub":client_id,"aud":issuer, + "exp":4102444800_u64,"jti":format!("client-assertion-{sequence}")}), + ) +} + +#[async_trait::async_trait] +impl EmaClientAssertionProvider for RecordingAssertionProvider { + async fn create_assertion( + &self, + server: &EmaAuthorizationServer, + ) -> Result> { + let sequence = { + let mut calls = self.calls.lock().unwrap(); + calls.push(( + server.issuer.clone(), + server.token_endpoint.clone(), + server.client_id.clone(), + )); + calls.len() + }; + match self.behavior { + AssertionBehavior::Error => return Err("client-assertion-provider-secret".into()), + AssertionBehavior::Empty(value) => return Ok(EmaClientAssertion::new(value.into())), + AssertionBehavior::Delay(delay) => tokio::time::sleep(delay).await, + AssertionBehavior::Pending => std::future::pending::<()>().await, + AssertionBehavior::Success => {} + } + Ok(EmaClientAssertion::new(client_assertion( + &server.client_id, + &server.issuer, + sequence, + ))) + } +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +enum AuthenticationMethod { + Public, + Basic, + Post, + Jwt, +} + +impl AuthenticationMethod { + fn configure( + self, + secret: &str, + provider: &Arc, + ) -> EmaClientAuthentication { + match self { + Self::Public => EmaClientAuthentication::None, + Self::Basic => { + EmaClientAuthentication::ClientSecretBasic(ClientSecret::new(secret.into())) + } + Self::Post => { + EmaClientAuthentication::ClientSecretPost(ClientSecret::new(secret.into())) + } + Self::Jwt => EmaClientAuthentication::JwtAssertion(provider.clone()), + } + } +} + +fn remove_client_authentication( + request: &OAuthHttpRequest, + method: AuthenticationMethod, + client_id: &str, + secret: &str, + encoded_basic: &str, + assertion: &str, +) -> BTreeMap { + let mut fields = authenticated_form(request); + let authorization = request.request.headers().get("authorization"); + if method == AuthenticationMethod::Basic { + let authorization = authorization.expect("Basic authentication must use a header"); + assert!(authorization.is_sensitive()); + let encoded = authorization + .to_str() + .unwrap() + .strip_prefix("Basic ") + .unwrap(); + assert_eq!(STANDARD.decode(encoded).unwrap(), encoded_basic.as_bytes()); + assert!(!fields.contains_key("client_id")); + } else { + assert!(authorization.is_none()); + assert_eq!(fields.remove("client_id").as_deref(), Some(client_id)); + } + assert_eq!( + fields.remove("client_secret").as_deref(), + (method == AuthenticationMethod::Post).then_some(secret) + ); + assert_eq!( + fields.remove("client_assertion_type").as_deref(), + (method == AuthenticationMethod::Jwt) + .then_some("urn:ietf:params:oauth:client-assertion-type:jwt-bearer") + ); + assert_eq!( + fields.remove("client_assertion").as_deref(), + (method == AuthenticationMethod::Jwt).then_some(assertion) + ); + fields +} + +#[tokio::test] +async fn client_authentication_is_endpoint_specific_and_never_mixed_with_grants() { + const METHODS: [AuthenticationMethod; 4] = [ + AuthenticationMethod::Public, + AuthenticationMethod::Basic, + AuthenticationMethod::Post, + AuthenticationMethod::Jwt, + ]; + // Colons, plus signs, spaces, percent signs, and non-ASCII bytes must be + // form-encoded individually before the HTTP Basic username/password join. + const IDP_CLIENT: &str = "idp:+ %é"; + const AS_CLIENT: &str = "mcp:+ %é"; + const IDP_SECRET: &str = "idp-secret:+ %é"; + const AS_SECRET: &str = "as-secret:+ %é"; + for idp_method in METHODS { + for resource_method in METHODS { + let provider = RecordingAssertionProvider::new(AssertionBehavior::Success); + let mut claims = claims(); + claims["client_id"] = json!(AS_CLIENT); + let response = jag(&claims); + let idp = MockHttp::new(200, response.clone()); + let resource = MockHttp::new(200, bearer()); + let refresh = RefreshToken::new("refresh-token".into()); + let token = EmaExchangeRequest::new( + EmaAuthorizationServer::new(IDP, IDP_TOKEN, IDP_CLIENT) + .with_client_authentication(idp_method.configure(IDP_SECRET, &provider)), + EmaAuthorizationServer::new(AS, AS_TOKEN, AS_CLIENT) + .with_client_authentication(resource_method.configure(AS_SECRET, &provider)), + RESOURCE, + &refresh, + ) + .exchange(&idp, &resource) + .await + .unwrap(); + assert_eq!(token.access_token.secret(), "resource-token"); + let idp_requests = idp.requests.lock().unwrap(); + let resource_requests = resource.requests.lock().unwrap(); + assert_eq!(idp_requests.len(), 1); + assert_eq!(resource_requests.len(), 1); + assert_eq!(idp_requests[0].request.uri(), IDP_TOKEN); + assert_eq!(resource_requests[0].request.uri(), AS_TOKEN); + let idp_fields = remove_client_authentication( + &idp_requests[0], + idp_method, + IDP_CLIENT, + IDP_SECRET, + "idp%3A%2B+%25%C3%A9:idp-secret%3A%2B+%25%C3%A9", + &client_assertion(IDP_CLIENT, IDP, 1), + ); + assert_eq!( + serde_json::to_value(idp_fields).unwrap(), + json!({ + "grant_type":"urn:ietf:params:oauth:grant-type:token-exchange", + "requested_token_type":ID_JAG_TOKEN_TYPE,"subject_token":"refresh-token", + "subject_token_type":"urn:ietf:params:oauth:token-type:refresh_token", + "audience":AS,"resource":RESOURCE + }) + ); + let resource_fields = remove_client_authentication( + &resource_requests[0], + resource_method, + AS_CLIENT, + AS_SECRET, + "mcp%3A%2B+%25%C3%A9:as-secret%3A%2B+%25%C3%A9", + &client_assertion( + AS_CLIENT, + AS, + 1 + usize::from(idp_method == AuthenticationMethod::Jwt), + ), + ); + assert_eq!( + serde_json::to_value(resource_fields).unwrap(), + json!({ + "grant_type":"urn:ietf:params:oauth:grant-type:jwt-bearer", + "assertion":response["access_token"] + }) + ); + let mut expected_calls = Vec::new(); + if idp_method == AuthenticationMethod::Jwt { + expected_calls.push((IDP.into(), IDP_TOKEN.into(), IDP_CLIENT.into())); + } + if resource_method == AuthenticationMethod::Jwt { + expected_calls.push((AS.into(), AS_TOKEN.into(), AS_CLIENT.into())); + } + assert_eq!(*provider.calls.lock().unwrap(), expected_calls); + } + } +} + +#[tokio::test(start_paused = true)] +async fn client_assertions_are_fresh_for_every_request_and_delayed_redemption() { + let provider = RecordingAssertionProvider::new(AssertionBehavior::Success); + let idp_server = EmaAuthorizationServer::new(IDP, IDP_TOKEN, "idp") + .with_client_authentication(EmaClientAuthentication::JwtAssertion(provider.clone())); + let resource_server = EmaAuthorizationServer::new(AS, AS_TOKEN, "mcp") + .with_client_authentication(EmaClientAuthentication::JwtAssertion(provider.clone())); + let refresh = RefreshToken::new("refresh-token".into()); + let mut assertions = HashSet::new(); + for attempt in 0..2 { + let idp = MockHttp::new(200, jag(&claims())); + let resource = MockHttp::new(200, bearer()); + let grant = EmaExchangeRequest::new( + idp_server.clone(), + resource_server.clone(), + RESOURCE, + &refresh, + ) + .exchange_id_jag(&idp) + .await + .unwrap(); + assert_eq!(provider.calls.lock().unwrap().len(), attempt * 2 + 1); + tokio::time::sleep(Duration::from_secs(60)).await; + assert_eq!(provider.calls.lock().unwrap().len(), attempt * 2 + 1); + grant.exchange(&resource).await.unwrap(); + assert_eq!(provider.calls.lock().unwrap().len(), attempt * 2 + 2); + for http in [&idp, &resource] { + let requests = http.requests.lock().unwrap(); + let mut fields = form(&requests[0]); + assert!(assertions.insert(fields.remove("client_assertion").unwrap())); + } + } + assert_eq!(assertions.len(), 4); +} + +#[tokio::test] +async fn invalid_static_client_secrets_fail_before_any_http_or_signing() { + for stage in [Idp, ResourceServer] { + for method in [AuthenticationMethod::Basic, AuthenticationMethod::Post] { + for secret in ["", " \t"] { + let provider = RecordingAssertionProvider::new(AssertionBehavior::Success); + let valid = EmaClientAuthentication::JwtAssertion(provider.clone()); + let invalid = method.configure(secret, &provider); + let (idp_auth, resource_auth) = match stage { + Idp => (invalid, valid), + _ => (valid, invalid), + }; + let idp = MockHttp::default(); + let resource = MockHttp::default(); + let error = EmaExchangeRequest::new( + EmaAuthorizationServer::new(IDP, IDP_TOKEN, "idp") + .with_client_authentication(idp_auth), + EmaAuthorizationServer::new(AS, AS_TOKEN, "mcp") + .with_client_authentication(resource_auth), + RESOURCE, + &RefreshToken::new("refresh-token".into()), + ) + .exchange(&idp, &resource) + .await + .unwrap_err(); + assert!(matches!(error, EmaError::InvalidRequest(_))); + assert!(idp.requests.lock().unwrap().is_empty()); + assert!(resource.requests.lock().unwrap().is_empty()); + assert!(provider.calls.lock().unwrap().is_empty()); + } + } + } +} + +#[tokio::test] +async fn a_grant_that_expires_while_signing_is_never_sent_to_the_resource_server() { + use std::sync::atomic::{AtomicU64, Ordering}; + + let provider = RecordingAssertionProvider::new(AssertionBehavior::Success); + let idp = MockHttp::new(200, jag(&claims())); + let resource = MockHttp::default(); + let mut grant = EmaExchangeRequest::new( + EmaAuthorizationServer::new(IDP, IDP_TOKEN, "idp"), + EmaAuthorizationServer::new(AS, AS_TOKEN, "mcp") + .with_client_authentication(EmaClientAuthentication::JwtAssertion(provider.clone())), + RESOURCE, + &RefreshToken::new("refresh-token".into()), + ) + .exchange_id_jag(&idp) + .await + .unwrap(); + grant.expires_at = 100; + // The clock crosses expiration between the checks before and after signing. + let now = AtomicU64::new(99); + let error = grant + .exchange_with_clock(&resource, || Ok(now.fetch_add(1, Ordering::Relaxed))) + .await + .unwrap_err(); + assert!(matches!(error, EmaError::InvalidRequest(_))); + assert_eq!(provider.calls.lock().unwrap().len(), 1); + assert!(resource.requests.lock().unwrap().is_empty()); +} + +#[tokio::test] +async fn assertion_provider_failures_and_empty_results_never_reach_http_or_escape_errors() { + for stage in [Idp, ResourceServer] { + for behavior in [ + AssertionBehavior::Error, + AssertionBehavior::Empty(""), + AssertionBehavior::Empty(" \t"), + ] { + let provider = RecordingAssertionProvider::new(behavior); + let auth = EmaClientAuthentication::JwtAssertion(provider.clone()); + let (idp_auth, resource_auth) = match stage { + Idp => (auth, EmaClientAuthentication::None), + _ => (EmaClientAuthentication::None, auth), + }; + let idp = MockHttp::new(200, jag(&claims())); + let resource = MockHttp::default(); + let error = EmaExchangeRequest::new( + EmaAuthorizationServer::new(IDP, IDP_TOKEN, "idp") + .with_client_authentication(idp_auth), + EmaAuthorizationServer::new(AS, AS_TOKEN, "mcp") + .with_client_authentication(resource_auth), + RESOURCE, + &RefreshToken::new("refresh-token".into()), + ) + .exchange(&idp, &resource) + .await + .unwrap_err(); + if matches!(behavior, AssertionBehavior::Error) { + assert_eq!(error, EmaError::RequestFailed(stage)); + } else { + assert!(matches!(error, EmaError::InvalidRequest(_))); + } + assert!(!format!("{error:?} {error}").contains("client-assertion-provider-secret")); + assert!(std::error::Error::source(&error).is_none()); + assert_eq!(provider.calls.lock().unwrap().len(), 1); + assert_eq!( + idp.requests.lock().unwrap().len(), + usize::from(stage == ResourceServer) + ); + assert!(resource.requests.lock().unwrap().is_empty()); + } + } +} + +#[tokio::test(start_paused = true)] +async fn signing_and_http_share_one_deadline_at_both_endpoints() { + struct DelayedHttp(Mutex>); + impl OAuthHttpClient for DelayedHttp { + fn execute(&self, request: OAuthHttpRequest) -> OAuthHttpClientFuture<'_> { + self.0.lock().unwrap().push(request); + Box::pin(async { + tokio::time::sleep(Duration::from_secs(10)).await; + Err("http-adapter-secret".into()) + }) + } + } + for stage in [Idp, ResourceServer] { + for behavior in [ + AssertionBehavior::Pending, + AssertionBehavior::Delay(Duration::from_secs(25)), + ] { + let provider = RecordingAssertionProvider::new(behavior); + let auth = EmaClientAuthentication::JwtAssertion(provider.clone()); + let (idp_auth, resource_auth) = match stage { + Idp => (auth, EmaClientAuthentication::None), + _ => (EmaClientAuthentication::None, auth), + }; + let ready = MockHttp::new(200, jag(&claims())); + let delayed = DelayedHttp(Mutex::new(Vec::new())); + let (idp, resource): (&dyn OAuthHttpClient, &dyn OAuthHttpClient) = match stage { + Idp => (&delayed, &ready), + _ => (&ready, &delayed), + }; + let refresh = RefreshToken::new("refresh-token".into()); + let request = EmaExchangeRequest::new( + EmaAuthorizationServer::new(IDP, IDP_TOKEN, "idp") + .with_client_authentication(idp_auth), + EmaAuthorizationServer::new(AS, AS_TOKEN, "mcp") + .with_client_authentication(resource_auth), + RESOURCE, + &refresh, + ); + let start = tokio::time::Instant::now(); + let result = + tokio::time::timeout(Duration::from_secs(31), request.exchange(idp, resource)) + .await + .expect("signing must share the SDK's token request deadline"); + assert_eq!(result.unwrap_err(), EmaError::RequestFailed(stage)); + assert_eq!(start.elapsed(), Duration::from_secs(30)); + assert_eq!(provider.calls.lock().unwrap().len(), 1); + assert_eq!( + delayed.0.lock().unwrap().len(), + usize::from(matches!(behavior, AssertionBehavior::Delay(_))) + ); + assert_eq!( + ready.requests.lock().unwrap().len(), + usize::from(stage == ResourceServer) + ); + } + } +} + +#[tokio::test] +async fn rejected_client_authentication_never_falls_back_or_retries() { + for stage in [Idp, ResourceServer] { + for method in [ + AuthenticationMethod::Basic, + AuthenticationMethod::Post, + AuthenticationMethod::Jwt, + ] { + let provider = RecordingAssertionProvider::new(AssertionBehavior::Success); + let failure = json!({"error":"invalid_client","error_description":"client-authentication-secret"}); + let idp = if stage == Idp { + MockHttp::new(401, failure.clone()) + } else { + MockHttp::new(200, jag(&claims())) + }; + let resource = MockHttp::new(401, failure); + let error = EmaExchangeRequest::new( + EmaAuthorizationServer::new(IDP, IDP_TOKEN, "idp") + .with_client_authentication(method.configure("idp-client-secret", &provider)), + EmaAuthorizationServer::new(AS, AS_TOKEN, "mcp") + .with_client_authentication(method.configure("as-client-secret", &provider)), + RESOURCE, + &RefreshToken::new("refresh-token".into()), + ) + .exchange(&idp, &resource) + .await + .unwrap_err(); + assert_eq!( + error, + EmaError::OAuthRejected { + stage, + status: 401, + code: "invalid_client" + } + ); + assert!(!format!("{error:?} {error}").contains("client-authentication-secret")); + assert_eq!(idp.requests.lock().unwrap().len(), 1); + assert_eq!( + resource.requests.lock().unwrap().len(), + usize::from(stage == ResourceServer) + ); + let expected_calls = if method == AuthenticationMethod::Jwt { + 1 + usize::from(stage == ResourceServer) + } else { + 0 + }; + assert_eq!(provider.calls.lock().unwrap().len(), expected_calls); + } + } +} + +#[test] +fn client_authentication_debug_output_redacts_credentials_and_provider_state() { + const SECRET: &str = "client-authentication-secret-sentinel"; + let provider = RecordingAssertionProvider::new(AssertionBehavior::Empty(SECRET)); + let assertion = EmaClientAssertion::new(SECRET.into()); + assert!(!format!("{assertion:?}").contains(SECRET)); + for authentication in [ + EmaClientAuthentication::ClientSecretBasic(ClientSecret::new(SECRET.into())), + EmaClientAuthentication::ClientSecretPost(ClientSecret::new(SECRET.into())), + EmaClientAuthentication::JwtAssertion(provider.clone()), + ] { + assert!(!format!("{authentication:?}").contains(SECRET)); + let server = EmaAuthorizationServer::new(IDP, IDP_TOKEN, "idp") + .with_client_authentication(authentication); + let refresh = RefreshToken::new("refresh-token".into()); + let request = EmaExchangeRequest::new(server.clone(), server.clone(), RESOURCE, &refresh); + assert!(!format!("{server:?} {request:?}").contains(SECRET)); + assert!(!format!("{server:?} {request:?}").contains("private-query")); + } +} diff --git a/docs/OAUTH_SUPPORT.md b/docs/OAUTH_SUPPORT.md index 82a1fe918..5d32be348 100644 --- a/docs/OAUTH_SUPPORT.md +++ b/docs/OAUTH_SUPPORT.md @@ -15,6 +15,7 @@ This document describes the OAuth 2.1 authorization implementation for Model Con - Automatic token refresh - Authorized HTTP Client implementation - Injectable OAuth HTTP client for custom network environments +- Opt-in EMA/XAA refresh-token and ID-JAG exchanges for registered public and confidential clients ## Usage Guide @@ -294,6 +295,130 @@ match oauth_state.request_scope_upgrade("admin:write", MCP_REDIRECT_URI).await { } ``` +## Enterprise-managed authorization (EMA/XAA) + +The example requires the `rmcp` features `auth-enterprise-managed`, `client`, +`reqwest` (TLS), and `transport-streamable-http-client-reqwest`, plus `oauth2` +version 5. Call the async function from a Tokio runtime. + +The exchange profile has these requirements and limits: + +- Each authorization server has its own approved client registration and explicit + `EmaClientAuthentication`: `None`, `ClientSecretBasic`, `ClientSecretPost`, or + `JwtAssertion`. The SDK does not select methods from metadata or fall back to a + different method after a failure. +- Input is an enterprise IdP refresh token. The requested MCP resource must match + the ID-JAG's sole `resource` claim; scope may be omitted or narrowed. +- RAR (`authorization_details`) and DPoP are not supported. Nonempty authorization + details are rejected at both exchange stages. +- Redemption consumes the SDK's ID-JAG handle and does not retry automatically. + This is an SDK safety choice, not a protocol requirement that ID-JAGs be single-use. + +The [ID-JAG draft recommends confidential clients](https://datatracker.ietf.org/doc/html/draft-ietf-oauth-identity-assertion-authz-grant-04#section-9.1). +The example below uses confidential clients with `client_secret_basic` at both +servers. Use `None` only where that server permits a public client registration. +Discovery, server approval, SSO, credential storage, and reauthentication remain +the application's responsibility. Client-side ID-JAG checks validate structure and +bindings, not signatures; the resource authorization server verifies signatures. + +```rust no_run +use oauth2::{ClientSecret, RefreshToken}; +use rmcp::{ + ServiceExt, + model::ClientInfo, + transport::{ + StreamableHttpClientTransport, + auth::{ + default_oauth_http_client, + enterprise::{EmaAuthorizationServer, EmaClientAuthentication, EmaExchangeRequest}, + }, + streamable_http_client::StreamableHttpClientTransportConfig, + }, +}; + +async fn connect( + refresh: &RefreshToken, + idp_client_secret: ClientSecret, + resource_client_secret: ClientSecret, +) -> Result<(), Box> { + let resource = "https://mcp.example/mcp"; + let http = default_oauth_http_client()?; + let idp = EmaAuthorizationServer::new( + "https://idp.example", "https://idp.example/token", "idp-client", + ) + .with_client_authentication(EmaClientAuthentication::ClientSecretBasic(idp_client_secret)); + let resource_as = EmaAuthorizationServer::new( + "https://as.example", "https://as.example/token", "mcp-client", + ) + .with_client_authentication(EmaClientAuthentication::ClientSecretBasic(resource_client_secret)); + let token = EmaExchangeRequest::new(idp, resource_as, resource, refresh) + .with_scopes(["files.read"]) + .exchange(&http, &http) + .await?; + let transport = StreamableHttpClientTransport::from_config( + StreamableHttpClientTransportConfig::with_uri(resource) + .auth_header(token.access_token.secret()), + ); + let client = ClientInfo::default().serve(transport).await?; + client.list_tools(Default::default()).await?; + client.cancel().await?; + Ok(()) +} +``` + +`auth_header` takes the token without a `Bearer ` prefix. Use it only with the +approved resource and never log it. This transport uses a fixed token; obtain a +new token and reconnect when it expires or is rejected. + +For a registration using JWT client authentication, implement +`EmaClientAssertionProvider` with your application's signer and configure +`JwtAssertion` on that server. The provider example below also uses `async-trait` +version 0.1; `AppSigner` represents your application's existing signing service. + +```rust ignore +use std::sync::Arc; +use rmcp::transport::auth::enterprise::{ + EmaAuthorizationServer, EmaClientAssertion, EmaClientAssertionProvider, + EmaClientAuthentication, +}; + +struct AppAssertionProvider { + signer: AppSigner, +} + +#[async_trait::async_trait] +impl EmaClientAssertionProvider for AppAssertionProvider { + async fn create_assertion( + &self, + server: &EmaAuthorizationServer, + ) -> Result> { + // Sign a new assertion for this registration and approved server. + let jwt = self.signer.sign_client_assertion( + &server.client_id, &server.issuer, &server.token_endpoint, + ).await?; + Ok(EmaClientAssertion::new(jwt)) + } +} + +let resource_as = EmaAuthorizationServer::new( + "https://as.example", "https://as.example/token", "mcp-client", +) +.with_client_authentication(EmaClientAuthentication::JwtAssertion(Arc::new( + AppAssertionProvider { signer }, +))); +``` + +The SDK calls the provider before each token request, including delayed +`EmaIdJag::exchange` redemption. Sign a fresh, short-lived assertion with a unique +`jti`, the registered client ID in `iss` and `sub`, and the server's approved +audience in `aud`. Your signer owns the keys and algorithm; the client assertion +is separate from the ID-JAG grant. Signing and HTTP share a 30-second deadline. + +The factory honors per-request redirect policy with the SDK's default reqwest +settings. For custom proxy, CA, or remote-execution policy, implement +`OAuthHttpClient`; use separate adapters for the IdP and resource AS when their +network policies differ. + ## Complete Examples - **Authorization Code client**: [`examples/clients/src/auth/oauth_client.rs`](../examples/clients/src/auth/oauth_client.rs) From add9cbeab56725547e7ea3bf3fb78f59de0806f1 Mon Sep 17 00:00:00 2001 From: Dale Seo <5466341+DaleSeo@users.noreply.github.com> Date: Tue, 8 Sep 2026 10:38:30 -0400 Subject: [PATCH 18/29] ci: bump codeql-action to 4.37.9 and group updates (#1240) --- .github/dependabot.yml | 4 ++++ .github/workflows/codeql.yml | 6 +++--- 2 files changed, 7 insertions(+), 3 deletions(-) diff --git a/.github/dependabot.yml b/.github/dependabot.yml index 0fa34ce69..168b680a6 100644 --- a/.github/dependabot.yml +++ b/.github/dependabot.yml @@ -21,6 +21,10 @@ updates: # Mark PRs as CI related change. - T-CI open-pull-requests-limit: 3 + groups: + codeql-action: + patterns: + - "github/codeql-action*" commit-message: prefix: "chore" include: "scope" diff --git a/.github/workflows/codeql.yml b/.github/workflows/codeql.yml index 4ecf3484b..83e1ce4d8 100644 --- a/.github/workflows/codeql.yml +++ b/.github/workflows/codeql.yml @@ -27,13 +27,13 @@ jobs: uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 - name: Initialize CodeQL - uses: github/codeql-action/init@ff2f1c621b7f889edc0d3c761ac2e6a3f8cdb0dd # v4.37.7 + uses: github/codeql-action/init@cdf488f595d80d6e07e03d4674febd5ab45fa938 # v4.37.9 with: languages: ${{ matrix.language }} config-file: ./.github/codeql/codeql-config.yml - name: Autobuild - uses: github/codeql-action/autobuild@ff2f1c621b7f889edc0d3c761ac2e6a3f8cdb0dd # v4.37.7 + uses: github/codeql-action/autobuild@cdf488f595d80d6e07e03d4674febd5ab45fa938 # v4.37.9 - name: Perform CodeQL Analysis - uses: github/codeql-action/analyze@ff2f1c621b7f889edc0d3c761ac2e6a3f8cdb0dd # v4.37.7 + uses: github/codeql-action/analyze@cdf488f595d80d6e07e03d4674febd5ab45fa938 # v4.37.9 From 04c2f3fac90dbd2ec4ae0fbd9082fababd32dab9 Mon Sep 17 00:00:00 2001 From: Dale Seo <5466341+DaleSeo@users.noreply.github.com> Date: Tue, 8 Sep 2026 19:04:15 -0400 Subject: [PATCH 19/29] fix(auth): unify refresh checks and error handling (#1236) * fix(auth): unify refresh checks and error handling Claude-Session: https://claude.ai/code/session_011pHFfoTygeG84mCDXzCcmw * test(auth): add live authorization-server checks Claude-Session: https://claude.ai/code/session_011pHFfoTygeG84mCDXzCcmw --- crates/rmcp/src/transport/auth.rs | 118 +++++++-- .../common/auth/streamable_http_client.rs | 118 ++++++++- crates/rmcp/tests/test_live_oauth_refresh.rs | 233 ++++++++++++++++++ scripts/keycloak-oauth-fixture.sh | 56 +++++ 4 files changed, 502 insertions(+), 23 deletions(-) create mode 100644 crates/rmcp/tests/test_live_oauth_refresh.rs create mode 100755 scripts/keycloak-oauth-fixture.sh diff --git a/crates/rmcp/src/transport/auth.rs b/crates/rmcp/src/transport/auth.rs index 8c9e41216..759044b0c 100644 --- a/crates/rmcp/src/transport/auth.rs +++ b/crates/rmcp/src/transport/auth.rs @@ -294,6 +294,7 @@ impl CredentialRefreshGuard { /// Implementations of this trait can provide custom storage backends /// for OAuth2 credentials, such as file-based storage, keychain integration, /// or database storage. +/// /// Return [`AuthError::CredentialStoreError`] for backend or locking failures /// so they remain distinct from errors requiring reauthorization. #[async_trait] @@ -2248,12 +2249,19 @@ impl AuthorizationManager { .as_ref() .ok_or_else(|| AuthError::InternalError("OAuth client not configured".to_string()))?; - let refresh_guard = self.credential_store.acquire_refresh_guard().await?; + // Held for the rest of this function so the load, the exchange, and the + // save stay inside one guarded section. + let _refresh_guard = self.credential_store.acquire_refresh_guard().await?; let stored = self.credential_store.load().await?; let stored_credentials = stored.ok_or(AuthError::AuthorizationRequired)?; - if refresh_guard.is_some() - && stored_credentials.client_id != oauth_client.client_id().as_str() - { + // Refreshing with another client's stored token would put that token on a + // request authenticated as this client. + if stored_credentials.client_id != oauth_client.client_id().as_str() { + tracing::warn!( + stored_client_id = stored_credentials.client_id.as_str(), + configured_client_id = oauth_client.client_id().as_str(), + "stored credentials belong to a different client; reauthorization required" + ); return Err(AuthError::AuthorizationRequired); } let current_credentials = stored_credentials @@ -2271,8 +2279,8 @@ impl AuthorizationManager { // RFC 8707: the resource indicator is required on token requests, including refreshes .add_extra_param("resource", self.oauth_resource().await); let mut refresh_scopes = stored_credentials.granted_scopes; - let authoritative_scopes = refresh_guard.is_some().then(|| refresh_scopes.clone()); self.add_offline_access_if_supported(&mut refresh_scopes); + let requested_scopes = refresh_scopes.clone(); for scope in refresh_scopes { refresh_request = refresh_request.add_scope(Scope::new(scope)); } @@ -2298,10 +2306,12 @@ impl AuthorizationManager { token_result.set_refresh_token(Some(refresh_token_value)); } - let granted_scopes: Vec = match (token_result.scopes(), authoritative_scopes) { - (Some(scopes), _) => scopes.iter().map(|s| s.to_string()).collect(), - (None, Some(scopes)) => scopes, - (None, None) => self.current_scopes.read().await.clone(), + let response_scopes = token_result + .scopes() + .map(|scopes| scopes.iter().map(|s| s.to_string()).collect()); + let granted_scopes = { + let current = self.current_scopes.read().await; + Self::resolve_granted_scopes(response_scopes, &requested_scopes, ¤t) }; *self.current_scopes.write().await = granted_scopes.clone(); @@ -8515,17 +8525,17 @@ mod tests { } } - fn refresh_store() -> RefreshStore { + async fn refresh_store() -> RefreshStore { let credentials = StoredCredentials::new( "my-client".into(), Some(make_token_response_with_refresh("old-token", "old-refresh")), vec!["read".into()], Some(AuthorizationManager::now_epoch_secs()), ); + let credential_store = InMemoryCredentialStore::new(); + credential_store.save(credentials).await.unwrap(); RefreshStore { - credentials: InMemoryCredentialStore { - credentials: Arc::new(tokio::sync::RwLock::new(Some(credentials))), - }, + credentials: credential_store, lock: Arc::new(Mutex::new(())), events: Arc::new(StdMutex::new(Vec::new())), guard_requested: Arc::new(Semaphore::new(0)), @@ -8604,7 +8614,7 @@ mod tests { #[tokio::test] async fn refresh_guard_spans_load_exchange_and_completed_save() { - let store = refresh_store(); + let store = refresh_store().await; let manager = refresh_manager(store.clone(), refresh_http_client(&store)).await; manager.refresh_token().await.unwrap(); @@ -8631,7 +8641,7 @@ mod tests { #[tokio::test] async fn concurrent_refreshes_wait_for_save_and_use_the_latest_token() { - let mut store = refresh_store(); + let mut store = refresh_store().await; let save_gate = Arc::new(Semaphore::new(0)); store.save_gate = Some(save_gate.clone()); let http_client = refresh_http_client(&store); @@ -8677,8 +8687,8 @@ mod tests { } #[tokio::test] - async fn guarded_refresh_rejects_credentials_for_another_client() { - let store = refresh_store(); + async fn refresh_rejects_credentials_for_another_client() { + let store = refresh_store().await; let mut credentials = store.credentials.load().await.unwrap().unwrap(); credentials.client_id = "other-client".into(); store.credentials.save(credentials).await.unwrap(); @@ -8693,6 +8703,78 @@ mod tests { assert!(store.lock.try_lock().is_ok()); } + #[tokio::test] + async fn refresh_rejects_credentials_for_another_client_without_a_guard() { + let (base_url, captured) = start_token_server().await; + let mut manager = manager_with_metadata(Some(AuthorizationMetadata { + authorization_endpoint: format!("{base_url}/authorize"), + token_endpoint: format!("{base_url}/token"), + ..Default::default() + })) + .await; + manager.configure_client(test_client_config()).unwrap(); + manager + .credential_store + .save(StoredCredentials::new( + "other-client".into(), + Some(make_token_response_with_refresh("old-token", "old-refresh")), + vec!["read".into()], + Some(AuthorizationManager::now_epoch_secs()), + )) + .await + .unwrap(); + + let error = manager.refresh_token().await.unwrap_err(); + + assert!( + matches!(error, AuthError::AuthorizationRequired), + "a client mismatch must require reauthorization, got: {error:?}" + ); + assert!( + captured.lock().unwrap().is_none(), + "a client mismatch must be caught before the refresh token leaves the process" + ); + } + + #[tokio::test] + async fn refresh_without_a_guard_keeps_stored_scopes_when_response_omits_them() { + // start_token_server answers without a `scope`, matching a provider that + // grants the request in full. + let (base_url, _captured) = start_token_server().await; + let mut manager = manager_with_metadata(Some(AuthorizationMetadata { + authorization_endpoint: format!("{base_url}/authorize"), + token_endpoint: format!("{base_url}/token"), + ..Default::default() + })) + .await; + manager.configure_client(test_client_config()).unwrap(); + manager + .credential_store + .save(StoredCredentials::new( + "my-client".into(), + Some(make_token_response_with_refresh("old-token", "old-refresh")), + vec!["read".into()], + Some(AuthorizationManager::now_epoch_secs()), + )) + .await + .unwrap(); + *manager.current_scopes.write().await = vec!["stale".into()]; + + manager.refresh_token().await.unwrap(); + + let saved = manager.credential_store.load().await.unwrap().unwrap(); + assert_eq!( + saved.granted_scopes, + ["read"], + "the stored grant outranks the per-process scope cache" + ); + assert_eq!( + manager.get_current_scopes().await, + ["read"], + "the refreshed grant must replace the stale scope cache" + ); + } + #[rstest] #[case("guard", 0)] #[case("load", 0)] @@ -8702,7 +8784,7 @@ mod tests { #[case] phase: &'static str, #[case] provider_requests: usize, ) { - let mut store = refresh_store(); + let mut store = refresh_store().await; store.fail_at = Some(phase); let http_client = refresh_http_client(&store); let manager = refresh_manager(store.clone(), http_client.clone()).await; diff --git a/crates/rmcp/src/transport/common/auth/streamable_http_client.rs b/crates/rmcp/src/transport/common/auth/streamable_http_client.rs index 97432f90e..47b6caf55 100644 --- a/crates/rmcp/src/transport/common/auth/streamable_http_client.rs +++ b/crates/rmcp/src/transport/common/auth/streamable_http_client.rs @@ -19,7 +19,10 @@ where /// 401 propagates as [`StreamableHttpError::AuthRequired`] carrying the /// `WWW-Authenticate` challenge for the caller to authorize with; /// - a token the server rejects (e.g. revoked) → one silent refresh, one - /// retry, then the challenge propagates. + /// retry, then the challenge propagates; + /// - a refresh that fails for any other reason (credential store, network, + /// provider) → that error propagates so the caller can retry instead of + /// being sent through a new authorization. async fn call_reacting_to_challenges( &self, auth_token: Option, @@ -54,11 +57,13 @@ where match refreshed { Ok(fresh_token) if fresh_token != sent_token => call(Some(fresh_token)).await, Ok(_) => Err(StreamableHttpError::AuthRequired(challenge)), - Err(error @ AuthError::CredentialStoreError(_)) => Err(error.into()), - Err(error) => { - debug!("token refresh after server rejection failed: {error}"); + // `try_refresh_or_reauth` already reports the cases that need a + // new authorization; anything else is retryable or infrastructural. + Err(AuthError::AuthorizationRequired) => { + debug!("token refresh after server rejection requires authorization"); Err(StreamableHttpError::AuthRequired(challenge)) } + Err(error) => Err(error.into()), } } result => result, @@ -212,11 +217,16 @@ where #[cfg(all(test, feature = "transport-streamable-http-client-reqwest"))] mod tests { + use std::sync::Arc; + + use oauth2::{AccessToken, RefreshToken, basic::BasicTokenType}; + use super::*; use crate::transport::{ auth::{ AuthorizationManager, AuthorizationMetadata, CredentialRefreshGuard, CredentialStore, - StoredCredentials, + InMemoryCredentialStore, OAuthHttpClient, OAuthHttpClientFuture, OAuthHttpRequest, + OAuthTokenResponse, StoredCredentials, VendorExtraTokenFields, }, streamable_http_client::AuthRequiredError, }; @@ -269,4 +279,102 @@ mod tests { StreamableHttpError::Auth(AuthError::CredentialStoreError(message)) if message == "guard unavailable")); } + + struct UnreachableTokenEndpoint; + + impl OAuthHttpClient for UnreachableTokenEndpoint { + fn execute(&self, _: OAuthHttpRequest) -> OAuthHttpClientFuture<'_> { + Box::pin(async { Err("token endpoint unreachable".into()) }) + } + } + + struct RejectingTokenEndpoint; + + impl OAuthHttpClient for RejectingTokenEndpoint { + fn execute(&self, _: OAuthHttpRequest) -> OAuthHttpClientFuture<'_> { + Box::pin(async { + Ok(oauth2::http::Response::builder() + .status(400) + .header("content-type", "application/json") + .body(br#"{"error":"invalid_grant"}"#.to_vec()) + .unwrap()) + }) + } + } + + /// A manager holding a refresh token the given token endpoint will answer for. + async fn manager_with_stored_refresh_token( + token_endpoint: Arc, + ) -> AuthorizationManager { + let mut manager = AuthorizationManager::new_with_oauth_http_client( + "https://mcp.example.com/mcp", + token_endpoint, + ) + .await + .unwrap(); + manager.set_metadata(AuthorizationMetadata { + authorization_endpoint: "https://auth.example.com/authorize".into(), + token_endpoint: "https://auth.example.com/token".into(), + ..Default::default() + }); + manager.configure_client_id("client").unwrap(); + + let mut token_response = OAuthTokenResponse::new( + AccessToken::new("old-token".into()), + BasicTokenType::Bearer, + VendorExtraTokenFields::default(), + ); + token_response.set_refresh_token(Some(RefreshToken::new("stored-refresh".into()))); + let store = InMemoryCredentialStore::new(); + store + .save(StoredCredentials::new( + "client".into(), + Some(token_response), + vec![], + None, + )) + .await + .unwrap(); + manager.set_credential_store(store); + manager + } + + /// Drive one call whose server answer is a 401 challenge. + async fn challenge_once(manager: AuthorizationManager) -> StreamableHttpError { + AuthClient::new(reqwest::Client::new(), manager) + .call_reacting_to_challenges(Some("old-token".into()), |_| async { + Err::<(), _>(StreamableHttpError::AuthRequired(AuthRequiredError::new( + "Bearer".into(), + ))) + }) + .await + .unwrap_err() + } + + #[tokio::test] + async fn reactive_refresh_propagates_retryable_refresh_failure() { + let manager = manager_with_stored_refresh_token(Arc::new(UnreachableTokenEndpoint)).await; + + let error = challenge_once(manager).await; + + assert!( + matches!( + error, + StreamableHttpError::Auth(AuthError::TokenRefreshFailed(_)) + ), + "a retryable refresh failure must reach the caller instead of asking for a new authorization, got: {error:?}" + ); + } + + #[tokio::test] + async fn reactive_refresh_reports_a_rejected_refresh_token_as_a_challenge() { + let manager = manager_with_stored_refresh_token(Arc::new(RejectingTokenEndpoint)).await; + + let error = challenge_once(manager).await; + + assert!( + matches!(error, StreamableHttpError::AuthRequired(_)), + "a definitively rejected refresh token must surface the challenge, got: {error:?}" + ); + } } diff --git a/crates/rmcp/tests/test_live_oauth_refresh.rs b/crates/rmcp/tests/test_live_oauth_refresh.rs new file mode 100644 index 000000000..543abe537 --- /dev/null +++ b/crates/rmcp/tests/test_live_oauth_refresh.rs @@ -0,0 +1,233 @@ +//! Live authorization-server checks for the refresh path. +//! +//! These are `#[ignore]`d, so they never run in CI. Provision the authorization +//! server with `./scripts/keycloak-oauth-fixture.sh`, then run them with +//! `cargo test -p rmcp --all-features --test test_live_oauth_refresh -- --ignored`. +//! `KC_BASE` overrides the server location for both the script and these tests. +#![cfg(feature = "auth")] + +use std::sync::Arc; + +use rmcp::transport::auth::{ + AuthError, AuthorizationManager, AuthorizationMetadata, CredentialRefreshGuard, + CredentialStore, InMemoryCredentialStore, OAuthClientConfig, OAuthTokenResponse, + StoredCredentials, +}; +use tokio::sync::Mutex; + +const REALM: &str = "rmcp"; +const CLIENT_ID: &str = "rmcp-client"; +const CLIENT_SECRET: &str = "rmcp-secret"; + +fn kc_base() -> String { + std::env::var("KC_BASE").unwrap_or_else(|_| "http://localhost:8081".to_string()) +} + +fn token_endpoint() -> String { + format!("{}/realms/{REALM}/protocol/openid-connect/token", kc_base()) +} + +/// Ask Keycloak for a genuine token pair through the direct access grant, so the +/// stored credentials hold a refresh token the server will actually honor. +async fn issue_real_credentials() -> OAuthTokenResponse { + let form = format!( + "client_id={CLIENT_ID}&client_secret={CLIENT_SECRET}\ + &username=alice&password=alice-pw&grant_type=password&scope=openid+profile" + ); + let body = reqwest::Client::new() + .post(token_endpoint()) + .header("content-type", "application/x-www-form-urlencoded") + .body(form) + .send() + .await + .expect("keycloak unreachable") + .text() + .await + .unwrap(); + serde_json::from_str(&body).unwrap_or_else(|e| panic!("unexpected token response {body}: {e}")) +} + +fn metadata() -> AuthorizationMetadata { + // `AuthorizationMetadata` is `#[non_exhaustive]`, so fill it field by field. + let mut metadata = AuthorizationMetadata::default(); + metadata.authorization_endpoint = + format!("{}/realms/{REALM}/protocol/openid-connect/auth", kc_base()); + metadata.token_endpoint = token_endpoint(); + metadata +} + +fn client_config() -> OAuthClientConfig { + let mut config = OAuthClientConfig::new(CLIENT_ID, "http://localhost/callback"); + config.client_secret = Some(CLIENT_SECRET.to_string()); + config +} + +async fn manager_with_store(store: S) -> AuthorizationManager { + let mut manager = AuthorizationManager::new(kc_base()).await.unwrap(); + manager.set_metadata(metadata()); + manager.configure_client(client_config()).unwrap(); + manager.set_credential_store(store); + manager +} + +fn stored(client_id: &str, token: OAuthTokenResponse) -> StoredCredentials { + StoredCredentials::new( + client_id.to_string(), + Some(token), + vec!["openid".into(), "profile".into()], + None, + ) +} + +/// A store that serializes refreshes the way a shared on-disk store would. +#[derive(Clone, Default)] +struct GuardedStore { + inner: InMemoryCredentialStore, + lock: Arc>, +} + +#[async_trait::async_trait] +impl CredentialStore for GuardedStore { + async fn load(&self) -> Result, AuthError> { + self.inner.load().await + } + + async fn save(&self, credentials: StoredCredentials) -> Result<(), AuthError> { + self.inner.save(credentials).await + } + + async fn clear(&self) -> Result<(), AuthError> { + self.inner.clear().await + } + + async fn acquire_refresh_guard(&self) -> Result, AuthError> { + Ok(Some(CredentialRefreshGuard::new( + self.lock.clone().lock_owned().await, + ))) + } +} + +#[tokio::test] +#[ignore = "requires a live Keycloak"] +async fn live_refresh_rotates_the_stored_token() { + use oauth2::TokenResponse; + + let issued = issue_real_credentials().await; + let original_refresh = issued.refresh_token().unwrap().secret().clone(); + let store = InMemoryCredentialStore::new(); + store.save(stored(CLIENT_ID, issued)).await.unwrap(); + let manager = manager_with_store(store.clone()).await; + + let refreshed = manager.refresh_token().await.expect("live refresh failed"); + + let saved = store.load().await.unwrap().unwrap(); + let saved_token = saved.token_response.unwrap(); + assert_ne!( + saved_token.refresh_token().unwrap().secret(), + &original_refresh, + "keycloak rotates refresh tokens, so the store must hold the new one" + ); + assert_eq!( + saved_token.access_token().secret(), + refreshed.access_token().secret(), + "the saved credentials must match what the caller received" + ); + assert!( + saved.granted_scopes.contains(&"openid".to_string()), + "granted scopes should come from the provider response, got: {:?}", + saved.granted_scopes + ); +} + +#[tokio::test] +#[ignore = "requires a live Keycloak"] +async fn live_refresh_rejects_credentials_for_another_client() { + use oauth2::TokenResponse; + + let issued = issue_real_credentials().await; + let untouched_refresh = issued.refresh_token().unwrap().secret().clone(); + let store = InMemoryCredentialStore::new(); + store + .save(stored("some-other-client", issued)) + .await + .unwrap(); + let manager = manager_with_store(store.clone()).await; + + let error = manager.refresh_token().await.unwrap_err(); + + assert!( + matches!(error, AuthError::AuthorizationRequired), + "a client mismatch must require reauthorization, got: {error:?}" + ); + let saved = store.load().await.unwrap().unwrap(); + assert_eq!( + saved + .token_response + .unwrap() + .refresh_token() + .unwrap() + .secret(), + &untouched_refresh, + "a rejected refresh must leave the stored token untouched" + ); +} + +#[tokio::test] +#[ignore = "requires a live Keycloak"] +async fn live_concurrent_refreshes_survive_refresh_token_rotation() { + use oauth2::TokenResponse; + + let issued = issue_real_credentials().await; + let store = GuardedStore::default(); + store.save(stored(CLIENT_ID, issued)).await.unwrap(); + let first = manager_with_store(store.clone()).await; + let second = manager_with_store(store.clone()).await; + + // Without the guard the second caller would reuse the refresh token the first + // one already consumed, and Keycloak would answer `invalid_grant`. + let (a, b) = tokio::join!( + tokio::spawn(async move { first.refresh_token().await }), + tokio::spawn(async move { second.refresh_token().await }) + ); + + let a = a.unwrap().expect("first concurrent refresh failed"); + let b = b.unwrap().expect("second concurrent refresh failed"); + assert_ne!( + a.access_token().secret(), + b.access_token().secret(), + "each caller performs its own exchange, so the tokens must differ" + ); +} + +/// Deterministic proof that this realm really does invalidate a rotated refresh +/// token, which is what makes the coordination above load-bearing. +#[tokio::test] +#[ignore = "requires a live Keycloak with revokeRefreshToken enabled"] +async fn live_reusing_a_rotated_refresh_token_is_rejected() { + let issued = issue_real_credentials().await; + + let first_store = InMemoryCredentialStore::new(); + first_store + .save(stored(CLIENT_ID, issued.clone())) + .await + .unwrap(); + manager_with_store(first_store) + .await + .refresh_token() + .await + .expect("the first refresh should succeed"); + + // A second caller that never saw the rotation still holds the consumed token. + let stale_store = InMemoryCredentialStore::new(); + stale_store.save(stored(CLIENT_ID, issued)).await.unwrap(); + let error = manager_with_store(stale_store) + .await + .refresh_token() + .await + .unwrap_err(); + + assert!( + matches!(error, AuthError::TokenRefreshRejected(_)), + "reusing a rotated refresh token must be rejected, got: {error:?}" + ); +} diff --git a/scripts/keycloak-oauth-fixture.sh b/scripts/keycloak-oauth-fixture.sh new file mode 100755 index 000000000..50c9b85b8 --- /dev/null +++ b/scripts/keycloak-oauth-fixture.sh @@ -0,0 +1,56 @@ +#!/usr/bin/env bash +# ============================================================================= +# keycloak-oauth-fixture.sh — Live OAuth fixture for rmcp refresh tests +# +# Starts a throwaway Keycloak and provisions the realm that +# crates/rmcp/tests/test_live_oauth_refresh.rs expects. Those tests are +# #[ignore]d, so nothing here runs in CI. +# +# Provision: ./scripts/keycloak-oauth-fixture.sh +# Run tests: cargo test -p rmcp --all-features \ +# --test test_live_oauth_refresh -- --ignored +# Tear down: docker rm -f kc-rmcp-test +# +# Requires Docker. Override the port with KC_BASE (default localhost:8081); +# the tests read the same variable. +# ============================================================================= +set -euo pipefail + +KC=${KC_BASE:-http://localhost:8081} +PORT=${KC##*:} +CONTAINER=kc-rmcp-test + +docker rm -f "$CONTAINER" >/dev/null 2>&1 || true +docker run -d --name "$CONTAINER" -p "$PORT:8080" \ + -e KC_BOOTSTRAP_ADMIN_USERNAME=admin -e KC_BOOTSTRAP_ADMIN_PASSWORD=admin \ + quay.io/keycloak/keycloak:26.0 start-dev >/dev/null + +echo "waiting for keycloak at $KC ..." +until curl -sf -o /dev/null "$KC/realms/master/.well-known/openid-configuration"; do + sleep 3 +done + +token=$(curl -s -X POST "$KC/realms/master/protocol/openid-connect/token" \ + -d client_id=admin-cli -d username=admin -d password=admin -d grant_type=password | + python3 -c 'import sys,json;print(json.load(sys.stdin)["access_token"])') + +provision() { + curl -s -X POST "$KC/admin/realms$1" \ + -H "Authorization: Bearer $token" -H "Content-Type: application/json" \ + -d "$2" -o /dev/null -w " ${1:-/} -> %{http_code}\n" +} + +# revokeRefreshToken makes refresh tokens single-use. The concurrency test +# depends on it: without the refresh guard the second caller replays a consumed +# token and Keycloak answers invalid_grant. +provision "" '{"realm":"rmcp","enabled":true,"revokeRefreshToken":true,"refreshTokenMaxReuse":0}' +provision "/rmcp/clients" '{"clientId":"rmcp-client","secret":"rmcp-secret","publicClient":false, + "directAccessGrantsEnabled":true,"standardFlowEnabled":true, + "redirectUris":["http://localhost/callback"]}' +# The profile fields and empty requiredActions keep the direct access grant from +# failing with "Account is not fully set up". +provision "/rmcp/users" '{"username":"alice","enabled":true,"emailVerified":true, + "email":"alice@example.com","firstName":"Alice","lastName":"Example","requiredActions":[], + "credentials":[{"type":"password","value":"alice-pw","temporary":false}]}' + +echo "ready" From 7d14f3a8a2419dc91dc439f3c308e327d4e9f0c7 Mon Sep 17 00:00:00 2001 From: Dale Seo <5466341+DaleSeo@users.noreply.github.com> Date: Tue, 8 Sep 2026 19:04:40 -0400 Subject: [PATCH 20/29] feat(macros): reject empty tool_router (#1233) Closes #1174 Claude-Session: https://claude.ai/code/session_017tyTJfwGpE8fidbZSRi6uU --- crates/rmcp-macros/src/lib.rs | 32 ++++++++++++ crates/rmcp-macros/src/tool_router.rs | 72 +++++++++++++++++++++++++++ crates/rmcp/tests/test_tool_macros.rs | 44 ++++++++++++++++ 3 files changed, 148 insertions(+) diff --git a/crates/rmcp-macros/src/lib.rs b/crates/rmcp-macros/src/lib.rs index 156e53b4a..176ead7cb 100644 --- a/crates/rmcp-macros/src/lib.rs +++ b/crates/rmcp-macros/src/lib.rs @@ -56,6 +56,7 @@ pub fn tool(attr: TokenStream, input: TokenStream) -> TokenStream { /// | `router` | `Ident` | The name of the router function to be generated. Defaults to `tool_router`. | /// | `vis` | `Visibility` | The visibility of the generated router function. Defaults to empty. | /// | `server_handler` | `flag` | When set, also emits `#[::rmcp::tool_handler]` on `impl ServerHandler for Self` so you can omit a separate `#[tool_handler]` block. | +/// | `allow_empty` | `flag` | When set, accepts an impl block with no `#[tool]` fn. Without it, an empty router is a compile error. | /// /// ## Example /// @@ -122,6 +123,37 @@ pub fn tool(attr: TokenStream, input: TokenStream) -> TokenStream { /// } /// } /// ``` +/// +/// ### Empty routers +/// +/// Collecting tools is this attribute's whole purpose, so an impl block with no `#[tool]` fn is a +/// compile error rather than a router that silently serves nothing. Pass `allow_empty` when that +/// is what you want: +/// +/// ```rust,ignore +/// #[tool_router(allow_empty)] +/// impl MyToolHandler {} +/// ``` +/// +/// The usual way to hit this by accident is a `macro_rules!` helper *inside* the impl block. An +/// attribute macro receives the unexpanded item, so `#[tool]` fns produced by such a helper are +/// invisible to `#[tool_router]`. Let the `macro_rules!` emit the whole annotated impl instead: +/// +/// ```rust,ignore +/// macro_rules! define_tools { +/// ($($name:ident => $description:literal),* $(,)?) => { +/// #[tool_router] +/// impl MyToolHandler { +/// $( +/// #[tool(description = $description)] +/// async fn $name(&self) -> String { stringify!($name).to_owned() } +/// )* +/// } +/// }; +/// } +/// +/// define_tools!(my_tool => "what my tool does"); +/// ``` #[proc_macro_attribute] pub fn tool_router(attr: TokenStream, input: TokenStream) -> TokenStream { tool_router::tool_router(attr.into(), input.into()) diff --git a/crates/rmcp-macros/src/tool_router.rs b/crates/rmcp-macros/src/tool_router.rs index edc8630b4..6efd6927f 100644 --- a/crates/rmcp-macros/src/tool_router.rs +++ b/crates/rmcp-macros/src/tool_router.rs @@ -17,6 +17,8 @@ pub struct ToolRouterAttribute { /// When set, also emit `#[::rmcp::tool_handler]` on `impl ServerHandler for Self` so callers /// can skip a separate `#[tool_handler]` block (expanded in a later macro pass). pub server_handler: bool, + /// When set, accept an impl block with no `#[tool]` fn instead of reporting an error. + pub allow_empty: bool, } impl Default for ToolRouterAttribute { @@ -25,6 +27,7 @@ impl Default for ToolRouterAttribute { router: format_ident!("tool_router"), vis: None, server_handler: false, + allow_empty: false, } } } @@ -35,6 +38,7 @@ pub fn tool_router(attr: TokenStream, input: TokenStream) -> syn::Result(input)?; // find all function marked with `#[rmcp::tool]` @@ -58,6 +62,19 @@ pub fn tool_router(attr: TokenStream, input: TokenStream) -> syn::Result syn::Result<()> { + let input = quote! { + impl Probe { + #[tool(description = "probe")] + async fn probe(&self) -> String { "probed".to_owned() } + } + }; + let generated = tool_router(TokenStream::new(), input)?.to_string(); + assert!(generated.contains("with_route"), "{generated}"); + Ok(()) + } + + #[test] + fn tool_router_allow_empty_generates_a_router_without_routes() -> syn::Result<()> { + let generated = tool_router(quote! { allow_empty }, quote! { impl Probe {} })?.to_string(); + assert!(generated.contains("fn tool_router"), "{generated}"); + assert!(!generated.contains("with_route"), "{generated}"); + Ok(()) + } } diff --git a/crates/rmcp/tests/test_tool_macros.rs b/crates/rmcp/tests/test_tool_macros.rs index 9b9530aa8..4975109cb 100644 --- a/crates/rmcp/tests/test_tool_macros.rs +++ b/crates/rmcp/tests/test_tool_macros.rs @@ -573,3 +573,47 @@ fn test_manual_get_info_not_overridden() { "manual resources should be preserved" ); } + +/// Server whose tools come from a `macro_rules!` helper wrapping the whole annotated impl. +#[derive(Debug, Clone)] +struct MacroGeneratedServer; + +macro_rules! define_tools { + ($($name:ident => $description:literal),* $(,)?) => { + #[tool_router] + impl MacroGeneratedServer { + $( + #[tool(description = $description)] + async fn $name(&self) -> String { + stringify!($name).to_owned() + } + )* + } + }; +} + +define_tools!(probe => "what a capability would own"); + +#[test] +fn test_macro_rules_around_the_impl_registers_tools() { + let tools = MacroGeneratedServer::tool_router().list_all(); + + assert_eq!(tools.len(), 1); + assert_eq!(tools[0].name, "probe"); + assert_eq!( + tools[0].description.as_deref(), + Some("what a capability would own") + ); +} + +/// Server that opts in to a router with no tools. +#[derive(Debug, Clone)] +struct EmptyRouterServer; + +#[tool_router(allow_empty)] +impl EmptyRouterServer {} + +#[test] +fn test_allow_empty_builds_a_router_without_tools() { + assert!(EmptyRouterServer::tool_router().list_all().is_empty()); +} From f9a68bfa5cad065727bf847086010d5bf7f2028d Mon Sep 17 00:00:00 2001 From: Dale Seo <5466341+DaleSeo@users.noreply.github.com> Date: Tue, 8 Sep 2026 19:05:04 -0400 Subject: [PATCH 21/29] feat: add ServerHandler::negotiate_initialize (#1247) --- crates/rmcp/src/handler/server.rs | 66 ++++++++++-- crates/rmcp/src/model.rs | 84 ++++++++++++++- crates/rmcp/src/service/server.rs | 4 +- .../test_protocol_version_negotiation.rs | 100 +++++++++++++++++- examples/servers/src/common/counter.rs | 5 +- 5 files changed, 246 insertions(+), 13 deletions(-) diff --git a/crates/rmcp/src/handler/server.rs b/crates/rmcp/src/handler/server.rs index b5f70be3f..70985527c 100644 --- a/crates/rmcp/src/handler/server.rs +++ b/crates/rmcp/src/handler/server.rs @@ -321,16 +321,59 @@ macro_rules! server_handler_methods { context: RequestContext, ) -> impl Future> + MaybeSendFuture + '_ { context.peer.set_peer_info(request.clone()); + std::future::ready(self.negotiate_initialize(&request)) + } + /// Build the `initialize` response for `request`, negotiating the + /// protocol version against [`Self::supported_protocol_versions`]. + /// + /// This is the whole body of the default [`Self::initialize`] minus its + /// `set_peer_info` side effect, so a server that overrides `initialize` + /// to add its own can call this instead of restating the negotiation + /// rule: + /// + /// ``` + /// use rmcp::{ + /// ErrorData as McpError, RoleServer, ServerHandler, + /// model::{InitializeRequestParams, InitializeResult, ServerInfo}, + /// service::RequestContext, + /// }; + /// + /// struct MyServer; + /// + /// impl ServerHandler for MyServer { + /// fn get_info(&self) -> ServerInfo { + /// ServerInfo::default() + /// } + /// + /// async fn initialize( + /// &self, + /// request: InitializeRequestParams, + /// context: RequestContext, + /// ) -> Result { + /// // ... record telemetry, register the peer, etc. + /// context.peer.set_peer_info(request.clone()); + /// self.negotiate_initialize(&request) + /// } + /// } + /// ``` + /// + /// # Errors + /// + /// Returns [`ErrorCode::UNSUPPORTED_PROTOCOL_VERSION`] when this server + /// supports no version that still has an `initialize` handshake. + /// + /// [`ErrorCode::UNSUPPORTED_PROTOCOL_VERSION`]: crate::model::ErrorCode::UNSUPPORTED_PROTOCOL_VERSION + fn negotiate_initialize( + &self, + request: &InitializeRequestParams, + ) -> Result { let mut info = self.get_info(); - let negotiated = negotiate_protocol_version( + info.protocol_version = negotiate_protocol_version( &request.protocol_version, std::mem::take(&mut info.protocol_version), &self.supported_protocol_versions(), - ); - std::future::ready(negotiated.map(|version| { - info.protocol_version = version; - info - })) + )?; + Ok(info) } /// Return the protocol versions supported by this server. /// @@ -339,6 +382,10 @@ macro_rules! server_handler_methods { /// list is advertised by [`Self::discover`], bounds what `initialize` /// negotiation may agree to, and is what per-request versions are /// validated against. + /// + /// To support everything up to some ceiling, use + /// [`ProtocolVersion::known_up_to`] rather than filtering + /// [`ProtocolVersion::KNOWN_VERSIONS`] by hand. fn supported_protocol_versions(&self) -> Cow<'static, [ProtocolVersion]> { Cow::Borrowed(ProtocolVersion::KNOWN_VERSIONS) } @@ -621,6 +668,13 @@ macro_rules! impl_server_handler_for_wrapper { (**self).initialize(request, context) } + fn negotiate_initialize( + &self, + request: &InitializeRequestParams, + ) -> Result { + (**self).negotiate_initialize(request) + } + fn supported_protocol_versions(&self) -> Cow<'static, [ProtocolVersion]> { (**self).supported_protocol_versions() } diff --git a/crates/rmcp/src/model.rs b/crates/rmcp/src/model.rs index 6a1409870..be911f7f9 100644 --- a/crates/rmcp/src/model.rs +++ b/crates/rmcp/src/model.rs @@ -177,7 +177,7 @@ impl ProtocolVersion { /// First protocol version that requires SEP-2243 standard HTTP headers. pub const STANDARD_HEADERS: Self = Self::V_2026_07_28; - /// All protocol versions known to this SDK. + /// All protocol versions known to this SDK, oldest first. pub const KNOWN_VERSIONS: &[Self] = &[ Self::V_2024_11_05, Self::V_2025_03_26, @@ -190,6 +190,43 @@ impl ProtocolVersion { pub fn as_str(&self) -> &str { &self.0 } + + /// The known versions up to and including `max`, oldest first. + /// + /// Servers that implement every revision up to some ceiling can return + /// this from `supported_protocol_versions` instead of filtering + /// [`Self::KNOWN_VERSIONS`] by hand. `max` itself need not be a known + /// version; the result is empty when it predates all of them. + /// + /// The result borrows from [`Self::KNOWN_VERSIONS`], so call it directly + /// in the method body — it needs no `static` and no `LazyLock`: + /// + /// ```rust,ignore + /// const MAX_SUPPORTED: ProtocolVersion = ProtocolVersion::V_2025_11_25; + /// + /// fn supported_protocol_versions(&self) -> Cow<'static, [ProtocolVersion]> { + /// Cow::Borrowed(ProtocolVersion::known_up_to(&MAX_SUPPORTED)) + /// } + /// ``` + /// + /// ``` + /// # use rmcp::model::ProtocolVersion; + /// assert_eq!( + /// ProtocolVersion::known_up_to(&ProtocolVersion::V_2025_06_18), + /// &[ + /// ProtocolVersion::V_2024_11_05, + /// ProtocolVersion::V_2025_03_26, + /// ProtocolVersion::V_2025_06_18, + /// ], + /// ); + /// ``` + pub fn known_up_to(max: &Self) -> &'static [Self] { + let count = Self::KNOWN_VERSIONS + .iter() + .take_while(|version| version.as_str() <= max.as_str()) + .count(); + &Self::KNOWN_VERSIONS[..count] + } } impl Serialize for ProtocolVersion { @@ -4643,6 +4680,51 @@ mod tests { use super::*; + #[test] + fn known_versions_are_ordered_oldest_first() { + // `known_up_to` walks the list as a sorted prefix. + assert!( + ProtocolVersion::KNOWN_VERSIONS + .windows(2) + .all(|pair| pair[0].as_str() < pair[1].as_str()) + ); + } + + #[test] + fn known_up_to_includes_the_ceiling_itself() { + assert_eq!( + ProtocolVersion::known_up_to(&ProtocolVersion::V_2024_11_05), + &[ProtocolVersion::V_2024_11_05] + ); + } + + #[test] + fn known_up_to_the_newest_version_yields_every_known_version() { + assert_eq!( + ProtocolVersion::known_up_to(&ProtocolVersion::V_2026_07_28), + ProtocolVersion::KNOWN_VERSIONS + ); + } + + #[test] + fn known_up_to_an_unknown_ceiling_stops_at_the_versions_below_it() { + let unknown = ProtocolVersion(Cow::Borrowed("2025-07-01")); + assert_eq!( + ProtocolVersion::known_up_to(&unknown), + &[ + ProtocolVersion::V_2024_11_05, + ProtocolVersion::V_2025_03_26, + ProtocolVersion::V_2025_06_18, + ] + ); + } + + #[test] + fn known_up_to_a_ceiling_below_every_known_version_is_empty() { + let ancient = ProtocolVersion(Cow::Borrowed("1999-01-01")); + assert!(ProtocolVersion::known_up_to(&ancient).is_empty()); + } + #[cfg(feature = "transport-streamable-http-client")] #[test] fn transport_closed_marker_accepts_only_the_process_local_token() { diff --git a/crates/rmcp/src/service/server.rs b/crates/rmcp/src/service/server.rs index 70a148641..29e46907a 100644 --- a/crates/rmcp/src/service/server.rs +++ b/crates/rmcp/src/service/server.rs @@ -499,7 +499,9 @@ pub(crate) fn negotiate_protocol_version( server_supported, )); }; - tracing::warn!( + // Falling back is the designed answer for a pinned client, and stateless + // HTTP re-runs it on every request, so this is not a warning. + tracing::debug!( client_requested = %client_requested, server_fallback = %legacy_fallback, "client requested a protocol version unavailable over initialize; falling back to server default" diff --git a/crates/rmcp/tests/test_protocol_version_negotiation.rs b/crates/rmcp/tests/test_protocol_version_negotiation.rs index 717a63f18..aa4607a1c 100644 --- a/crates/rmcp/tests/test_protocol_version_negotiation.rs +++ b/crates/rmcp/tests/test_protocol_version_negotiation.rs @@ -5,13 +5,19 @@ #![cfg(not(feature = "local"))] #![cfg(feature = "client")] -use std::borrow::Cow; +use std::{ + borrow::Cow, + sync::{ + Arc, + atomic::{AtomicUsize, Ordering}, + }, +}; use rmcp::{ ClientHandler, ErrorData, RoleServer, ServerHandler, ServiceExt, model::{ - ClientInfo, ErrorCode, InitializeRequestParams, InitializeResult, ProtocolVersion, - ServerInfo, + ClientCapabilities, ClientInfo, ErrorCode, Implementation, InitializeRequestParams, + InitializeResult, ProtocolVersion, ServerInfo, }, service::{ClientInitializeError, RequestContext}, }; @@ -232,3 +238,91 @@ async fn narrowed_server_caps_even_when_it_overrides_initialize() { "the handshake layer should not raise the version above what the server supports" ); } + +/// Overrides `initialize` to run a side effect, then delegates the version +/// answer back to the SDK with [`ServerHandler::negotiate_initialize`]. +#[derive(Debug, Clone, Default)] +struct DelegatingServer { + initializations: Arc, +} + +impl ServerHandler for DelegatingServer { + fn get_info(&self) -> ServerInfo { + ServerInfo::default() + } + + fn supported_protocol_versions(&self) -> Cow<'static, [ProtocolVersion]> { + Cow::Borrowed(HANDSHAKE_VERSIONS) + } + + async fn initialize( + &self, + request: InitializeRequestParams, + context: RequestContext, + ) -> Result { + self.initializations.fetch_add(1, Ordering::Relaxed); + context.peer.set_peer_info(request.clone()); + self.negotiate_initialize(&request) + } +} + +fn initialize_params(protocol_version: ProtocolVersion) -> InitializeRequestParams { + let mut params = InitializeRequestParams::new( + ClientCapabilities::default(), + Implementation::new("test-client", "0.0.0"), + ); + params.protocol_version = protocol_version; + params +} + +#[test] +fn negotiate_initialize_echoes_a_supported_version() { + let result = NarrowedServer + .negotiate_initialize(&initialize_params(ProtocolVersion::V_2025_06_18)) + .expect("a supported handshake version should negotiate"); + assert_eq!(result.protocol_version, ProtocolVersion::V_2025_06_18); +} + +#[test] +fn negotiate_initialize_caps_at_supported_versions() { + let result = NarrowedServer + .negotiate_initialize(&initialize_params(ProtocolVersion::V_2026_07_28)) + .expect("an unsupported version should fall back rather than fail"); + assert_eq!(result.protocol_version, ProtocolVersion::V_2025_11_25); +} + +#[test] +fn negotiate_initialize_keeps_the_rest_of_get_info() { + let server = NarrowedServer; + let result = server + .negotiate_initialize(&initialize_params(ProtocolVersion::V_2026_07_28)) + .expect("an unsupported version should fall back rather than fail"); + assert_eq!(result.capabilities, server.get_info().capabilities); +} + +#[test] +fn negotiate_initialize_rejects_when_no_handshake_version_is_supported() { + let error = ModernOnlyServer + .negotiate_initialize(&initialize_params(ProtocolVersion::V_2026_07_28)) + .expect_err("a server with no handshake version cannot answer initialize"); + assert_eq!(error.code, ErrorCode::UNSUPPORTED_PROTOCOL_VERSION); +} + +#[tokio::test] +async fn delegating_server_negotiates_like_the_default_initialize() { + let negotiated = + negotiated_version_with(DelegatingServer::default(), ProtocolVersion::V_2026_07_28).await; + assert_eq!( + negotiated, + ProtocolVersion::V_2025_11_25, + "an override that delegates should answer what the default initialize would" + ); +} + +#[tokio::test] +async fn delegating_server_still_runs_its_own_side_effect() { + let server = DelegatingServer::default(); + let initializations = Arc::clone(&server.initializations); + negotiated_version_with(server, ProtocolVersion::V_2025_06_18).await; + assert_eq!(initializations.load(Ordering::Relaxed), 1); +} diff --git a/examples/servers/src/common/counter.rs b/examples/servers/src/common/counter.rs index c6602770f..3bd16173c 100644 --- a/examples/servers/src/common/counter.rs +++ b/examples/servers/src/common/counter.rs @@ -270,7 +270,7 @@ impl ServerHandler for Counter { async fn initialize( &self, - _request: InitializeRequestParams, + request: InitializeRequestParams, context: RequestContext, ) -> Result { if let Some(http_request_part) = context.extensions.get::() { @@ -278,7 +278,8 @@ impl ServerHandler for Counter { let initialize_uri = &http_request_part.uri; tracing::info!(?initialize_headers, %initialize_uri, "initialize from http server"); } - Ok(self.get_info()) + context.peer.set_peer_info(request.clone()); + self.negotiate_initialize(&request) } } From 6fa40536deb44ceaedad28f783444c4700207270 Mon Sep 17 00:00:00 2001 From: Dale Seo <5466341+DaleSeo@users.noreply.github.com> Date: Tue, 8 Sep 2026 19:05:29 -0400 Subject: [PATCH 22/29] fix: resolve clippy warnings across workspace (#1195) * fix: resolve clippy warnings across workspace * ci: also run clippy on rmcp with all features except local --all-features enables the local feature, which cfg-gates out tower.rs and most of the test suite, so the existing clippy step never actually lints them. --- .github/workflows/ci.yml | 10 +- conformance/src/bin/client.rs | 5 + .../transport/streamable_http_server/tower.rs | 202 ++++++++---------- crates/rmcp/tests/test_custom_headers.rs | 90 ++++---- crates/rmcp/tests/test_prompt_macros.rs | 6 +- crates/rmcp/tests/test_prompt_routers.rs | 6 - crates/rmcp/tests/test_sampling.rs | 2 +- .../test_sep_2260_request_association.rs | 9 +- crates/rmcp/tests/test_task.rs | 6 +- crates/rmcp/tests/test_tool_routers.rs | 6 - .../rmcp/tests/test_unix_socket_transport.rs | 86 ++++---- examples/clients/src/progress_client.rs | 32 +-- examples/clients/src/sampling_stdio.rs | 5 + examples/servers/src/cimd_auth_streamhttp.rs | 20 +- examples/servers/src/common/progress_demo.rs | 4 +- examples/servers/src/completion_stdio.rs | 30 +-- .../servers/src/complex_auth_streamhttp.rs | 8 +- .../servers/src/elicitation_enum_inference.rs | 4 +- examples/servers/src/elicitation_stdio.rs | 4 +- examples/servers/src/prompt_stdio.rs | 26 +-- examples/servers/src/sampling_stdio.rs | 5 +- 21 files changed, 277 insertions(+), 289 deletions(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index c71d9d502..340cf9a9d 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -61,9 +61,17 @@ jobs: - uses: Swatinem/rust-cache@6323deb102c322ba6fcbdcafc7e3dddab59af2b6 # v2 - - name: Run clippy + - name: Run clippy (all features) run: cargo clippy --all-targets --all-features -- -D warnings + - name: Run clippy (all features except local) + run: | + FEATURES=$(cargo metadata --no-deps --format-version 1 \ + | jq -r '[.packages[] | select(.name == "rmcp") | .features | keys[] + | select(startswith("__") | not) + | select(. != "local")] | join(",")') + cargo clippy --package rmcp --all-targets --no-default-features --features "$FEATURES" -- -D warnings + semver: name: SemVer Check runs-on: ubuntu-latest diff --git a/conformance/src/bin/client.rs b/conformance/src/bin/client.rs index ea98a8de6..92e4a945a 100644 --- a/conformance/src/bin/client.rs +++ b/conformance/src/bin/client.rs @@ -1,3 +1,8 @@ +#![expect( + deprecated, + reason = "The conformance suite still exercises deprecated sampling scenarios" +)] + use anyhow::Context; use oauth2::{ClientSecret, RefreshToken}; use rmcp::{ diff --git a/crates/rmcp/src/transport/streamable_http_server/tower.rs b/crates/rmcp/src/transport/streamable_http_server/tower.rs index e15dfd434..f03014e02 100644 --- a/crates/rmcp/src/transport/streamable_http_server/tower.rs +++ b/crates/rmcp/src/transport/streamable_http_server/tower.rs @@ -55,6 +55,24 @@ use crate::{ pub(crate) const DEFAULT_MAX_REQUEST_BODY_BYTES: usize = 4 * 1024 * 1024; const STATELESS_STREAM_CHANNEL_CAPACITY: usize = 16; +struct ErrorResponse(Box); + +impl ErrorResponse { + fn into_response(self) -> BoxResponse { + *self.0 + } +} + +impl From for ErrorResponse { + fn from(response: BoxResponse) -> Self { + Self(Box::new(response)) + } +} + +type HttpResult = Result; +type RestoreResultSender = tokio::sync::watch::Sender>; +type PendingRestores = Arc>>; + #[non_exhaustive] #[derive(Debug, Clone)] pub struct StreamableHttpServerConfig { @@ -247,10 +265,6 @@ impl StreamableHttpServerConfig { } } -#[expect( - clippy::result_large_err, - reason = "BoxResponse is intentionally large; matches other handlers in this file" -)] /// Validates the `MCP-Protocol-Version` header on incoming HTTP requests. /// /// Per the MCP 2025-06-18 spec: @@ -259,7 +273,7 @@ impl StreamableHttpServerConfig { fn validate_protocol_version_header( headers: &http::HeaderMap, allow_unknown: bool, -) -> Result<(), BoxResponse> { +) -> HttpResult<()> { if let Some(value) = headers.get(HEADER_MCP_PROTOCOL_VERSION) { let version_str = value.to_str().map_err(|_| { Response::builder() @@ -284,7 +298,8 @@ fn validate_protocol_version_header( ))) .boxed(), ) - .expect("valid response")); + .expect("valid response") + .into()); } } Ok(()) @@ -349,28 +364,24 @@ impl> Service for NegotiatingStatelessHttpSer } } -#[expect( - clippy::result_large_err, - reason = "BoxResponse is intentionally large; matches other handlers in this file" -)] // SEP-2567: sessions are removed from the discover lifecycle. Validate // protocol-version consistency, then classify the request with the shared // lifecycle helper. fn is_legacy_request( message: Option<&ClientJsonRpcMessage>, headers: &HeaderMap, -) -> Result { +) -> HttpResult { let has_per_request_version = message.is_some_and(message_has_per_request_protocol_version); validate_protocol_version_header(headers, has_per_request_version)?; if let Some(message) = message { - if let ClientJsonRpcMessage::Request(req) = message { - if let ClientRequest::InitializeRequest(init) = &req.request { - validate_header_matches_init_body( - headers, - init.params.protocol_version.as_str(), - Some(req.id.clone()), - )?; - } + if let ClientJsonRpcMessage::Request(req) = message + && let ClientRequest::InitializeRequest(init) = &req.request + { + validate_header_matches_init_body( + headers, + init.params.protocol_version.as_str(), + Some(req.id.clone()), + )?; } validate_request_protocol_version_meta(headers, message)?; } @@ -430,10 +441,10 @@ async fn persist_and_forward_event( output: &mut Option>, ) -> Result<(), EventStoreError> { event.event_id = Some(event_store.store_event(stream_id, &event).await?); - if let Some(sender) = output { - if sender.send(event).await.is_err() { - *output = None; - } + if let Some(sender) = output + && sender.send(event).await.is_err() + { + *output = None; } Ok(()) } @@ -464,16 +475,12 @@ fn invalid_params_jsonrpc_response( .expect("valid response") } -#[expect( - clippy::result_large_err, - reason = "BoxResponse is intentionally large; matches other handlers in this file" -)] /// Absent header is allowed; the first initialize round-trip may legitimately omit it. fn validate_header_matches_init_body( headers: &http::HeaderMap, body_version: &str, request_id: Option, -) -> Result<(), BoxResponse> { +) -> HttpResult<()> { let Some(header_value) = headers.get(HEADER_MCP_PROTOCOL_VERSION) else { return Ok(()); }; @@ -494,19 +501,16 @@ fn validate_header_matches_init_body( format!( "Invalid Request: MCP-Protocol-Version header ({header_str}) does not match initialize params.protocolVersion ({body_version})" ), - )); + ) + .into()); } Ok(()) } -#[expect( - clippy::result_large_err, - reason = "BoxResponse is intentionally large; matches other handlers in this file" -)] fn validate_request_protocol_version_meta( headers: &HeaderMap, message: &ClientJsonRpcMessage, -) -> Result<(), BoxResponse> { +) -> HttpResult<()> { let ClientJsonRpcMessage::Request(request) = message else { return Ok(()); }; @@ -530,7 +534,8 @@ fn validate_request_protocol_version_meta( "Invalid params: request _meta is missing or has malformed required fields: {}", missing.join(", ") ), - )); + ) + .into()); } return Ok(()); }; @@ -538,7 +543,8 @@ fn validate_request_protocol_version_meta( return Err(header_mismatch_jsonrpc_response( Some(request.id.clone()), "request _meta protocolVersion requires MCP-Protocol-Version header", - )); + ) + .into()); }; if header_version != meta_version.as_str() { return Err(header_mismatch_jsonrpc_response( @@ -546,7 +552,8 @@ fn validate_request_protocol_version_meta( format!( "MCP-Protocol-Version header ({header_version}) does not match request _meta protocolVersion ({meta_version})" ), - )); + ) + .into()); } Ok(()) } @@ -557,15 +564,11 @@ fn validate_request_protocol_version_meta( /// HTTP 400 / JSON-RPC `-32020` before handler dispatch. `server/discover` /// is included so the seam aligns with the per-POST header contract; its /// body-metadata rule is preserved unchanged. -#[expect( - clippy::result_large_err, - reason = "BoxResponse is intentionally large; matches other handlers in this file" -)] fn validate_required_protocol_header( config: &StreamableHttpServerConfig, headers: &HeaderMap, message: &ClientJsonRpcMessage, -) -> Result<(), BoxResponse> { +) -> HttpResult<()> { if !config.stateless_protocol_metadata_required { return Ok(()); } @@ -583,7 +586,8 @@ fn validate_required_protocol_header( Err(header_mismatch_jsonrpc_response( Some(request.id.clone()), "Missing MCP-Protocol-Version header for request requiring per-request protocol metadata", - )) + ) + .into()) } /// When `stateless_protocol_metadata_required` is enabled in stateless mode, @@ -593,14 +597,10 @@ fn validate_required_protocol_header( /// `server/discover` (whose body-metadata rule is already enforced by /// `validate_request_protocol_version_meta`), notifications, and other message /// kinds are exempt. -#[expect( - clippy::result_large_err, - reason = "BoxResponse is intentionally large; matches other handlers in this file" -)] fn validate_required_protocol_meta( config: &StreamableHttpServerConfig, message: &ClientJsonRpcMessage, -) -> Result<(), BoxResponse> { +) -> HttpResult<()> { if !config.stateless_protocol_metadata_required { return Ok(()); } @@ -619,7 +619,8 @@ fn validate_required_protocol_meta( Err(invalid_params_jsonrpc_response( Some(request.id.clone()), "Invalid params: request requires protocolVersion in request _meta", - )) + ) + .into()) } fn jsonrpc_http_status(message: &ServerJsonRpcMessage) -> http::StatusCode { @@ -640,7 +641,7 @@ fn jsonrpc_http_status(message: &ServerJsonRpcMessage) -> http::StatusCode { fn jsonrpc_message_response( message: ServerJsonRpcMessage, map_protocol_status: bool, -) -> Result { +) -> HttpResult { let status = if map_protocol_status { jsonrpc_http_status(&message) } else { @@ -674,15 +675,11 @@ fn header_mismatch_jsonrpc_response( /// The `initialize` handshake is exempt: clients emit these headers only after the /// version has been negotiated. `tool_schema` supplies the called tool's input schema /// so annotated `Mcp-Param-*` headers can be checked (no schema => those are skipped). -#[expect( - clippy::result_large_err, - reason = "BoxResponse is intentionally large; matches other handlers in this file" -)] fn validate_standard_headers( headers: &HeaderMap, message: &ClientJsonRpcMessage, tool_schema: impl Fn(&str) -> Option>, -) -> Result<(), BoxResponse> { +) -> HttpResult<()> { let version_requires_headers = headers .get(HEADER_MCP_PROTOCOL_VERSION) .and_then(|value| value.to_str().ok()) @@ -715,7 +712,7 @@ fn validate_standard_headers( .and_then(|name| name.as_str()) .and_then(tool_schema); if let Err(reason) = mcp_headers::validate_request_headers(headers, &value, schema.as_deref()) { - return Err(header_mismatch_jsonrpc_response(request_id, reason)); + return Err(header_mismatch_jsonrpc_response(request_id, reason).into()); } Ok(()) } @@ -839,10 +836,7 @@ fn bad_request_response(message: &str) -> BoxResponse { .expect("failed to build bad request response") } -fn parse_host_header( - uri: &http::Uri, - headers: &HeaderMap, -) -> Result { +fn parse_host_header(uri: &http::Uri, headers: &HeaderMap) -> HttpResult { if let Some(host) = headers.get(http::header::HOST) { let host_str = host .to_str() @@ -873,23 +867,20 @@ fn validate_dns_rebinding_headers( uri: &http::Uri, headers: &HeaderMap, config: &StreamableHttpServerConfig, -) -> Result<(), BoxResponse> { +) -> HttpResult<()> { let host = parse_host_header(uri, headers)?; if !host_is_allowed(&host, &config.allowed_hosts) { tracing::warn!( host = ?host, "rejected request with disallowed Host header (possible DNS rebinding attempt)", ); - return Err(forbidden_response("Forbidden: Host header is not allowed")); + return Err(forbidden_response("Forbidden: Host header is not allowed").into()); } validate_origin_header(headers, &config.allowed_origins)?; Ok(()) } -fn validate_origin_header( - headers: &HeaderMap, - allowed_origins: &[String], -) -> Result<(), BoxResponse> { +fn validate_origin_header(headers: &HeaderMap, allowed_origins: &[String]) -> HttpResult<()> { if allowed_origins.is_empty() { return Ok(()); } @@ -914,9 +905,7 @@ fn validate_origin_header( origin = ?origin, "rejected request with disallowed Origin header (possible cross-origin attack)", ); - return Err(forbidden_response( - "Forbidden: Origin header is not allowed", - )); + return Err(forbidden_response("Forbidden: Origin header is not allowed").into()); } Ok(()) } @@ -1012,9 +1001,7 @@ pub struct StreamableHttpService { /// same unknown session ID wait for the first restore to complete rather /// than racing to replay the initialize handshake. `None` when no external /// session store is configured (avoids allocating the map). - pending_restores: Option< - Arc>>>>, - >, + pending_restores: Option, /// Caches tool input schemas by name for SEP-2243 `Mcp-Param-*` validation. /// Populated lazily via `get_tool` so the service factory runs at most once /// per tool name. `None` value means the tool exposes no schema. @@ -1067,10 +1054,9 @@ where /// `result` defaults to `false` (failure / cancellation). Only the success path /// needs to set it to `true` before returning. struct PendingRestoreGuard { - pending_restores: - Arc>>>>, + pending_restores: PendingRestores, session_id: SessionId, - watch_tx: tokio::sync::watch::Sender>, + watch_tx: RestoreResultSender, /// The value that will be broadcast to waiting tasks on drop. result: bool, } @@ -1098,12 +1084,10 @@ where session_manager: Arc, config: StreamableHttpServerConfig, ) -> Self { - let pending_restores = config.session_store.is_some().then(|| { - Arc::new(tokio::sync::RwLock::new(HashMap::< - SessionId, - tokio::sync::watch::Sender>, - >::new())) - }); + let pending_restores = config + .session_store + .is_some() + .then(|| Arc::new(tokio::sync::RwLock::new(HashMap::new()))); Self { config, session_manager, @@ -1130,19 +1114,18 @@ where tokio::spawn(async move { let mut sender = Some(sender); - if let Some(retry) = retry { - if let Err(error) = persist_and_forward_event( + if let Some(retry) = retry + && let Err(error) = persist_and_forward_event( event_store.as_ref(), &stream_id, ServerSseMessage::retry(retry), &mut sender, ) .await - { - tracing::error!(%stream_id, %error, "failed to persist SSE priming event"); - request_ct.cancel(); - return; - } + { + tracing::error!(%stream_id, %error, "failed to persist SSE priming event"); + request_ct.cancel(); + return; } let mut first = first; @@ -1214,7 +1197,7 @@ where service: S, mut request: crate::model::JsonRpcRequest, parts: http::request::Parts, - ) -> Result { + ) -> HttpResult { let peer_info = Self::peer_info_for_stateless_request(&request, &parts.headers); request.request.extensions_mut().insert(parts); let (transport, mut receiver) = @@ -1279,10 +1262,10 @@ where /// per name to read its `ServerHandler::get_tool` definition. Used to /// validate SEP-2243 `Mcp-Param-*` headers against the request body. fn tool_schema(&self, name: &str) -> Option> { - if let Ok(cache) = self.tool_schemas.read() { - if let Some(schema) = cache.get(name) { - return schema.clone(); - } + if let Ok(cache) = self.tool_schemas.read() + && let Some(schema) = cache.get(name) + { + return schema.clone(); } let schema = self .get_service() @@ -1472,23 +1455,15 @@ where Some(init_done_tx), ); - if let Err(e) = self - .session_manager + self.session_manager .initialize_session(session_id, restore_init) .await - .map_err(|e| std::io::Error::other(e.to_string())) - { - return Err(e); - } + .map_err(|e| std::io::Error::other(e.to_string()))?; - if let Err(e) = self - .session_manager + self.session_manager .accept_message(session_id, restore_initialized) .await - .map_err(|e| std::io::Error::other(e.to_string())) - { - return Err(e); - } + .map_err(|e| std::io::Error::other(e.to_string()))?; if init_done_rx.await.is_err() { return Err(std::io::Error::other( @@ -1513,7 +1488,7 @@ where if let Err(response) = validate_dns_rebinding_headers(request.uri(), request.headers(), &self.config) { - return response; + return response.into_response(); } let method = request.method().clone(); let supports_stateless_replay = self.session_manager.event_store().is_some(); @@ -1540,10 +1515,10 @@ where }; match result { Ok(response) => response, - Err(response) => response, + Err(response) => response.into_response(), } } - async fn handle_get(&self, request: Request) -> Result + async fn handle_get(&self, request: Request) -> HttpResult where B: Body + Send + 'static, B::Error: Display, @@ -1686,7 +1661,7 @@ where )) } - async fn handle_post(&self, request: Request) -> Result + async fn handle_post(&self, request: Request) -> HttpResult where B: Body + Send + 'static, B::Error: Display, @@ -1838,7 +1813,7 @@ where let stored_init_params = match &mut message { ClientJsonRpcMessage::Request(req) => { let ClientRequest::InitializeRequest(init_req) = &req.request else { - return Err(unexpected_message_response("initialize request")); + return Err(unexpected_message_response("initialize request").into()); }; // Reject mismatched MCP-Protocol-Version header before binding the session to anything. validate_header_matches_init_body( @@ -1856,7 +1831,7 @@ where stored_init_params } _ => { - return Err(unexpected_message_response("initialize request")); + return Err(unexpected_message_response("initialize request").into()); } }; let service = self @@ -2013,7 +1988,8 @@ where std::io::ErrorKind::UnexpectedEof, "no response message received from handler", ), - )); + ) + .into()); }; tracing::trace!(?message); if matches!( @@ -2045,7 +2021,7 @@ where } } - async fn handle_delete(&self, request: Request) -> Result + async fn handle_delete(&self, request: Request) -> HttpResult where B: Body + Send + 'static, B::Error: Display, diff --git a/crates/rmcp/tests/test_custom_headers.rs b/crates/rmcp/tests/test_custom_headers.rs index 736dce18e..cb1018269 100644 --- a/crates/rmcp/tests/test_custom_headers.rs +++ b/crates/rmcp/tests/test_custom_headers.rs @@ -380,10 +380,10 @@ async fn test_mcp_custom_headers_sent_to_server() -> anyhow::Result<()> { let mut headers_map = HashMap::new(); for (name, value) in headers.iter() { let name_str = name.as_str(); - if name_str.starts_with("x-") { - if let Ok(v) = value.to_str() { - headers_map.insert(name_str.to_string(), v.to_string()); - } + if name_str.starts_with("x-") + && let Ok(v) = value.to_str() + { + headers_map.insert(name_str.to_string(), v.to_string()); } } @@ -392,48 +392,48 @@ async fn test_mcp_custom_headers_sent_to_server() -> anyhow::Result<()> { stored.extend(headers_map); // Parse the MCP request - if let Ok(json_body) = serde_json::from_slice::(&body) { - if let Some(method) = json_body.get("method").and_then(|m| m.as_str()) { - if method == "initialize" { - state.initialize_called.notify_one(); - // Return a valid MCP initialize response with session header - let response = json!({ - "jsonrpc": "2.0", - "id": json_body.get("id"), - "result": { - "protocolVersion": "2024-11-05", - "capabilities": {}, - "serverInfo": { - "name": "test-server", - "version": "1.0.0" - } + if let Ok(json_body) = serde_json::from_slice::(&body) + && let Some(method) = json_body.get("method").and_then(|m| m.as_str()) + { + if method == "initialize" { + state.initialize_called.notify_one(); + // Return a valid MCP initialize response with session header + let response = json!({ + "jsonrpc": "2.0", + "id": json_body.get("id"), + "result": { + "protocolVersion": "2024-11-05", + "capabilities": {}, + "serverInfo": { + "name": "test-server", + "version": "1.0.0" } - }); - return ( - StatusCode::OK, - [ - (http::header::CONTENT_TYPE, "application/json"), - ( - http::HeaderName::from_static("mcp-session-id"), - "test-session-123", - ), - ], - response.to_string(), - ); - } else if method == "notifications/initialized" { - // For initialized notification, return 202 Accepted - return ( - StatusCode::ACCEPTED, - [ - (http::header::CONTENT_TYPE, "application/json"), - ( - http::HeaderName::from_static("mcp-session-id"), - "test-session-123", - ), - ], - String::new(), - ); - } + } + }); + return ( + StatusCode::OK, + [ + (http::header::CONTENT_TYPE, "application/json"), + ( + http::HeaderName::from_static("mcp-session-id"), + "test-session-123", + ), + ], + response.to_string(), + ); + } else if method == "notifications/initialized" { + // For initialized notification, return 202 Accepted + return ( + StatusCode::ACCEPTED, + [ + (http::header::CONTENT_TYPE, "application/json"), + ( + http::HeaderName::from_static("mcp-session-id"), + "test-session-123", + ), + ], + String::new(), + ); } } diff --git a/crates/rmcp/tests/test_prompt_macros.rs b/crates/rmcp/tests/test_prompt_macros.rs index 7a00249a4..642ae88da 100644 --- a/crates/rmcp/tests/test_prompt_macros.rs +++ b/crates/rmcp/tests/test_prompt_macros.rs @@ -4,14 +4,12 @@ use std::sync::Arc; use rmcp::{ - ClientHandler, RoleServer, ServerHandler, ServiceExt, + ClientHandler, ServerHandler, ServiceExt, handler::server::{router::prompt::PromptRouter, wrapper::Parameters}, model::{ - ClientInfo, ContentBlock, GetPromptRequestParams, GetPromptResult, ListPromptsResult, - PaginatedRequestParams, PromptMessage, Role, + ClientInfo, ContentBlock, GetPromptRequestParams, GetPromptResult, PromptMessage, Role, }, prompt, prompt_handler, prompt_router, - service::RequestContext, }; use schemars::JsonSchema; use serde::{Deserialize, Serialize}; diff --git a/crates/rmcp/tests/test_prompt_routers.rs b/crates/rmcp/tests/test_prompt_routers.rs index eecdb1e4a..e73057947 100644 --- a/crates/rmcp/tests/test_prompt_routers.rs +++ b/crates/rmcp/tests/test_prompt_routers.rs @@ -20,12 +20,6 @@ struct Request { fields: HashMap, } -#[derive(Debug, schemars::JsonSchema, serde::Deserialize, serde::Serialize)] -struct Sum { - a: i32, - b: i32, -} - #[rmcp::prompt_router(router = "test_router")] impl TestHandler { #[rmcp::prompt] diff --git a/crates/rmcp/tests/test_sampling.rs b/crates/rmcp/tests/test_sampling.rs index b108ff412..3c82fd42b 100644 --- a/crates/rmcp/tests/test_sampling.rs +++ b/crates/rmcp/tests/test_sampling.rs @@ -375,7 +375,7 @@ fn test_tool_result_content_requires_content() { #[case::array(serde_json::json!([{ "city": "SF", "temp": 72 }, { "city": "NY", "temp": 65 }]))] #[case::string(serde_json::json!("sunny"))] #[case::integer(serde_json::json!(42))] -#[case::float(serde_json::json!(3.14))] +#[case::float(serde_json::json!(3.5))] #[case::boolean(serde_json::json!(true))] fn tool_result_content_round_trips_non_object_structured_content( #[case] structured: serde_json::Value, diff --git a/crates/rmcp/tests/test_sep_2260_request_association.rs b/crates/rmcp/tests/test_sep_2260_request_association.rs index b0e7af905..b02cff037 100644 --- a/crates/rmcp/tests/test_sep_2260_request_association.rs +++ b/crates/rmcp/tests/test_sep_2260_request_association.rs @@ -1,5 +1,8 @@ #![cfg(all(feature = "server", feature = "client", not(feature = "local")))] -#![allow(deprecated)] +#![expect( + deprecated, + reason = "This test verifies request association for the deprecated sampling API" +)] use std::sync::{Arc, Mutex}; @@ -18,9 +21,11 @@ use tokio::{ sync::oneshot, }; +type RequestResultSender = oneshot::Sender>; + #[derive(Clone)] struct SamplingServer { - outside: Arc>>>>, + outside: Arc>>, } impl ServerHandler for SamplingServer { diff --git a/crates/rmcp/tests/test_task.rs b/crates/rmcp/tests/test_task.rs index ea1a2595e..55ba4ccb2 100644 --- a/crates/rmcp/tests/test_task.rs +++ b/crates/rmcp/tests/test_task.rs @@ -13,9 +13,9 @@ use rmcp::{ use serde_json::json; #[derive(Debug, serde::Deserialize, rmcp::schemars::JsonSchema)] -pub struct SumArgs { - pub a: i32, - pub b: i32, +struct SumArgs { + a: i32, + b: i32, } #[derive(Clone)] diff --git a/crates/rmcp/tests/test_tool_routers.rs b/crates/rmcp/tests/test_tool_routers.rs index d2bbe8687..12a5d4000 100644 --- a/crates/rmcp/tests/test_tool_routers.rs +++ b/crates/rmcp/tests/test_tool_routers.rs @@ -21,12 +21,6 @@ struct Request { fields: HashMap, } -#[derive(Debug, schemars::JsonSchema, serde::Deserialize, serde::Serialize)] -struct Sum { - a: i32, - b: i32, -} - #[rmcp::tool_router(router = test_router_1)] impl TestHandler { #[rmcp::tool] diff --git a/crates/rmcp/tests/test_unix_socket_transport.rs b/crates/rmcp/tests/test_unix_socket_transport.rs index 4c4ad52f1..d0e5397a3 100644 --- a/crates/rmcp/tests/test_unix_socket_transport.rs +++ b/crates/rmcp/tests/test_unix_socket_transport.rs @@ -35,10 +35,10 @@ async fn mcp_handler( let mut headers_map = HashMap::new(); for (name, value) in headers.iter() { let name_str = name.as_str(); - if name_str.starts_with("x-") || name_str == "host" { - if let Ok(v) = value.to_str() { - headers_map.insert(name_str.to_string(), v.to_string()); - } + if (name_str.starts_with("x-") || name_str == "host") + && let Ok(v) = value.to_str() + { + headers_map.insert(name_str.to_string(), v.to_string()); } } @@ -46,46 +46,46 @@ async fn mcp_handler( stored.extend(headers_map); drop(stored); - if let Ok(json_body) = serde_json::from_slice::(&body) { - if let Some(method) = json_body.get("method").and_then(|m| m.as_str()) { - if method == "initialize" { - state.initialize_called.notify_one(); - let response = json!({ - "jsonrpc": "2.0", - "id": json_body.get("id"), - "result": { - "protocolVersion": "2024-11-05", - "capabilities": {}, - "serverInfo": { - "name": "test-unix-server", - "version": "1.0.0" - } + if let Ok(json_body) = serde_json::from_slice::(&body) + && let Some(method) = json_body.get("method").and_then(|m| m.as_str()) + { + if method == "initialize" { + state.initialize_called.notify_one(); + let response = json!({ + "jsonrpc": "2.0", + "id": json_body.get("id"), + "result": { + "protocolVersion": "2024-11-05", + "capabilities": {}, + "serverInfo": { + "name": "test-unix-server", + "version": "1.0.0" } - }); - return ( - StatusCode::OK, - [ - (http::header::CONTENT_TYPE, "application/json"), - ( - http::HeaderName::from_static("mcp-session-id"), - "unix-test-session", - ), - ], - response.to_string(), - ); - } else if method == "notifications/initialized" { - return ( - StatusCode::ACCEPTED, - [ - (http::header::CONTENT_TYPE, "application/json"), - ( - http::HeaderName::from_static("mcp-session-id"), - "unix-test-session", - ), - ], - String::new(), - ); - } + } + }); + return ( + StatusCode::OK, + [ + (http::header::CONTENT_TYPE, "application/json"), + ( + http::HeaderName::from_static("mcp-session-id"), + "unix-test-session", + ), + ], + response.to_string(), + ); + } else if method == "notifications/initialized" { + return ( + StatusCode::ACCEPTED, + [ + (http::header::CONTENT_TYPE, "application/json"), + ( + http::HeaderName::from_static("mcp-session-id"), + "unix-test-session", + ), + ], + String::new(), + ); } } diff --git a/examples/clients/src/progress_client.rs b/examples/clients/src/progress_client.rs index a9f68c139..ad00dea32 100644 --- a/examples/clients/src/progress_client.rs +++ b/examples/clients/src/progress_client.rs @@ -100,10 +100,10 @@ impl ProgressAwareClient { } fn stop_tracking(&self) { - if let Ok(mut tracker_opt) = self.tracker.lock() { - if let Some(tracker) = tracker_opt.take() { - tracker.print_summary(); - } + if let Ok(mut tracker_opt) = self.tracker.lock() + && let Some(tracker) = tracker_opt.take() + { + tracker.print_summary(); } } } @@ -114,10 +114,10 @@ impl ClientHandler for ProgressAwareClient { params: ProgressNotificationParam, _context: NotificationContext, ) { - if let Ok(tracker_opt) = self.tracker.lock() { - if let Some(tracker) = tracker_opt.as_ref() { - tracker.handle_progress(¶ms); - } + if let Ok(tracker_opt) = self.tracker.lock() + && let Some(tracker) = tracker_opt.as_ref() + { + tracker.handle_progress(¶ms); } } @@ -182,10 +182,10 @@ async fn test_stdio_transport(records: u32) -> Result<()> { .call_tool(CallToolRequestParams::new("stream_processor")) .await?; - if let Some(content) = tool_result.content.first() { - if let Some(text) = content.as_text() { - tracing::info!("Processing completed: {}", text.text); - } + if let Some(content) = tool_result.content.first() + && let Some(text) = content.as_text() + { + tracing::info!("Processing completed: {}", text.text); } service.cancel().await?; @@ -236,10 +236,10 @@ async fn test_http_transport(http_url: &str, records: u32) -> Result<()> { .call_tool(CallToolRequestParams::new("stream_processor")) .await?; - if let Some(content) = tool_result.content.first() { - if let Some(text) = content.as_text() { - tracing::info!("processing completed: {}", text.text); - } + if let Some(content) = tool_result.content.first() + && let Some(text) = content.as_text() + { + tracing::info!("processing completed: {}", text.text); } client.cancel().await?; diff --git a/examples/clients/src/sampling_stdio.rs b/examples/clients/src/sampling_stdio.rs index cc7c5f153..107958624 100644 --- a/examples/clients/src/sampling_stdio.rs +++ b/examples/clients/src/sampling_stdio.rs @@ -1,3 +1,8 @@ +#![expect( + deprecated, + reason = "This example demonstrates the deprecated MCP sampling API" +)] + use anyhow::Result; use rmcp::{ ClientHandler, ServiceExt, diff --git a/examples/servers/src/cimd_auth_streamhttp.rs b/examples/servers/src/cimd_auth_streamhttp.rs index 6a634d883..a81297715 100644 --- a/examples/servers/src/cimd_auth_streamhttp.rs +++ b/examples/servers/src/cimd_auth_streamhttp.rs @@ -151,16 +151,16 @@ async fn fetch_and_validate_client_metadata(client_id_url: &str) -> Result(progress_token.clone()) else { return Err(McpError::internal_error( - format!("Invalid format of the progress token"), + "Invalid format of the progress token", None, )); }; diff --git a/examples/servers/src/completion_stdio.rs b/examples/servers/src/completion_stdio.rs index 812ed31a8..ffaeee524 100644 --- a/examples/servers/src/completion_stdio.rs +++ b/examples/servers/src/completion_stdio.rs @@ -123,23 +123,23 @@ impl SqlQueryServer { .collect(); // If no uppercase letters found, just use first letter - if first_chars.is_empty() && !candidate.is_empty() { - if let Some(first) = candidate.chars().next() { - first_chars.push(first.to_lowercase().next().unwrap_or('\0')); - } + if first_chars.is_empty() + && let Some(first) = candidate.chars().next() + { + first_chars.push(first.to_lowercase().next().unwrap_or('\0')); } } // Special case: if query is 2 chars and we only got 1 char, try matching first 2 letters - if query_chars.len() == 2 && first_chars.len() == 1 { - if let Some(first) = candidate.chars().nth(0) { - if let Some(second) = candidate.chars().nth(1) { - first_chars = vec![ - first.to_lowercase().next().unwrap_or('\0'), - second.to_lowercase().next().unwrap_or('\0'), - ]; - } - } + if query_chars.len() == 2 + && first_chars.len() == 1 + && let Some(first) = candidate.chars().next() + && let Some(second) = candidate.chars().nth(1) + { + first_chars = vec![ + first.to_lowercase().next().unwrap_or('\0'), + second.to_lowercase().next().unwrap_or('\0'), + ]; } if query_chars.len() != first_chars.len() { @@ -193,7 +193,7 @@ impl SqlQueryServer { } } -#[prompt_router] +#[prompt_router(router = "prompt_router")] impl SqlQueryServer { #[prompt(name = "sql_query", description = "Smart SQL query builder")] async fn sql_query( @@ -308,7 +308,7 @@ impl SqlQueryServer { } } -#[prompt_handler] +#[prompt_handler(router = self.prompt_router)] impl ServerHandler for SqlQueryServer { fn get_info(&self) -> ServerInfo { ServerInfo::new( diff --git a/examples/servers/src/complex_auth_streamhttp.rs b/examples/servers/src/complex_auth_streamhttp.rs index 34c4b1584..bf763d384 100644 --- a/examples/servers/src/complex_auth_streamhttp.rs +++ b/examples/servers/src/complex_auth_streamhttp.rs @@ -69,10 +69,10 @@ impl McpOAuthStore { redirect_uri: &str, ) -> Option { let clients = self.clients.read().await; - if let Some(client) = clients.get(client_id) { - if client.redirect_uri.contains(&redirect_uri.to_string()) { - return Some(client.clone()); - } + if let Some(client) = clients.get(client_id) + && client.redirect_uri == redirect_uri + { + return Some(client.clone()); } None } diff --git a/examples/servers/src/elicitation_enum_inference.rs b/examples/servers/src/elicitation_enum_inference.rs index bed5c1db6..50d2844c6 100644 --- a/examples/servers/src/elicitation_enum_inference.rs +++ b/examples/servers/src/elicitation_enum_inference.rs @@ -97,7 +97,7 @@ struct ElicitationEnumFormServer { tool_router: ToolRouter, } -#[tool_router] +#[tool_router(router = tool_router)] impl ElicitationEnumFormServer { pub fn new() -> Self { Self { @@ -153,7 +153,7 @@ impl ElicitationEnumFormServer { } } -#[tool_handler] +#[tool_handler(router = self.tool_router)] impl ServerHandler for ElicitationEnumFormServer { fn get_info(&self) -> ServerInfo { ServerInfo::new(ServerCapabilities::builder().enable_tools().build()) diff --git a/examples/servers/src/elicitation_stdio.rs b/examples/servers/src/elicitation_stdio.rs index d506a9c7f..16b4773b0 100644 --- a/examples/servers/src/elicitation_stdio.rs +++ b/examples/servers/src/elicitation_stdio.rs @@ -64,7 +64,7 @@ impl Default for ElicitationServer { } } -#[tool_router] +#[tool_router(router = tool_router)] impl ElicitationServer { #[tool(description = "Greet user with name collection")] async fn greet_user( @@ -145,7 +145,7 @@ impl ElicitationServer { } } -#[tool_handler] +#[tool_handler(router = self.tool_router)] impl ServerHandler for ElicitationServer { fn get_info(&self) -> ServerInfo { ServerInfo::new(ServerCapabilities::builder().enable_tools().build()) diff --git a/examples/servers/src/prompt_stdio.rs b/examples/servers/src/prompt_stdio.rs index 7b6e28532..6ef24e937 100644 --- a/examples/servers/src/prompt_stdio.rs +++ b/examples/servers/src/prompt_stdio.rs @@ -112,7 +112,7 @@ impl Default for PromptServer { } } -#[prompt_router] +#[prompt_router(router = "prompt_router")] impl PromptServer { /// Simple greeting prompt without parameters #[prompt( @@ -305,17 +305,17 @@ impl PromptServer { ]; // Add tried solutions if any - if let Some(tried) = args.tried_solutions { - if !tried.is_empty() { - messages.push(PromptMessage::new_text( - Role::User, - format!("I've already tried: {}", tried.join(", ")), - )); - messages.push(PromptMessage::new_text( - Role::Assistant, - "I see you've already attempted some solutions. Let me suggest different approaches.", - )); - } + if let Some(tried) = args.tried_solutions + && !tried.is_empty() + { + messages.push(PromptMessage::new_text( + Role::User, + format!("I've already tried: {}", tried.join(", ")), + )); + messages.push(PromptMessage::new_text( + Role::Assistant, + "I see you've already attempted some solutions. Let me suggest different approaches.", + )); } messages.push(PromptMessage::new_text( @@ -361,7 +361,7 @@ impl PromptServer { } } -#[prompt_handler] +#[prompt_handler(router = self.prompt_router)] impl ServerHandler for PromptServer { fn get_info(&self) -> ServerInfo { ServerInfo::new(ServerCapabilities::builder().enable_prompts().build()).with_instructions( diff --git a/examples/servers/src/sampling_stdio.rs b/examples/servers/src/sampling_stdio.rs index be230add0..2be7f5d46 100644 --- a/examples/servers/src/sampling_stdio.rs +++ b/examples/servers/src/sampling_stdio.rs @@ -1,4 +1,7 @@ -#![allow(deprecated)] +#![expect( + deprecated, + reason = "This example demonstrates the deprecated MCP sampling API" +)] use std::sync::Arc; use anyhow::Result; From e8a0d7099aad673377b1d0e83bf192df218b2bd0 Mon Sep 17 00:00:00 2001 From: "dependabot[bot]" <49699333+dependabot[bot]@users.noreply.github.com> Date: Tue, 8 Sep 2026 20:24:09 -0400 Subject: [PATCH 23/29] chore(deps): bump taiki-e/install-action from 2.86.7 to 2.87.6 (#1250) Bumps [taiki-e/install-action](https://github.com/taiki-e/install-action) from 2.86.7 to 2.87.6. - [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/b6ff580856c41316412a0b9b60540fbc6f8c82cc...7b8d4719ee4aaa279bdf55df38dacb9ebfe12a6c) --- updated-dependencies: - dependency-name: taiki-e/install-action dependency-version: 2.87.6 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 340cf9a9d..da175293e 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -87,7 +87,7 @@ jobs: - uses: Swatinem/rust-cache@6323deb102c322ba6fcbdcafc7e3dddab59af2b6 # v2 - name: Install cargo-semver-checks - uses: taiki-e/install-action@b6ff580856c41316412a0b9b60540fbc6f8c82cc # v2.86.7 + uses: taiki-e/install-action@7b8d4719ee4aaa279bdf55df38dacb9ebfe12a6c # v2.87.6 with: tool: cargo-semver-checks @@ -141,7 +141,7 @@ jobs: - uses: Swatinem/rust-cache@6323deb102c322ba6fcbdcafc7e3dddab59af2b6 # v2 - name: Install cargo-public-api - uses: taiki-e/install-action@b6ff580856c41316412a0b9b60540fbc6f8c82cc # v2.86.7 + uses: taiki-e/install-action@7b8d4719ee4aaa279bdf55df38dacb9ebfe12a6c # v2.87.6 with: tool: cargo-public-api From 744b9f904c7f17d589326e61ddbe126cb9d58888 Mon Sep 17 00:00:00 2001 From: "dependabot[bot]" <49699333+dependabot[bot]@users.noreply.github.com> Date: Tue, 8 Sep 2026 20:24:44 -0400 Subject: [PATCH 24/29] chore(deps): bump crate-ci/typos (#1249) Bumps [crate-ci/typos](https://github.com/crate-ci/typos) from 1a51d4b5a03bb97576af186c813af67e9137ba7c to d43b6c087ac471e2ea7b8af622ff15f05c0c365b. - [Release notes](https://github.com/crate-ci/typos/releases) - [Changelog](https://github.com/crate-ci/typos/blob/master/CHANGELOG.md) - [Commits](https://github.com/crate-ci/typos/compare/1a51d4b5a03bb97576af186c813af67e9137ba7c...d43b6c087ac471e2ea7b8af622ff15f05c0c365b) --- updated-dependencies: - dependency-name: crate-ci/typos dependency-version: d43b6c087ac471e2ea7b8af622ff15f05c0c365b dependency-type: direct:production ... Signed-off-by: dependabot[bot] Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com> --- .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 da175293e..e08521904 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -214,7 +214,7 @@ jobs: steps: - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 - name: Spell Check Repo - uses: crate-ci/typos@1a51d4b5a03bb97576af186c813af67e9137ba7c # master + uses: crate-ci/typos@d43b6c087ac471e2ea7b8af622ff15f05c0c365b # master msrv: name: Check MSRV From 46db531df975ba1bd44fceabf2b3ff9cc9a21514 Mon Sep 17 00:00:00 2001 From: ump45nose <52391318+ump45nose@users.noreply.github.com> Date: Thu, 10 Sep 2026 09:24:29 +0800 Subject: [PATCH 25/29] fix(sse): saturate exponential reconnect backoff to avoid overflow panic (#1231) * fix(sse): saturate exponential reconnect backoff ExponentialBackoff::retry computed the reconnect multiplier with 2u32.pow(current_times). With max_times unset, current_times can reach the bit width, panicking in debug builds and wrapping to a zero delay in release builds for long-lived SSE clients. Use saturating_pow and Duration::saturating_mul so the delay stays monotonic and panic-free. * fix(sse): cap exponential reconnect backoff at a bounded max delay Saturating the multiplier alone can still yield decades-long sleeps once current_times reaches the bit width, pinning the stream in tokio::time::sleep without reconnecting or terminating. Add an optional max_delay (default 30s) that clamps the computed delay, keeping the backoff monotonic and panic-free while guaranteeing the client retries. * fix(sse): default max_delay to None to preserve unbounded backoff Per maintainer feedback, leave ExponentialBackoff::default() unbounded so the fix stays a pure overflow bug-fix. The saturating multiplier removes the debug panic / release wrap from #1198, while max_delay stays opt-in for callers that want a bounded reconnect delay. --- .../src/transport/common/client_side_sse.rs | 88 ++++++++++++++++++- .../src/transport/streamable_http_client.rs | 2 + 2 files changed, 89 insertions(+), 1 deletion(-) diff --git a/crates/rmcp/src/transport/common/client_side_sse.rs b/crates/rmcp/src/transport/common/client_side_sse.rs index e668d63df..2ee37aaf1 100644 --- a/crates/rmcp/src/transport/common/client_side_sse.rs +++ b/crates/rmcp/src/transport/common/client_side_sse.rs @@ -205,6 +205,11 @@ impl Default for FixedInterval { pub struct ExponentialBackoff { pub max_times: Option, pub base_duration: Duration, + /// Optional upper bound on a single reconnect delay. `None` (the default) preserves the + /// pre-existing unbounded doubling behavior. Once the multiplier saturates near the bit + /// width that can still produce very long sleeps, so callers that need the client to + /// actually reconnect can set `Some(...)` to clamp the delay. + pub max_delay: Option, } impl ExponentialBackoff { @@ -216,6 +221,7 @@ impl Default for ExponentialBackoff { Self { max_times: None, base_duration: Self::DEFAULT_DURATION, + max_delay: None, } } } @@ -227,7 +233,16 @@ impl SseRetryPolicy for ExponentialBackoff { { return None; } - Some(self.base_duration * (2u32.pow(current_times as u32))) + // `current_times` is unbounded when `max_times` is unset, so the exponent can reach + // the bit width. Saturate the multiplier at `u32::MAX` and use saturating multiplication + // for the base duration so the delay stays monotonic and panic-free instead of an + // overflow panic (debug) or a wrapped-to-zero backoff (release). + let multiplier = 2u32.saturating_pow(current_times as u32); + let delay = self.base_duration.saturating_mul(multiplier); + Some(match self.max_delay { + Some(max_delay) => delay.min(max_delay), + None => delay, + }) } } @@ -775,4 +790,75 @@ mod tests { assert!(stream.next().await.is_none()); assert_eq!(attempts.load(Ordering::Relaxed), 0); } + + #[test] + fn exponential_backoff_saturates_at_high_retry_counts() { + // With `max_times` unset, `current_times` can reach the bit width. The old + // `2u32.pow(current_times)` panicked in debug builds and wrapped in release; + // the saturating implementation must return a monotonic, non-zero delay instead. + let policy = ExponentialBackoff { + max_times: None, + base_duration: Duration::from_millis(1), + max_delay: None, + }; + let mut previous = Duration::ZERO; + for current_times in [31usize, 32, 63, 64, 100] { + let delay = policy + .retry(current_times) + .expect("unbounded policy never gives up"); + assert!( + !delay.is_zero(), + "delay must stay non-zero at {current_times}" + ); + assert!( + delay >= previous, + "delay must stay monotonic at {current_times}" + ); + previous = delay; + } + } + + #[test] + fn exponential_backoff_caps_delay_at_max_delay() { + // An explicit cap keeps the unbounded doubling policy from producing decades-long + // sleeps once the multiplier saturates. The delay must grow monotonically, stop at + // the configured ceiling, and never exceed it. + let policy = ExponentialBackoff { + max_times: None, + base_duration: Duration::from_secs(1), + max_delay: Some(Duration::from_secs(30)), + }; + let mut previous = Duration::ZERO; + for current_times in [0usize, 1, 2, 3, 4, 5, 10, 32, 64, 100] { + let delay = policy + .retry(current_times) + .expect("unbounded policy never gives up"); + assert!( + delay >= previous, + "delay must stay monotonic at {current_times}" + ); + assert!( + delay <= Duration::from_secs(30), + "delay must respect max_delay at {current_times}" + ); + previous = delay; + } + // Beyond the ceiling the delay stays pinned at max_delay. + assert_eq!( + policy.retry(100).expect("never gives up"), + Duration::from_secs(30) + ); + } + + #[test] + fn exponential_backoff_respects_max_times() { + let policy = ExponentialBackoff { + max_times: Some(3), + base_duration: Duration::from_millis(1), + max_delay: None, + }; + assert!(policy.retry(0).is_some()); + assert!(policy.retry(2).is_some()); + assert!(policy.retry(3).is_none()); + } } diff --git a/crates/rmcp/src/transport/streamable_http_client.rs b/crates/rmcp/src/transport/streamable_http_client.rs index d702bc1ca..6f27abf90 100644 --- a/crates/rmcp/src/transport/streamable_http_client.rs +++ b/crates/rmcp/src/transport/streamable_http_client.rs @@ -2293,6 +2293,7 @@ mod tests { Arc::new(ExponentialBackoff { max_times: Some(1), base_duration: Duration::ZERO, + max_delay: None, }), ); let mut stream = std::pin::pin!(stream); @@ -2397,6 +2398,7 @@ mod tests { Arc::new(ExponentialBackoff { max_times: Some(1), base_duration: Duration::ZERO, + max_delay: None, }), ); From b2c90b1edc85d0ef2b47d344e9505b9732b06d9b Mon Sep 17 00:00:00 2001 From: Jack Amadeo Date: Thu, 10 Sep 2026 10:05:55 -0400 Subject: [PATCH 26/29] ci: pin commitlint dependencies (#1217) --- .github/workflows/ci.yml | 8 +++----- 1 file changed, 3 insertions(+), 5 deletions(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index e08521904..1edfe69b2 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -31,12 +31,10 @@ jobs: node-version: '22' - name: Install commitlint - run: | - npm install --save-dev @commitlint/cli @commitlint/config-conventional - echo "module.exports = {extends: ['@commitlint/config-conventional']}" > commitlint.config.js - + run: npm install --global @commitlint/cli@20.4.3 @commitlint/config-conventional@20.4.3 + - name: Lint commit messages - run: npx commitlint --from ${{ github.event.pull_request.base.sha }} --to ${{ github.event.pull_request.head.sha }} --verbose + run: commitlint --extends @commitlint/config-conventional --from ${{ github.event.pull_request.base.sha }} --to ${{ github.event.pull_request.head.sha }} --verbose fmt: name: Code Formatting From 3e636cab26c013eca5131103c03d20237f12c4df Mon Sep 17 00:00:00 2001 From: "github-actions[bot]" <41898282+github-actions[bot]@users.noreply.github.com> Date: Thu, 10 Sep 2026 11:23:32 -0400 Subject: [PATCH 27/29] chore: release v3.3.0 (#1252) Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com> --- Cargo.toml | 6 +++--- crates/rmcp-macros/CHANGELOG.md | 6 ++++++ crates/rmcp/CHANGELOG.md | 18 ++++++++++++++++++ 3 files changed, 27 insertions(+), 3 deletions(-) diff --git a/Cargo.toml b/Cargo.toml index b3a0da365..261b77e83 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -4,13 +4,13 @@ default-members = ["crates/rmcp", "crates/rmcp-macros"] resolver = "2" [workspace.dependencies] -rmcp = { version = "3.2.0", path = "./crates/rmcp" } -rmcp-macros = { version = "3.2.0", path = "./crates/rmcp-macros" } +rmcp = { version = "3.3.0", path = "./crates/rmcp" } +rmcp-macros = { version = "3.3.0", path = "./crates/rmcp-macros" } [workspace.package] edition = "2024" rust-version = "1.88" -version = "3.2.0" +version = "3.3.0" authors = ["4t145 "] license = "Apache-2.0" repository = "https://github.com/modelcontextprotocol/rust-sdk/" diff --git a/crates/rmcp-macros/CHANGELOG.md b/crates/rmcp-macros/CHANGELOG.md index 4f59b7513..67154ba86 100644 --- a/crates/rmcp-macros/CHANGELOG.md +++ b/crates/rmcp-macros/CHANGELOG.md @@ -7,6 +7,12 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] +## [3.3.0](https://github.com/modelcontextprotocol/rust-sdk/compare/rmcp-macros-v3.2.0...rmcp-macros-v3.3.0) - 2026-09-10 + +### Added + +- *(macros)* reject empty tool_router ([#1233](https://github.com/modelcontextprotocol/rust-sdk/pull/1233)) + ## [3.2.0](https://github.com/modelcontextprotocol/rust-sdk/compare/rmcp-macros-v3.1.4...rmcp-macros-v3.2.0) - 2026-08-31 ### Added diff --git a/crates/rmcp/CHANGELOG.md b/crates/rmcp/CHANGELOG.md index 86ab5b146..5465cf210 100644 --- a/crates/rmcp/CHANGELOG.md +++ b/crates/rmcp/CHANGELOG.md @@ -7,6 +7,24 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] +## [3.3.0](https://github.com/modelcontextprotocol/rust-sdk/compare/rmcp-v3.2.0...rmcp-v3.3.0) - 2026-09-10 + +### Added + +- add ServerHandler::negotiate_initialize ([#1247](https://github.com/modelcontextprotocol/rust-sdk/pull/1247)) +- *(macros)* reject empty tool_router ([#1233](https://github.com/modelcontextprotocol/rust-sdk/pull/1233)) +- *(auth)* add enterprise refresh-token and ID-JAG exchanges ([#1234](https://github.com/modelcontextprotocol/rust-sdk/pull/1234)) + +### Fixed + +- *(sse)* saturate exponential reconnect backoff to avoid overflow panic ([#1231](https://github.com/modelcontextprotocol/rust-sdk/pull/1231)) +- resolve clippy warnings across workspace ([#1195](https://github.com/modelcontextprotocol/rust-sdk/pull/1195)) +- *(auth)* unify refresh checks and error handling ([#1236](https://github.com/modelcontextprotocol/rust-sdk/pull/1236)) + +### Other + +- *(deps)* update process-wrap requirement from 9.0 to 10.0 ([#1229](https://github.com/modelcontextprotocol/rust-sdk/pull/1229)) + ## [3.2.0](https://github.com/modelcontextprotocol/rust-sdk/compare/rmcp-v3.1.4...rmcp-v3.2.0) - 2026-08-31 ### Added From 0e1184b47645d5eb64d1df3bb84067b1d4a53340 Mon Sep 17 00:00:00 2001 From: Jacob Magar Date: Sun, 13 Sep 2026 11:36:12 -0400 Subject: [PATCH 28/29] feat: preserve typed custom responses on rmcp 3.3 and enforce Origin validation --- 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 +- .../common/auth/streamable_http_client.rs | 4 + .../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 | 205 ++++++++-- .../transport/streamable_http_server/tower.rs | 24 +- crates/rmcp/src/transport/worker.rs | 90 +++- crates/rmcp/tests/test_custom_headers.rs | 55 +++ crates/rmcp/tests/test_custom_request.rs | 383 +++++++++++++++++- 15 files changed, 953 insertions(+), 103 deletions(-) diff --git a/crates/rmcp/src/service.rs b/crates/rmcp/src/service.rs index b6cc5e538..84c73912c 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: @@ -275,6 +281,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 { @@ -533,8 +556,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, @@ -542,11 +565,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; @@ -608,7 +631,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 = @@ -692,12 +715,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, @@ -844,6 +907,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, @@ -857,16 +958,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(), @@ -911,7 +1018,7 @@ impl Peer { .send(PeerSinkMessage::Request { request, id: id.clone(), - responder, + responder: wrap_responder(responder), }) .await .is_err() @@ -948,6 +1055,7 @@ impl Peer { request, options, Some((sender, channel_capacity)), + PendingResponder::Standard, ) .await?; Ok((handle, receiver)) @@ -1347,8 +1455,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 @@ -1361,7 +1468,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)] @@ -1379,7 +1486,7 @@ where #[derive(Debug)] enum Event { ProxyMessage(PeerSinkMessage), - PeerMessage(RxJsonRpcMessage), + PeerMessage(RawRxJsonRpcMessage), ToSink(TxJsonRpcMessage), SendTaskResult(SendTaskResult), ResponseSendTaskResult(Result<(), tokio::task::JoinError>), @@ -1397,7 +1504,7 @@ where continue } } - m = transport.receive() => { + m = transport.receive_raw() => { if let Some(m) = m { Event::PeerMessage(m) } else { @@ -1445,7 +1552,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 { @@ -1463,9 +1570,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) => { @@ -1500,6 +1607,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())); { @@ -1609,9 +1720,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"); @@ -1642,8 +1753,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( @@ -1694,10 +1804,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, .. })) => { @@ -1715,10 +1822,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 29e46907a..51ad5445a 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 fdedc9c3a..31c23b551 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/auth/streamable_http_client.rs b/crates/rmcp/src/transport/common/auth/streamable_http_client.rs index 47b6caf55..165fccdb5 100644 --- a/crates/rmcp/src/transport/common/auth/streamable_http_client.rs +++ b/crates/rmcp/src/transport/common/auth/streamable_http_client.rs @@ -77,6 +77,10 @@ where { type Error = C::Error; + fn preserves_raw_responses() -> bool { + C::preserves_raw_responses() + } + async fn delete_session( &self, uri: std::sync::Arc, diff --git a/crates/rmcp/src/transport/common/client_side_sse.rs b/crates/rmcp/src/transport/common/client_side_sse.rs index 2ee37aaf1..2be78ca15 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>; @@ -386,7 +386,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<'_>, @@ -421,7 +421,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 6f27abf90..2f4d21fc9 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::{ @@ -114,6 +114,36 @@ fn cache_tools_from_response( } } +fn cache_tools_from_raw_response( + cache: &mut HashMap>, + message: &mut RawRxJsonRpcMessage, + protocol_version: &ProtocolVersion, +) { + if protocol_version < &ProtocolVersion::STANDARD_HEADERS { + return; + } + if let crate::model::JsonRpcMessage::Response(response) = message + && let Some(tools) = response + .result + .get_mut("tools") + .and_then(serde_json::Value::as_array_mut) + { + tools.retain(|value| { + // Preserve the original JSON, including extension fields. Malformed + // tools remain for the normal response decoder to reject. + let Ok(tool) = serde_json::from_value::(value.clone()) else { + return true; + }; + if let Err(reason) = mcp_headers::validate_param_header_annotations(&tool.input_schema) { + tracing::warn!(tool = %tool.name, "rejecting invalid x-mcp-header annotations: {reason}"); + return false; + } + cache.insert(tool.name.to_string(), tool.input_schema); + true + }); + } +} + fn negotiate_version_headers( init_response: &ServerJsonRpcMessage, base: HashMap, @@ -245,7 +275,7 @@ pub enum StreamableHttpProtocolError { MissingSessionIdInResponse, } -#[expect( +#[allow( clippy::large_enum_variant, reason = "boxing the streaming response would add an allocation to the common response path" )] @@ -253,6 +283,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 +293,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 +308,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 +319,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 +348,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 +363,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 +420,11 @@ 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 this backend preserves response result bodies as raw JSON. + fn preserves_raw_responses() -> bool { + false + } fn post_message( &self, uri: Arc, @@ -649,17 +693,17 @@ impl StreamableHttpClientWorker { } } - fn server_response_id(message: &ServerJsonRpcMessage) -> Option<&RequestId> { + fn server_response_id(message: &RawRxJsonRpcMessage) -> 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( pending_stream_response_ids: &mut HashSet, - message: &ServerJsonRpcMessage, + message: &RawRxJsonRpcMessage, ) -> Option { let response_id = Self::server_response_id(message)?; if let Some(id) = pending_stream_response_ids.take(response_id) { @@ -670,7 +714,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 +807,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 +837,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 +855,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, @@ -832,9 +878,10 @@ impl StreamableHttpClientWorker { } async fn execute_sse_stream( - sse_stream: impl Stream>> - + Send, - sse_worker_tx: tokio::sync::mpsc::Sender, + sse_stream: impl Stream< + Item = Result, StreamableHttpError>, + > + Send, + sse_worker_tx: tokio::sync::mpsc::Sender>, origin: InboundStreamOrigin, close_on_response: bool, ct: CancellationToken, @@ -855,12 +902,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 +934,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 +1065,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 +1099,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 +1222,7 @@ impl Worker for StreamableHttpClientWorker { StartPost(WorkerSendRequest>), PostResult(PostResult), RecoveryTimeout, - ServerMessage(ServerJsonRpcMessage), + ServerMessage(RawRxJsonRpcMessage), StreamResult { request_id: Option, result: Result<(), StreamableHttpError>, @@ -1639,6 +1689,19 @@ impl Worker for StreamableHttpClientWorker { context.send_to_handler(message).await?; Ok(()) } + Ok(StreamableHttpPostResponse::RawJson(mut raw_message, ..)) => { + if matches!(&message, ClientJsonRpcMessage::Request(request) + if matches!(&request.request, ClientRequest::ListToolsRequest(_))) + { + cache_tools_from_raw_response( + &mut tool_header_cache, + &mut raw_message, + &version, + ); + } + context.send_to_handler(raw_message).await?; + Ok(()) + } Ok(StreamableHttpPostResponse::Sse(stream, ..)) => { let stream_request_id = request_id; let sse_stream = Self::response_sse_to_jsonrpc( @@ -1718,7 +1781,7 @@ impl Worker for StreamableHttpClientWorker { ) { drop(request_stream_cancellations.remove(&request_id)); } - cache_tools_from_response( + cache_tools_from_raw_response( &mut tool_header_cache, &mut json_rpc_message, &negotiated_version, @@ -2159,9 +2222,9 @@ mod tests { deprecated, reason = "Sampling is deprecated by SEP-2577 but remains the canonical restricted request" )] - fn sampling_request_message(id: i64) -> ServerJsonRpcMessage { + fn sampling_request_message(id: i64) -> RawRxJsonRpcMessage { use crate::model::{CreateMessageRequest, CreateMessageRequestParams, SamplingMessage}; - ServerJsonRpcMessage::request( + RawRxJsonRpcMessage::::request( ServerRequest::CreateMessageRequest(CreateMessageRequest::new( CreateMessageRequestParams::new(vec![SamplingMessage::user_text("hi")], 16), )), @@ -2175,8 +2238,9 @@ mod tests { InboundStreamOrigin::Unassociated, InboundStreamOrigin::OutboundRequest(RequestId::Number(3)), ] { - let response = ServerJsonRpcMessage::response( - ServerResult::ListToolsResult(ListToolsResult::default()), + let response = RawRxJsonRpcMessage::::response( + serde_json::to_value(ServerResult::ListToolsResult(ListToolsResult::default())) + .unwrap(), NumberOrString::Number(1), ); let stream = futures::stream::iter([Ok(sampling_request_message(9)), Ok(response)]); @@ -2191,7 +2255,7 @@ mod tests { .await .expect("stream completes"); - let ServerJsonRpcMessage::Request(request) = + let crate::model::JsonRpcMessage::Request(request) = rx.recv().await.expect("request forwarded") else { panic!("expected request first"); @@ -2204,7 +2268,7 @@ mod tests { // Responses are correlated by JSON-RPC id; no marker needed or added. assert!(matches!( rx.recv().await.expect("response forwarded"), - ServerJsonRpcMessage::Response(_) + crate::model::JsonRpcMessage::Response(_) )); } } @@ -2254,8 +2318,9 @@ mod tests { .lock() .expect("lock reconnects") .push((session_id.map(|id| id.to_string()), last_event_id)); - let response = ServerJsonRpcMessage::response( - ServerResult::ListToolsResult(ListToolsResult::default()), + let response = RawRxJsonRpcMessage::::response( + serde_json::to_value(ServerResult::ListToolsResult(ListToolsResult::default())) + .expect("serialize response"), NumberOrString::Number(1), ); Ok(futures::stream::once(async move { @@ -2300,7 +2365,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()))] @@ -2351,8 +2416,9 @@ mod tests { .expect("lock reconnects") .push((session_id.map(|id| id.to_string()), last_event_id)); let request = sampling_request_message(9); - let response = ServerJsonRpcMessage::response( - ServerResult::ListToolsResult(ListToolsResult::default()), + let response = RawRxJsonRpcMessage::::response( + serde_json::to_value(ServerResult::ListToolsResult(ListToolsResult::default())) + .expect("serialize response"), NumberOrString::Number(1), ); // Stay open after the response, like a live connection, so the @@ -2419,7 +2485,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"); }; @@ -2430,7 +2497,7 @@ mod tests { ); assert!(matches!( rx.recv().await.expect("response forwarded"), - ServerJsonRpcMessage::Response(_) + crate::model::JsonRpcMessage::Response(_) )); } @@ -2482,6 +2549,36 @@ mod tests { ); } + #[test] + fn raw_tool_cache_preserves_extensions_and_filters_invalid_annotations() { + let mut valid = serde_json::to_value(tool( + "valid", + json!({"type": "string", "x-mcp-header": "Value"}), + )) + .unwrap(); + valid["vendorExtension"] = json!({"retained": true}); + let invalid = serde_json::to_value(tool( + "invalid", + json!({"type": "string", "x-mcp-header": ""}), + )) + .unwrap(); + let mut message = RawRxJsonRpcMessage::::response( + json!({"tools": [valid.clone(), invalid], "vendorResult": 42}), + NumberOrString::Number(1), + ); + let mut cache = HashMap::new(); + cache_tools_from_raw_response(&mut cache, &mut message, &ProtocolVersion::V_2026_07_28); + assert!(cache.contains_key("valid")); + assert!(!cache.contains_key("invalid")); + let crate::model::JsonRpcMessage::Response(response) = message else { + panic!("response") + }; + assert_eq!( + response.result, + json!({"tools": [valid], "vendorResult": 42}) + ); + } + #[test] fn cache_tools_preserves_pre_standard_header_results() { let invalid = tool("legacy", json!({ "type": "string", "x-mcp-header": "" })); @@ -2509,12 +2606,41 @@ mod tests { ); } + #[test] + fn raw_tool_cache_preserves_legacy_results_and_malformed_tools() { + let malformed = json!({"name": "broken", "inputSchema": "not-an-object"}); + let invalid = serde_json::to_value(tool("legacy", json!({"x-mcp-header": ""}))).unwrap(); + let original = json!({"tools": [invalid, malformed.clone()], "vendorResult": 42}); + let mut message = RawRxJsonRpcMessage::::response( + original.clone(), + NumberOrString::Number(1), + ); + let mut cache = HashMap::new(); + cache_tools_from_raw_response(&mut cache, &mut message, &ProtocolVersion::V_2025_11_25); + assert!(cache.is_empty()); + let crate::model::JsonRpcMessage::Response(response) = &message else { + panic!("response") + }; + assert_eq!(response.result, original); + + cache_tools_from_raw_response(&mut cache, &mut message, &ProtocolVersion::V_2026_07_28); + assert!(cache.is_empty()); + let crate::model::JsonRpcMessage::Response(response) = message else { + panic!("response") + }; + assert_eq!( + response.result, + json!({"tools": [malformed], "vendorResult": 42}) + ); + } + #[cfg(feature = "transport-streamable-http-client-reqwest")] #[test] fn clear_stream_response_pending_accepts_stringified_numeric_id() { let mut pending = HashSet::from([NumberOrString::Number(1)]); - let response = ServerJsonRpcMessage::response( - ServerResult::ListToolsResult(ListToolsResult::default()), + let response = RawRxJsonRpcMessage::::response( + serde_json::to_value(ServerResult::ListToolsResult(ListToolsResult::default())) + .unwrap(), NumberOrString::String("1".into()), ); @@ -2533,8 +2659,9 @@ mod tests { fn clear_stream_response_pending_prefers_exact_string_id() { let string_id = NumberOrString::String("1".into()); let mut pending = HashSet::from([NumberOrString::Number(1), string_id.clone()]); - let response = ServerJsonRpcMessage::response( - ServerResult::ListToolsResult(ListToolsResult::default()), + let response = RawRxJsonRpcMessage::::response( + serde_json::to_value(ServerResult::ListToolsResult(ListToolsResult::default())) + .unwrap(), string_id.clone(), ); diff --git a/crates/rmcp/src/transport/streamable_http_server/tower.rs b/crates/rmcp/src/transport/streamable_http_server/tower.rs index f03014e02..59ef1a03c 100644 --- a/crates/rmcp/src/transport/streamable_http_server/tower.rs +++ b/crates/rmcp/src/transport/streamable_http_server/tower.rs @@ -791,13 +791,22 @@ fn parse_origin_value(value: &str) -> Option { if value.eq_ignore_ascii_case("null") { return Some(NormalizedOrigin::Null); } + let (_, serialized_authority) = value.split_once("://")?; + if serialized_authority.contains(['/', '?', '#', '@']) { + return None; + } let uri = http::Uri::try_from(value).ok()?; let scheme = uri.scheme_str()?.to_ascii_lowercase(); let authority = uri.authority()?; + let port = authority.port_u16().or(match scheme.as_str() { + "http" => Some(80), + "https" => Some(443), + _ => None, + }); Some(NormalizedOrigin::Tuple { scheme, host: normalize_host(authority.host()), - port: authority.port_u16(), + port, }) } @@ -821,7 +830,7 @@ fn origin_is_allowed(origin: &NormalizedOrigin, allowed_origins: &[String]) -> b host: o_host, port: o_port, }, - ) => a_scheme == o_scheme && a_host == o_host && (a_port.is_none() || a_port == o_port), + ) => a_scheme == o_scheme && a_host == o_host && a_port == o_port, _ => false, }) } @@ -884,21 +893,26 @@ fn validate_origin_header(headers: &HeaderMap, allowed_origins: &[String]) -> Ht if allowed_origins.is_empty() { return Ok(()); } - let Some(origin_header) = headers.get(http::header::ORIGIN) else { + let mut origin_headers = headers.get_all(http::header::ORIGIN).iter(); + let Some(origin_header) = origin_headers.next() else { return Ok(()); }; + if origin_headers.next().is_some() { + tracing::warn!("rejected request with multiple Origin headers"); + return Err(forbidden_response("Forbidden: Multiple Origin headers").into()); + } let origin_str = origin_header .to_str() .inspect_err(|_| { tracing::warn!(origin = ?origin_header, "rejected request with non-UTF-8 Origin header"); }) - .map_err(|_| bad_request_response("Bad Request: Invalid Origin header encoding"))?; + .map_err(|_| forbidden_response("Forbidden: Invalid Origin header encoding"))?; let origin = parse_origin_value(origin_str).ok_or_else(|| { tracing::warn!( origin = origin_str, "rejected request with malformed Origin header", ); - bad_request_response("Bad Request: Invalid Origin header") + forbidden_response("Forbidden: Invalid Origin header") })?; if !origin_is_allowed(&origin, allowed_origins) { tracing::warn!( diff --git a/crates/rmcp/src/transport/worker.rs b/crates/rmcp/src/transport/worker.rs index e32da7d7d..8924ecad0 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 { @@ -64,6 +68,9 @@ pub trait Worker: Sized + Send + 'static { fn config(&self) -> WorkerConfig { WorkerConfig::default() } + fn preserves_raw_responses() -> bool { + false + } /// Return true to send this message through the separate control queue. /// /// Workers that opt in must read [`WorkerContext::control_from_handler_rx`] @@ -167,7 +174,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 +221,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 +284,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"); @@ -293,6 +332,7 @@ pub struct SendRequest { #[non_exhaustive] pub struct WorkerContext { pub to_handler_tx: tokio::sync::mpsc::Sender>, + raw_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>, @@ -320,11 +360,33 @@ 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 { + 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) + .map_err(WorkerQuitReason::ResponseSerialization)?, + }) + } + JsonRpcMessage::Notification(notification) => { + JsonRpcMessage::Notification(notification) + } + JsonRpcMessage::Error(error) => JsonRpcMessage::Error(error), + }; + self.raw_to_handler_tx .send(item) .await .map_err(|_| WorkerQuitReason::HandlerTerminated) @@ -343,6 +405,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 +463,16 @@ impl Transport for WorkerTransport { } } async fn receive(&mut self) -> Option> { + loop { + match decode_peer_response::(self.rx.recv().await?) { + 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_headers.rs b/crates/rmcp/tests/test_custom_headers.rs index cb1018269..65ff152af 100644 --- a/crates/rmcp/tests/test_custom_headers.rs +++ b/crates/rmcp/tests/test_custom_headers.rs @@ -1199,6 +1199,61 @@ mod origin_validation { assert_eq!(response.status(), http::StatusCode::FORBIDDEN); } + #[tokio::test] + async fn malformed_origin_is_forbidden() { + let service = service_with_allowed_origins(&["http://localhost:8080"]); + let response = service.handle(init_request(Some("not an origin"))).await; + assert_eq!(response.status(), http::StatusCode::FORBIDDEN); + } + + #[tokio::test] + async fn multiple_origin_headers_are_forbidden() { + let service = service_with_allowed_origins(&["http://localhost:8080"]); + let mut request = init_request(Some("http://localhost:8080")); + request.headers_mut().append( + http::header::ORIGIN, + "http://attacker.example".parse().unwrap(), + ); + let response = service.handle(request).await; + assert_eq!(response.status(), http::StatusCode::FORBIDDEN); + } + + #[tokio::test] + async fn allowlisted_origin_does_not_wildcard_an_unexpected_port() { + let service = service_with_allowed_origins(&["https://app.example"]); + let response = service + .handle(init_request(Some("https://app.example:9443"))) + .await; + assert_eq!(response.status(), http::StatusCode::FORBIDDEN); + } + + #[tokio::test] + async fn default_origin_ports_are_equivalent() { + for (allowed, presented) in [ + ("http://localhost", "http://localhost:80"), + ("https://localhost:443", "https://localhost"), + ] { + let service = service_with_allowed_origins(&[allowed]); + let response = service.handle(init_request(Some(presented))).await; + assert_eq!(response.status(), http::StatusCode::OK); + } + } + + #[tokio::test] + async fn origin_with_non_origin_components_is_forbidden() { + for origin in [ + "http://localhost:8080/", + "http://localhost:8080/evil", + "http://localhost:8080?query=1", + "http://user@localhost:8080", + "http://localhost:8080#fragment", + ] { + let service = service_with_allowed_origins(&["http://localhost:8080"]); + let response = service.handle(init_request(Some(origin))).await; + assert_eq!(response.status(), http::StatusCode::FORBIDDEN, "{origin}"); + } + } + #[tokio::test] async fn missing_origin_passes_through() { let service = service_with_allowed_origins(&["http://localhost:8080"]); diff --git a/crates/rmcp/tests/test_custom_request.rs b/crates/rmcp/tests/test_custom_request.rs index 66ee1ff99..7f2efbcae 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,352 @@ 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?) +} + +#[cfg(all( + feature = "auth", + feature = "transport-streamable-http-server", + feature = "transport-streamable-http-client-reqwest" +))] +#[tokio::test] +async fn typed_skills_request_survives_oauth_http_wrapper() -> anyhow::Result<()> { + use rmcp::transport::{ + auth::{AuthClient, AuthorizationManager}, + streamable_http_client::{StreamableHttpClientTransportConfig, StreamableHttpClientWorker}, + streamable_http_server::{ + StreamableHttpServerConfig, StreamableHttpService, session::local::LocalSessionManager, + }, + }; + let stop = tokio_util::sync::CancellationToken::new(); + let server: StreamableHttpService = + StreamableHttpService::new( + || Ok(TypedCustomRequestServer), + Default::default(), + StreamableHttpServerConfig::default().with_cancellation_token(stop.child_token()), + ); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await?; + let endpoint = format!("http://{}/mcp", listener.local_addr()?); + let router = axum::Router::new().nest_service("/mcp", server); + let shutdown = stop.clone(); + let task = tokio::spawn(async move { + axum::serve(listener, router) + .with_graceful_shutdown(shutdown.cancelled_owned()) + .await + }); + let result = tokio::time::timeout(Duration::from_secs(10), async { + let manager = AuthorizationManager::new(&endpoint).await?; + let auth = AuthClient::new(reqwest::Client::new(), manager); + let transport = StreamableHttpClientWorker::new( + auth, + StreamableHttpClientTransportConfig::with_uri(endpoint), + ); + let client = ().serve(transport).await?; + let response = client + .send_request_as::(ClientRequest::CustomRequest(CustomRequest::new( + "skills/list", + Some(json!({})), + ))) + .await; + client.cancel().await?; + let response = response?; + assert_eq!(response.skills, ["example"]); + assert_eq!( + response.meta["io.modelcontextprotocol/serverInfo"]["name"], + "skills" + ); + anyhow::Ok(()) + }) + .await; + stop.cancel(); + tokio::time::timeout(Duration::from_secs(5), task).await???; + result??; + Ok(()) +} + +#[cfg(all( + feature = "auth", + feature = "transport-streamable-http-server", + feature = "transport-streamable-http-client-reqwest" +))] +#[tokio::test] +async fn json_tool_listing_populates_parameter_headers_through_auth_client() -> anyhow::Result<()> { + use rmcp::{ + ClientServiceExt, + model::{ + CallToolRequestParams, CallToolResponse, CallToolResult, ListToolsResult, + PaginatedRequestParams, ServerCapabilities, ServerInfo, Tool, + }, + service::RequestContext, + transport::{ + auth::{AuthClient, AuthorizationManager}, + streamable_http_client::{ + StreamableHttpClientTransportConfig, StreamableHttpClientWorker, + }, + streamable_http_server::{ + StreamableHttpServerConfig, StreamableHttpService, + session::local::LocalSessionManager, + }, + }, + }; + + struct AnnotatedToolServer; + impl ServerHandler for AnnotatedToolServer { + fn supported_protocol_versions( + &self, + ) -> std::borrow::Cow<'static, [rmcp::model::ProtocolVersion]> { + std::borrow::Cow::Borrowed(&[rmcp::model::ProtocolVersion::V_2026_07_28]) + } + + fn get_info(&self) -> ServerInfo { + ServerInfo::new(ServerCapabilities::builder().enable_tools().build()) + } + + fn get_tool(&self, name: &str) -> Option { + (name == "deploy").then(|| { + Tool::new( + "deploy", + "Deploy in a region", + Arc::new( + json!({"type": "object", "properties": { + "region": {"type": "string", "x-mcp-header": "Region"} + }}) + .as_object() + .unwrap() + .clone(), + ), + ) + }) + } + + async fn list_tools( + &self, + _: Option, + _: RequestContext, + ) -> Result { + Ok(ListToolsResult::with_all_items(vec![ + self.get_tool("deploy").unwrap(), + ])) + } + + async fn call_tool( + &self, + _: CallToolRequestParams, + _: RequestContext, + ) -> Result { + Ok(CallToolResult::success(vec![]).into()) + } + } + + let observed_headers = Arc::new(Mutex::new(Vec::new())); + let capture = observed_headers.clone(); + let stop = tokio_util::sync::CancellationToken::new(); + let server: StreamableHttpService = + StreamableHttpService::new( + || Ok(AnnotatedToolServer), + Default::default(), + StreamableHttpServerConfig::default() + .with_legacy_session_mode(false) + .with_json_response(true) + .with_cancellation_token(stop.child_token()), + ); + let router = axum::Router::new() + .nest_service("/mcp", server) + .layer(axum::middleware::from_fn( + move |request: axum::extract::Request, next: axum::middleware::Next| { + let capture = capture.clone(); + async move { + if request + .headers() + .get("Mcp-Method") + .is_some_and(|method| method == "tools/call") + { + capture + .lock() + .await + .push(request.headers().get("Mcp-Param-Region").cloned()); + } + next.run(request).await + } + }, + )); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await?; + let endpoint = format!("http://{}/mcp", listener.local_addr()?); + let shutdown = stop.clone(); + let task = tokio::spawn(async move { + axum::serve(listener, router) + .with_graceful_shutdown(shutdown.cancelled_owned()) + .await + }); + let result = tokio::time::timeout(Duration::from_secs(10), async { + let manager = AuthorizationManager::new(&endpoint).await?; + let transport = StreamableHttpClientWorker::new( + AuthClient::new(reqwest::Client::new(), manager), + StreamableHttpClientTransportConfig::with_uri(endpoint), + ); + let client = () + .serve_with_lifecycle( + transport, + rmcp::service::ClientLifecycleMode::Discover { + preferred_versions: vec![rmcp::model::ProtocolVersion::V_2026_07_28], + }, + ) + .await?; + let listed = client.list_tools(None).await?; + assert_eq!(listed.tools.len(), 1); + let response = client + .call_tool( + CallToolRequestParams::new("deploy") + .with_arguments(json!({"region": "us-east-1"}).as_object().unwrap().clone()), + ) + .await; + client.cancel().await?; + response?; + assert_eq!( + *observed_headers.lock().await, + vec![Some(http::HeaderValue::from_static("us-east-1"))] + ); + anyhow::Ok(()) + }) + .await; + stop.cancel(); + tokio::time::timeout(Duration::from_secs(5), task).await???; + result??; + Ok(()) +} + +#[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>>, From 68e6f4a1f808e5f0f4393ca3196857e9e0a14ab9 Mon Sep 17 00:00:00 2001 From: Jacob Magar Date: Sun, 13 Sep 2026 13:49:47 -0400 Subject: [PATCH 29/29] fix(rmcp): harden typed response compatibility --- crates/rmcp/CHANGELOG.md | 9 + crates/rmcp/src/service.rs | 81 +++++++-- crates/rmcp/src/transport.rs | 27 ++- .../src/transport/streamable_http_client.rs | 9 + .../rmcp/tests/support/typed_stdio_server.sh | 20 +++ crates/rmcp/tests/test_custom_headers.rs | 36 +++- crates/rmcp/tests/test_custom_request.rs | 155 +++++++++++++++++- crates/rmcp/tests/test_typed_child_process.rs | 54 ++++++ .../rmcp/tests/test_unix_socket_transport.rs | 98 +++++++++++ 9 files changed, 465 insertions(+), 24 deletions(-) create mode 100644 crates/rmcp/tests/support/typed_stdio_server.sh create mode 100644 crates/rmcp/tests/test_typed_child_process.rs diff --git a/crates/rmcp/CHANGELOG.md b/crates/rmcp/CHANGELOG.md index 5465cf210..545d81c39 100644 --- a/crates/rmcp/CHANGELOG.md +++ b/crates/rmcp/CHANGELOG.md @@ -7,6 +7,15 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] +### Added + +- add typed custom-response requests that preserve extension result fields across supported built-in transports + +### Fixed + +- reject malformed, duplicate, non-UTF-8, and disallowed HTTP Origin headers consistently +- compare HTTP origins using normalized effective ports instead of treating an omitted configured port as a wildcard + ## [3.3.0](https://github.com/modelcontextprotocol/rust-sdk/compare/rmcp-v3.2.0...rmcp-v3.3.0) - 2026-09-10 ### Added diff --git a/crates/rmcp/src/service.rs b/crates/rmcp/src/service.rs index 84c73912c..32c59ff15 100644 --- a/crates/rmcp/src/service.rs +++ b/crates/rmcp/src/service.rs @@ -281,6 +281,12 @@ pub type RxJsonRpcMessage = JsonRpcMessage< ::PeerResp, ::PeerNot, >; +/// A received JSON-RPC message whose response result remains raw JSON. +/// +/// Requests and notifications retain the peer types defined by `R`; only a +/// successful response's result is represented as [`serde_json::Value`]. This +/// prevents extension fields from being lost to the role's core response union +/// before a typed request can deserialize them. pub type RawRxJsonRpcMessage = JsonRpcMessage<::PeerReq, serde_json::Value, ::PeerNot>; @@ -465,7 +471,7 @@ impl> DynService for S { } use std::{ - collections::{HashMap, VecDeque}, + collections::HashMap, ops::Deref, sync::{Arc, atomic::AtomicU64}, time::Duration, @@ -549,11 +555,12 @@ type ProgressTimeoutWatchers = Arc = (mpsc::Sender, usize); type SubscriptionChannelMap = HashMap>; -/// A handle to a remote request -/// -/// You can cancel it by call [`RequestHandle::cancel`] with a reason, +/// A handle to a remote request whose response resolves to `T`. /// -/// or wait for response by call [`RequestHandle::await_response`] +/// `T` defaults to the role's core peer-response union. Typed extension +/// requests created by [`Peer::send_request_as_with_option`] instead use the +/// caller's concrete result type. Call [`RequestHandle::cancel`] to cancel the +/// request with a reason, or [`RequestHandle::await_response`] to await it. #[derive(Debug)] #[non_exhaustive] pub struct RequestHandle::PeerResp> { @@ -753,6 +760,15 @@ impl PendingResponder { Self::Typed(responder) => responder(Ok(value)), } } + + fn send_standard_result(self, value: R::PeerResp) { + match self { + Self::Standard(responder) => { + let _ = responder.send(Ok(value)); + } + Self::Typed(responder) => responder(Err(ServiceError::RawResponseUnavailable)), + } + } } #[derive(Debug)] @@ -914,6 +930,14 @@ impl Peer { /// concrete type avoids ambiguous `#[serde(untagged)]` union matching. /// The caller is responsible for pairing the request method with its /// correct result type. + /// + /// # Errors + /// + /// Returns [`ServiceError::RawResponseUnavailable`] without sending the + /// request when the selected transport cannot preserve raw result JSON. + /// Returns [`ServiceError::ResponseDeserialization`] when the received + /// result does not deserialize as `T`. Transport, peer, timeout, and + /// cancellation errors are returned unchanged. pub async fn send_request_as(&self, request: R::Req) -> Result where T: serde::de::DeserializeOwned + Send + 'static, @@ -926,6 +950,9 @@ impl Peer { /// Send a typed request with the same timeout, metadata, progress, and /// cancellation lifecycle available to core requests. + /// + /// Raw-response support and error behavior are the same as + /// [`Peer::send_request_as`]. pub async fn send_request_as_with_option( &self, request: R::Req, @@ -1468,7 +1495,6 @@ 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 send_task_set = tokio::task::JoinSet::::new(); let mut response_send_tasks = tokio::task::JoinSet::<()>::new(); #[derive(Debug)] @@ -1487,16 +1513,14 @@ where enum Event { ProxyMessage(PeerSinkMessage), PeerMessage(RawRxJsonRpcMessage), + LegacyPeerMessage(RxJsonRpcMessage), ToSink(TxJsonRpcMessage), SendTaskResult(SendTaskResult), ResponseSendTaskResult(Result<(), tokio::task::JoinError>), } let quit_reason = loop { - let evt = if let Some(m) = batch_messages.pop_front() { - Event::PeerMessage(m) - } else { - tokio::select! { + let evt = tokio::select! { m = sink_proxy_rx.recv(), if !sink_proxy_rx.is_closed() => { if let Some(m) = m { Event::ToSink(m) @@ -1504,9 +1528,15 @@ where continue } } - m = transport.receive_raw() => { - if let Some(m) = m { - Event::PeerMessage(m) + m = async { + if T::preserves_raw_responses() { + transport.receive_raw().await.map(Event::PeerMessage) + } else { + transport.receive().await.map(Event::LegacyPeerMessage) + } + } => { + if let Some(event) = m { + event } else { // input stream closed tracing::info!("input stream terminated"); @@ -1544,7 +1574,31 @@ where tracing::info!("task cancelled"); break QuitReason::Cancelled } + }; + + let evt = match evt { + Event::LegacyPeerMessage(JsonRpcMessage::Response(JsonRpcResponse { + result, + id, + .. + })) => { + if let Some(responder) = + remove_pending_request(&mut local_responder_pool, &id) + { + responder.send_standard_result(result); + } + continue; + } + Event::LegacyPeerMessage(JsonRpcMessage::Request(request)) => { + Event::PeerMessage(JsonRpcMessage::Request(request)) + } + Event::LegacyPeerMessage(JsonRpcMessage::Notification(notification)) => { + Event::PeerMessage(JsonRpcMessage::Notification(notification)) + } + Event::LegacyPeerMessage(JsonRpcMessage::Error(error)) => { + Event::PeerMessage(JsonRpcMessage::Error(error)) } + event => event, }; tracing::trace!(?evt, "new event"); @@ -1825,6 +1879,7 @@ where responder.send_error(service_error); } } + Event::LegacyPeerMessage(_) => unreachable!("legacy messages are normalized above"), } }; diff --git a/crates/rmcp/src/transport.rs b/crates/rmcp/src/transport.rs index 31c23b551..bf45d287b 100644 --- a/crates/rmcp/src/transport.rs +++ b/crates/rmcp/src/transport.rs @@ -147,19 +147,32 @@ where /// 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. + /// Typed extension requests require this capability. Returning `true` is a + /// contract that [`Self::receive_raw`] returns the original JSON result for + /// every response path without first decoding it through the role response + /// union. A transport that cannot satisfy that contract must retain the + /// default `false`; typed requests then fail with + /// [`crate::service::ServiceError::RawResponseUnavailable`] before they are + /// sent. + /// + /// The built-in async-read/write, child-process, reqwest HTTP, Unix-socket + /// HTTP transports preserve raw responses. Authenticated HTTP wrappers + /// forward the wrapped backend's capability, and a `WorkerTransport` + /// forwards the capability declared by its worker. + /// Existing custom transports default to `false` for source and behavior + /// compatibility. 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. + /// A transport that returns `true` from [`Self::preserves_raw_responses`] + /// must override this method and retain raw results on every response path. + /// The default adapts the existing typed receive API for custom transports; + /// it does not make typed extension requests available because an extension + /// result already decoded through the role 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 { diff --git a/crates/rmcp/src/transport/streamable_http_client.rs b/crates/rmcp/src/transport/streamable_http_client.rs index 2f4d21fc9..e4e205a6d 100644 --- a/crates/rmcp/src/transport/streamable_http_client.rs +++ b/crates/rmcp/src/transport/streamable_http_client.rs @@ -422,6 +422,15 @@ pub trait StreamableHttpClient: Clone + Send + 'static { type Error: std::error::Error + Send + Sync + 'static; /// Whether this backend preserves response result bodies as raw JSON. + /// + /// Returning `true` is a contract that all JSON responses use + /// [`StreamableHttpPostResponse::RawJson`] and that every SSE response path + /// yields events whose JSON-RPC result has not first been decoded through + /// [`ServerResult`]. Backends that cannot satisfy both requirements must + /// retain the default `false`; typed requests then fail before being sent. + /// The built-in reqwest and Unix-socket backends satisfy this contract. + /// The authenticated HTTP wrapper forwards the wrapped backend's + /// capability. fn preserves_raw_responses() -> bool { false } diff --git a/crates/rmcp/tests/support/typed_stdio_server.sh b/crates/rmcp/tests/support/typed_stdio_server.sh new file mode 100644 index 000000000..85850380d --- /dev/null +++ b/crates/rmcp/tests/support/typed_stdio_server.sh @@ -0,0 +1,20 @@ +#!/bin/sh + +# Minimal line-delimited MCP server used to exercise TokioChildProcess without +# relying on a language-specific MCP implementation or package installation. +while IFS= read -r message; do + request_id=$(printf '%s\n' "$message" | sed -n 's/.*"id":\([^,}]*\).*/\1/p') + case "$message" in + *'"method":"initialize"'*) + printf '%s\n' "{\"jsonrpc\":\"2.0\",\"id\":${request_id},\"result\":{\"protocolVersion\":\"2025-03-26\",\"capabilities\":{},\"serverInfo\":{\"name\":\"typed-stdio-fixture\",\"version\":\"1.0.0\"}}}" + ;; + *'"method":"skills/list"'*) + printf '%s\n' "{\"jsonrpc\":\"2.0\",\"id\":${request_id},\"result\":{\"resultType\":\"complete\",\"skills\":[\"stdio-example\"],\"_meta\":{\"vendorExtension\":{\"retained\":true}}}}" + ;; + *'"method":"ping"'*) + printf '%s\n' "{\"jsonrpc\":\"2.0\",\"id\":${request_id},\"result\":{}}" + ;; + *'"method":"notifications/initialized"'*) + ;; + esac +done diff --git a/crates/rmcp/tests/test_custom_headers.rs b/crates/rmcp/tests/test_custom_headers.rs index 65ff152af..250f97dd4 100644 --- a/crates/rmcp/tests/test_custom_headers.rs +++ b/crates/rmcp/tests/test_custom_headers.rs @@ -1124,7 +1124,10 @@ async fn test_server_falls_back_to_uri_authority_when_host_header_missing() { #[cfg(all(feature = "transport-streamable-http-server", feature = "server"))] mod origin_validation { - use std::sync::Arc; + use std::sync::{ + Arc, + atomic::{AtomicUsize, Ordering}, + }; use bytes::Bytes; use http::{Method, Request, header::CONTENT_TYPE}; @@ -1149,9 +1152,19 @@ mod origin_validation { fn service_with_allowed_origins( origins: &[&str], + ) -> StreamableHttpService { + service_with_allowed_origins_and_counter(origins, Arc::new(AtomicUsize::new(0))) + } + + fn service_with_allowed_origins_and_counter( + origins: &[&str], + handler_creations: Arc, ) -> StreamableHttpService { StreamableHttpService::new( - || Ok(TestHandler), + move || { + handler_creations.fetch_add(1, Ordering::SeqCst); + Ok(TestHandler) + }, Arc::new(LocalSessionManager::default()), StreamableHttpServerConfig::default().with_allowed_origins(origins.iter().copied()), ) @@ -1206,6 +1219,25 @@ mod origin_validation { assert_eq!(response.status(), http::StatusCode::FORBIDDEN); } + #[tokio::test] + async fn non_utf8_origin_is_forbidden_before_handler_creation() { + let handler_creations = Arc::new(AtomicUsize::new(0)); + let service = service_with_allowed_origins_and_counter( + &["http://localhost:8080"], + handler_creations.clone(), + ); + let mut request = init_request(None); + request.headers_mut().insert( + http::header::ORIGIN, + http::HeaderValue::from_bytes(b"\xff").expect("opaque header value"), + ); + + let response = service.handle(request).await; + + assert_eq!(response.status(), http::StatusCode::FORBIDDEN); + assert_eq!(handler_creations.load(Ordering::SeqCst), 0); + } + #[tokio::test] async fn multiple_origin_headers_are_forbidden() { let service = service_with_allowed_origins(&["http://localhost:8080"]); diff --git a/crates/rmcp/tests/test_custom_request.rs b/crates/rmcp/tests/test_custom_request.rs index 7f2efbcae..b4056b5ff 100644 --- a/crates/rmcp/tests/test_custom_request.rs +++ b/crates/rmcp/tests/test_custom_request.rs @@ -1,5 +1,12 @@ #![cfg(not(feature = "local"))] -use std::{future::Future, sync::Arc, time::Duration}; +use std::{ + future::Future, + sync::{ + Arc, + atomic::{AtomicUsize, Ordering}, + }, + time::Duration, +}; use rmcp::{ ClientHandler, RoleClient, ServerHandler, ServiceExt, @@ -7,7 +14,7 @@ use rmcp::{ ClientRequest, ClientResult, CustomRequest, CustomResult, ErrorCode, ErrorData, PingRequest, ServerRequest, ServerResult, }, - service::{PeerRequestOptions, ServiceError}, + service::{PeerRequestOptions, ServiceError, serve_directly}, transport::Transport, }; use serde::Deserialize; @@ -44,6 +51,83 @@ async fn existing_transport_implementations_get_the_raw_receive_compatibility_de assert!(transport.receive_raw().await.is_none()); } +struct LegacyResponseTransport { + responses: tokio::sync::mpsc::UnboundedReceiver, + response_tx: tokio::sync::mpsc::UnboundedSender, + sends: Arc, +} + +impl Transport for LegacyResponseTransport { + type Error = std::convert::Infallible; + + fn send( + &mut self, + item: rmcp::service::TxJsonRpcMessage, + ) -> impl Future> + Send + 'static { + let response_tx = self.response_tx.clone(); + let sends = self.sends.clone(); + async move { + if let rmcp::model::ClientJsonRpcMessage::Request(request) = item { + sends.fetch_add(1, Ordering::SeqCst); + let _ = response_tx.send(rmcp::model::ServerJsonRpcMessage::response( + ServerResult::CustomResult(CustomResult::new(json!({}))), + request.id, + )); + } + Ok(()) + } + } + + async fn receive(&mut self) -> Option> { + self.responses.recv().await + } + + fn close(&mut self) -> impl Future> + Send { + std::future::ready(Ok(())) + } +} + +#[tokio::test] +async fn legacy_transport_preserves_standard_response_without_raw_round_trip() -> anyhow::Result<()> +{ + let (response_tx, responses) = tokio::sync::mpsc::unbounded_channel(); + let sends = Arc::new(AtomicUsize::new(0)); + let transport = LegacyResponseTransport { + responses, + response_tx, + sends: sends.clone(), + }; + let client = serve_directly::( + (), + transport, + Some(rmcp::model::ServerInfo::default().into()), + ); + + let typed = client + .send_request_as::(ClientRequest::CustomRequest(CustomRequest::new( + "skills/list", + None, + ))) + .await; + assert!(matches!(typed, Err(ServiceError::RawResponseUnavailable))); + assert_eq!(sends.load(Ordering::SeqCst), 0, "typed request was sent"); + + let standard = client + .send_request(ClientRequest::CustomRequest(CustomRequest::new( + "requests/custom-test", + None, + ))) + .await?; + assert!(matches!( + standard, + ServerResult::CustomResult(result) if result.0 == json!({}) + )); + assert_eq!(sends.load(Ordering::SeqCst), 1); + + client.cancel().await?; + Ok(()) +} + struct CustomRequestServer { receive_signal: Arc, payload: Arc>>, @@ -463,6 +547,73 @@ async fn typed_custom_requests_use_standard_timeout_and_cancellation() -> anyhow Ok(()) } +struct CancellationObservingServer { + started: Arc, + cancelled: Arc, +} + +impl ServerHandler for CancellationObservingServer { + async fn on_custom_request( + &self, + request: CustomRequest, + context: rmcp::service::RequestContext, + ) -> Result { + if request.method == "skills/list" { + return Ok(CustomResult::new(json!({ + "resultType": "complete", + "skills": ["after-cancellation"], + "_meta": {} + }))); + } + self.started.notify_one(); + context.ct.cancelled().await; + self.cancelled.notify_one(); + Ok(CustomResult::new(json!({"cancelled": true}))) + } +} + +#[tokio::test] +async fn typed_custom_request_handle_sends_cancellation_to_the_peer() -> anyhow::Result<()> { + let started = Arc::new(Notify::new()); + let cancelled = Arc::new(Notify::new()); + let (server_transport, client_transport) = tokio::io::duplex(4096); + let server_task = tokio::spawn({ + let started = started.clone(); + let cancelled = cancelled.clone(); + async move { + CancellationObservingServer { started, cancelled } + .serve(server_transport) + .await? + .waiting() + .await?; + anyhow::Ok(()) + } + }); + let client = ().serve(client_transport).await?; + let handle = client + .send_request_as_with_option::( + ClientRequest::CustomRequest(CustomRequest::new("skills/cancel", None)), + PeerRequestOptions::no_options(), + ) + .await?; + + tokio::time::timeout(Duration::from_secs(2), started.notified()).await?; + handle.cancel(Some("test cancellation".into())).await?; + tokio::time::timeout(Duration::from_secs(2), cancelled.notified()).await?; + + let following: SkillsListResult = client + .send_request_as(ClientRequest::CustomRequest(CustomRequest::new( + "skills/list", + None, + ))) + .await?; + assert_eq!(following.skills, ["after-cancellation"]); + + client.cancel().await?; + tokio::time::timeout(Duration::from_secs(2), server_task).await???; + Ok(()) +} + struct CustomRequestClient { receive_signal: Arc, payload: Arc>>, diff --git a/crates/rmcp/tests/test_typed_child_process.rs b/crates/rmcp/tests/test_typed_child_process.rs new file mode 100644 index 000000000..9ff26b816 --- /dev/null +++ b/crates/rmcp/tests/test_typed_child_process.rs @@ -0,0 +1,54 @@ +#![cfg(all( + unix, + not(feature = "local"), + feature = "transport-child-process", + feature = "client" +))] + +use rmcp::{ + ServiceExt, + model::{ClientRequest, CustomRequest, PingRequest, ServerResult}, + transport::{ConfigureCommandExt, TokioChildProcess}, +}; +use serde::Deserialize; + +#[derive(Debug, Deserialize, PartialEq)] +#[serde(rename_all = "camelCase")] +struct SkillsListResult { + result_type: String, + skills: Vec, + #[serde(rename = "_meta")] + meta: serde_json::Value, +} + +#[tokio::test] +async fn typed_extension_response_survives_child_process_and_preserves_correlation() +-> anyhow::Result<()> { + let fixture = concat!( + env!("CARGO_MANIFEST_DIR"), + "/tests/support/typed_stdio_server.sh" + ); + let transport = + TokioChildProcess::new(tokio::process::Command::new("sh").configure(|command| { + command.arg(fixture); + }))?; + let client = ().serve(transport).await?; + + let response: SkillsListResult = client + .send_request_as(ClientRequest::CustomRequest(CustomRequest::new( + "skills/list", + None, + ))) + .await?; + assert_eq!(response.result_type, "complete"); + assert_eq!(response.skills, ["stdio-example"]); + assert_eq!(response.meta["vendorExtension"]["retained"], true); + + let following = client + .send_request(ClientRequest::PingRequest(PingRequest::default())) + .await?; + assert!(matches!(following, ServerResult::EmptyResult(_))); + + client.cancel().await?; + Ok(()) +} diff --git a/crates/rmcp/tests/test_unix_socket_transport.rs b/crates/rmcp/tests/test_unix_socket_transport.rs index d0e5397a3..eaaf2d858 100644 --- a/crates/rmcp/tests/test_unix_socket_transport.rs +++ b/crates/rmcp/tests/test_unix_socket_transport.rs @@ -13,11 +13,13 @@ use http::{HeaderName, HeaderValue}; use hyper_util::rt::TokioIo; use rmcp::{ ServiceExt, + model::{ClientRequest, CustomRequest}, transport::{ StreamableHttpClientTransport, UnixSocketHttpClient, streamable_http_client::StreamableHttpClientTransportConfig, }, }; +use serde::Deserialize; use serde_json::json; use tokio::sync::Mutex; @@ -86,6 +88,27 @@ async fn mcp_handler( ], String::new(), ); + } else if method == "skills/list" { + let response = json!({ + "jsonrpc": "2.0", + "id": json_body.get("id"), + "result": { + "resultType": "complete", + "skills": ["unix-example"], + "_meta": {"vendorExtension": {"retained": true}} + } + }); + return ( + StatusCode::OK, + [ + (http::header::CONTENT_TYPE, "application/json"), + ( + http::HeaderName::from_static("mcp-session-id"), + "unix-test-session", + ), + ], + response.to_string(), + ); } } @@ -111,6 +134,81 @@ async fn mcp_handler( ) } +#[derive(Debug, Deserialize, PartialEq)] +#[serde(rename_all = "camelCase")] +struct SkillsListResult { + result_type: String, + skills: Vec, + #[serde(rename = "_meta")] + meta: serde_json::Value, +} + +struct TemporarySocketDirectory(std::path::PathBuf); + +impl TemporarySocketDirectory { + fn new() -> std::io::Result { + // Keep the pathname below the small sockaddr_un limit on macOS; its + // resolved system temporary directory can itself be very long. + let path = std::path::Path::new("/tmp").join(format!("rmcp-{}", uuid::Uuid::new_v4())); + std::fs::create_dir(&path)?; + Ok(Self(path)) + } +} + +impl Drop for TemporarySocketDirectory { + fn drop(&mut self) { + let _ = std::fs::remove_dir_all(&self.0); + } +} + +struct AbortServerOnDrop(tokio::task::JoinHandle<()>); + +impl Drop for AbortServerOnDrop { + fn drop(&mut self) { + self.0.abort(); + } +} + +/// Typed extension results must retain fields that the core response union does +/// not know about when transported over a Unix-domain HTTP connection. +#[tokio::test] +async fn test_unix_socket_typed_custom_response_preserves_extension_fields() -> anyhow::Result<()> { + let dir = TemporarySocketDirectory::new()?; + let socket_path = dir.0.join("mcp.sock"); + + let state = ServerState { + received_headers: Arc::new(Mutex::new(HashMap::new())), + initialize_called: Arc::new(tokio::sync::Notify::new()), + }; + let app = Router::new() + .route("/mcp", post(mcp_handler)) + .with_state(state); + let listener = tokio::net::UnixListener::bind(&socket_path)?; + let _server_guard = AbortServerOnDrop(spawn_unix_server(listener, app)); + + let socket_str = socket_path.to_str().expect("UTF-8 temporary path"); + let uri = "http://mcp-server.internal/mcp"; + let transport = StreamableHttpClientTransport::with_client( + UnixSocketHttpClient::new(socket_str, uri), + StreamableHttpClientTransportConfig::with_uri(uri), + ); + let client = ().serve(transport).await?; + + let response: SkillsListResult = client + .send_request_as(ClientRequest::CustomRequest(CustomRequest::new( + "skills/list", + None, + ))) + .await?; + + assert_eq!(response.result_type, "complete"); + assert_eq!(response.skills, ["unix-example"]); + assert_eq!(response.meta["vendorExtension"]["retained"], true); + + client.cancel().await?; + Ok(()) +} + /// Spawns an HTTP/1.1 server on a Unix socket using hyper directly. /// Avoids `axum::serve(UnixListener, ...)` which uses `spawn_local` on Linux. fn spawn_unix_server(