diff --git a/crates/rmcp/src/service/client.rs b/crates/rmcp/src/service/client.rs index 40d410d5..f71f2fb0 100644 --- a/crates/rmcp/src/service/client.rs +++ b/crates/rmcp/src/service/client.rs @@ -102,13 +102,18 @@ where .ok_or_else(|| ClientInitializeError::ConnectionClosed(context.to_string())) } +enum StartupResponse { + Response(Box, RequestId), + Error(ErrorData, Option), +} + /// Helper function to expect a response from the stream async fn expect_response( transport: &mut T, context: &str, service: &S, peer: Peer, -) -> Result<(ServerResult, RequestId), ClientInitializeError> +) -> Result where T: Transport, S: Service, @@ -118,11 +123,11 @@ where match message { // Expected message to complete the initialization ServerJsonRpcMessage::Response(JsonRpcResponse { id, result, .. }) => { - break Ok((result, id)); + break Ok(StartupResponse::Response(Box::new(result), id)); } // Handle JSON-RPC error responses ServerJsonRpcMessage::Error(error) => { - break Err(ClientInitializeError::JsonRpcError(error.error)); + break Ok(StartupResponse::Error(error.error, error.id)); } // Server could send logging messages before handshake ServerJsonRpcMessage::Notification(mut notification) => { @@ -550,6 +555,209 @@ pub enum ClientLifecycleMode { }, } +pub(crate) const DISCOVER_PROBE_HTTP_STATUS_KEY: &str = + "io.modelcontextprotocol/rmcp/discoverProbeHttpStatus"; + +#[derive(Debug)] +struct DiscoverStartupError { + error: ClientInitializeError, + failure: DiscoverProbeFailure, +} + +#[derive(Debug)] +enum DiscoverProbeFailure { + Other, + IncompatibleVersions(Vec), + HttpStatus(u16), + Rpc { + error: ErrorData, + requested_version: ProtocolVersion, + request_id: RequestId, + response_id: Option, + }, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum ProbeDisposition { + RetryLegacy, + Fail, +} + +fn classify_probe_failure(failure: &DiscoverProbeFailure) -> ProbeDisposition { + match failure { + DiscoverProbeFailure::IncompatibleVersions(versions) + if exclusively_historical_protocol_versions(versions) => + { + ProbeDisposition::RetryLegacy + } + DiscoverProbeFailure::HttpStatus(404 | 405) => ProbeDisposition::RetryLegacy, + DiscoverProbeFailure::Rpc { + error, + requested_version, + request_id, + response_id, + } if rpc_error_proves_legacy( + error, + requested_version, + request_id, + response_id.as_ref(), + ) => + { + ProbeDisposition::RetryLegacy + } + _ => ProbeDisposition::Fail, + } +} + +fn rpc_error_proves_legacy( + error: &ErrorData, + requested_version: &ProtocolVersion, + request_id: &RequestId, + response_id: Option<&RequestId>, +) -> bool { + let correlated = + response_id.is_some_and(|response_id| request_id.matches_response_id(response_id)); + + if error.code == crate::model::ErrorCode::METHOD_NOT_FOUND { + return correlated; + } + + if error.code != crate::model::ErrorCode(-32000) || (!correlated && response_id.is_some()) { + return false; + } + + let message = error.message.trim().to_ascii_lowercase(); + let normalized_message = message.strip_prefix("bad request: ").unwrap_or(&message); + if normalized_message == "no valid session id provided" { + return true; + } + + if !normalized_message.contains("unsupported protocol version") + || !normalized_message.contains(requested_version.as_str()) + { + return false; + } + + let supported_data = error.data.as_ref().and_then(|data| data.get("supported")); + let supported_from_data = match supported_data { + Some(value) => { + let Ok(versions) = serde_json::from_value::>(value.clone()) else { + return false; + }; + Some(versions) + } + None => None, + }; + let supported_from_message = historical_versions_from_message(normalized_message); + if normalized_message.contains("supported versions:") && supported_from_message.is_none() { + return false; + } + + let supported = match (supported_from_data, supported_from_message) { + (Some(data), Some(message)) + if data.len() == message.len() + && data.iter().all(|version| message.contains(version)) => + { + Some(data) + } + (Some(_), Some(_)) => None, + (Some(versions), None) | (None, Some(versions)) => Some(versions), + (None, None) => None, + }; + + supported + .as_deref() + .is_some_and(exclusively_historical_protocol_versions) +} + +fn internal_probe_http_status(error: &ErrorData) -> Option { + error + .data + .as_ref()? + .get(DISCOVER_PROBE_HTTP_STATUS_KEY)? + .as_u64() + .and_then(|status| u16::try_from(status).ok()) +} + +impl From for DiscoverStartupError { + fn from(error: ClientInitializeError) -> Self { + let failure = match &error { + ClientInitializeError::NoCompatibleProtocolVersion { + server_supported, .. + } => DiscoverProbeFailure::IncompatibleVersions(server_supported.clone()), + _ => DiscoverProbeFailure::Other, + }; + Self { error, failure } + } +} + +impl DiscoverStartupError { + fn json_rpc( + error: ErrorData, + requested_version: ProtocolVersion, + request_id: RequestId, + response_id: Option, + ) -> Self { + let failure = if let Some(status) = internal_probe_http_status(&error) { + DiscoverProbeFailure::HttpStatus(status) + } else { + DiscoverProbeFailure::Rpc { + error: error.clone(), + requested_version, + request_id, + response_id, + } + }; + Self { + error: ClientInitializeError::JsonRpcError(error), + failure, + } + } + + fn disposition(&self) -> ProbeDisposition { + classify_probe_failure(&self.failure) + } +} + +fn exclusively_historical_protocol_versions(versions: &[ProtocolVersion]) -> bool { + !versions.is_empty() + && versions.iter().all(|version| { + (ProtocolVersion::KNOWN_VERSIONS.contains(version) + && version < &ProtocolVersion::V_2026_07_28) + // Some deployed legacy servers also advertise this pre-release version. + || version.as_str() == "2024-10-07" + }) +} + +fn historical_versions_from_message(message: &str) -> Option> { + let (_, supported_versions) = message.split_once("supported versions:")?; + let supported_versions = supported_versions.split(')').next()?; + let mut versions = Vec::new(); + for candidate in supported_versions.split(',') { + let candidate = candidate + .trim() + .trim_matches(|character| matches!(character, '[' | ']' | '"' | '\'')); + let bytes = candidate.as_bytes(); + if bytes.len() != 10 + || bytes.get(4) != Some(&b'-') + || bytes.get(7) != Some(&b'-') + || bytes + .iter() + .enumerate() + .any(|(index, byte)| index != 4 && index != 7 && !byte.is_ascii_digit()) + { + return None; + } + let version = serde_json::from_value::(serde_json::Value::String( + candidate.to_owned(), + )) + .ok()?; + versions.push(version); + } + + (!versions.is_empty()).then_some(versions) +} + /// Client-specific lifecycle entry points. pub trait ClientServiceExt: Service + Sized { fn serve_with_lifecycle( @@ -676,7 +884,8 @@ where &client_info, preferred_versions, ) - .await?; + .await + .map_err(|error| error.error)?; } ClientLifecycleMode::Auto { preferred_versions, @@ -693,9 +902,7 @@ where .await; match discover_result { Ok(()) => {} - Err(ClientInitializeError::JsonRpcError(error)) - if error.code == crate::model::ErrorCode::METHOD_NOT_FOUND => - { + Err(error) if error.disposition() == ProbeDisposition::RetryLegacy => { let mut legacy_info = client_info; if let Some(version) = legacy_version { legacy_info.protocol_version = version; @@ -703,7 +910,7 @@ where legacy_startup(&service, &mut transport, &id_provider, &peer, legacy_info) .await?; } - Err(error) => return Err(error), + Err(error) => return Err(error.error), } } } @@ -739,7 +946,12 @@ where })?; let (response, response_id) = - expect_response(transport, "initialize response", service, peer.clone()).await?; + match expect_response(transport, "initialize response", service, peer.clone()).await? { + StartupResponse::Response(response, response_id) => (*response, response_id), + StartupResponse::Error(error, _) => { + return Err(ClientInitializeError::JsonRpcError(error)); + } + }; if !id.matches_response_id(&response_id) { return Err(ClientInitializeError::ConflictInitResponseId( @@ -773,13 +985,13 @@ async fn discover_startup( peer: &Peer, client_info: &ClientInfo, preferred_versions: Vec, -) -> Result<(), ClientInitializeError> +) -> Result<(), DiscoverStartupError> where S: Service, T: Transport + 'static, { if preferred_versions.is_empty() { - return Err(ClientInitializeError::NoPreferredProtocolVersion); + return Err(ClientInitializeError::NoPreferredProtocolVersion.into()); } let mut attempted = Vec::new(); @@ -805,13 +1017,20 @@ where ClientInitializeError::transport::(error, "send discover request") })?; - match expect_response(transport, "discover response", service, peer.clone()).await { - Ok((ServerResult::DiscoverResult(result), response_id)) => { + match expect_response(transport, "discover response", service, peer.clone()).await? { + StartupResponse::Response(response, response_id) => { + let result = match *response { + ServerResult::DiscoverResult(result) => result, + response => { + return Err( + ClientInitializeError::ExpectedInitResult(Some(response)).into() + ); + } + }; if !id.matches_response_id(&response_id) { - return Err(ClientInitializeError::ConflictInitResponseId( - id, - response_id, - )); + return Err( + ClientInitializeError::ConflictInitResponseId(id, response_id).into(), + ); } let Some(selected) = select_protocol_version(&preferred_versions, &result.supported_versions) @@ -819,7 +1038,8 @@ where return Err(ClientInitializeError::NoCompatibleProtocolVersion { client_supported: preferred_versions, server_supported: result.supported_versions, - }); + } + .into()); }; peer.set_peer_info(ServerInfo { protocol_version: selected.clone(), @@ -835,12 +1055,28 @@ where }); return Ok(()); } - Ok((response, _)) => { - return Err(ClientInitializeError::ExpectedInitResult(Some(response))); - } - Err(ClientInitializeError::JsonRpcError(error)) - if error.code == crate::model::ErrorCode::UNSUPPORTED_PROTOCOL_VERSION => - { + StartupResponse::Error(error, response_id) => { + if let Some(response_id) = response_id.as_ref() + && !id.matches_response_id(response_id) + { + return Err(ClientInitializeError::ConflictInitResponseId( + id, + response_id.clone(), + ) + .into()); + } + + if error.code != crate::model::ErrorCode::UNSUPPORTED_PROTOCOL_VERSION + || response_id.is_none() + { + return Err(DiscoverStartupError::json_rpc( + error, + candidate, + id, + response_id, + )); + } + let supported = error .data .as_ref() @@ -865,11 +1101,11 @@ where return Err(ClientInitializeError::NoCompatibleProtocolVersion { client_supported: preferred_versions, server_supported: supported, - }); + } + .into()); }; candidate = next; } - Err(error) => return Err(error), } } } @@ -2041,6 +2277,101 @@ where mod tests { use super::*; + fn rpc_probe_failure( + code: crate::model::ErrorCode, + message: &str, + response_id: Option, + ) -> DiscoverProbeFailure { + DiscoverProbeFailure::Rpc { + error: ErrorData::new(code, message.to_owned(), None), + requested_version: ProtocolVersion::V_2026_07_28, + request_id: RequestId::Number(7), + response_id, + } + } + + #[test] + fn probe_classifier_has_explicit_http_and_version_policy() { + let historical = DiscoverProbeFailure::IncompatibleVersions(vec![ + ProtocolVersion::V_2025_11_25, + ProtocolVersion::V_2025_06_18, + ]); + let future = DiscoverProbeFailure::IncompatibleVersions(vec![ + serde_json::from_value(serde_json::json!("2099-01-01")).unwrap(), + ]); + + assert_eq!( + classify_probe_failure(&historical), + ProbeDisposition::RetryLegacy + ); + assert_eq!(classify_probe_failure(&future), ProbeDisposition::Fail); + assert_eq!( + classify_probe_failure(&DiscoverProbeFailure::HttpStatus(404)), + ProbeDisposition::RetryLegacy + ); + assert_eq!( + classify_probe_failure(&DiscoverProbeFailure::HttpStatus(405)), + ProbeDisposition::RetryLegacy + ); + assert_eq!( + classify_probe_failure(&DiscoverProbeFailure::HttpStatus(401)), + ProbeDisposition::Fail + ); + } + + #[test] + fn probe_classifier_requires_correlation_for_method_not_found() { + let correlated = rpc_probe_failure( + crate::model::ErrorCode::METHOD_NOT_FOUND, + "Method not found", + Some(RequestId::Number(7)), + ); + let uncorrelated = rpc_probe_failure( + crate::model::ErrorCode::METHOD_NOT_FOUND, + "Method not found", + None, + ); + + assert_eq!( + classify_probe_failure(&correlated), + ProbeDisposition::RetryLegacy + ); + assert_eq!( + classify_probe_failure(&uncorrelated), + ProbeDisposition::Fail + ); + } + + #[test] + fn probe_classifier_accepts_only_known_null_id_prevalidation_errors() { + let missing_session = rpc_probe_failure( + crate::model::ErrorCode(-32000), + "Bad Request: No valid session ID provided", + None, + ); + let historical_versions = rpc_probe_failure( + crate::model::ErrorCode(-32000), + "Bad Request: Unsupported protocol version: 2026-07-28 \ + (supported versions: 2025-11-25, 2025-06-18)", + None, + ); + let arbitrary = rpc_probe_failure( + crate::model::ErrorCode(-32000), + "Internal gateway rejection", + None, + ); + + assert_eq!( + classify_probe_failure(&missing_session), + ProbeDisposition::RetryLegacy + ); + assert_eq!( + classify_probe_failure(&historical_versions), + ProbeDisposition::RetryLegacy + ); + assert_eq!(classify_probe_failure(&arbitrary), ProbeDisposition::Fail); + } + fn disconnected_peer() -> Peer { let (peer, receiver) = Peer::::new(Arc::new(AtomicU32RequestIdProvider::default()), None); 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 e2eeebb4..84bd7f5d 100644 --- a/crates/rmcp/src/transport/common/reqwest/streamable_http_client.rs +++ b/crates/rmcp/src/transport/common/reqwest/streamable_http_client.rs @@ -219,6 +219,22 @@ impl StreamableHttpClient for reqwest::Client { if status == reqwest::StatusCode::NOT_FOUND && session_was_attached { return Err(StreamableHttpError::SessionExpired); } + if matches!( + status, + reqwest::StatusCode::UNAUTHORIZED + | reqwest::StatusCode::FORBIDDEN + | reqwest::StatusCode::NOT_FOUND + | reqwest::StatusCode::METHOD_NOT_ALLOWED + ) { + let body = response + .text() + .await + .unwrap_or_else(|_| "".to_owned()); + return Err(StreamableHttpError::UnexpectedHttpStatus { + status: status.as_u16(), + body: Cow::Owned(body), + }); + } let content_type = response .headers() .get(reqwest::header::CONTENT_TYPE) @@ -262,9 +278,10 @@ impl StreamableHttpClient for reqwest::Client { ), } } - return Err(StreamableHttpError::UnexpectedServerResponse(Cow::Owned( - format!("HTTP {status}: {body}"), - ))); + return Err(StreamableHttpError::UnexpectedHttpStatus { + status: status.as_u16(), + body: Cow::Owned(body), + }); } match content_type.as_deref() { Some(ct) if ct.as_bytes().starts_with(EVENT_STREAM_MIME_TYPE.as_bytes()) => { diff --git a/crates/rmcp/src/transport/common/unix_socket.rs b/crates/rmcp/src/transport/common/unix_socket.rs index 5f995db2..8c94c989 100644 --- a/crates/rmcp/src/transport/common/unix_socket.rs +++ b/crates/rmcp/src/transport/common/unix_socket.rs @@ -274,9 +274,10 @@ impl StreamableHttpClient for UnixSocketHttpClient { .await .map(|c| String::from_utf8_lossy(&c.to_bytes()).into_owned()) .unwrap_or_else(|_| "".to_owned()); - return Err(StreamableHttpError::UnexpectedServerResponse(Cow::Owned( - format!("HTTP {status}: {body}"), - ))); + return Err(StreamableHttpError::UnexpectedHttpStatus { + status: status.as_u16(), + body: Cow::Owned(body), + }); } let content_type = response.headers().get(http::header::CONTENT_TYPE).cloned(); diff --git a/crates/rmcp/src/transport/streamable_http_client.rs b/crates/rmcp/src/transport/streamable_http_client.rs index 5194ba96..6a1fc8ba 100644 --- a/crates/rmcp/src/transport/streamable_http_client.rs +++ b/crates/rmcp/src/transport/streamable_http_client.rs @@ -19,7 +19,7 @@ use super::common::client_side_sse::{ use crate::{ RoleClient, model::{ - ClientJsonRpcMessage, ClientNotification, ClientRequest, ErrorData, GetMeta, + ClientJsonRpcMessage, ClientNotification, ClientRequest, ErrorCode, ErrorData, GetMeta, InitializedNotification, JsonObject, ProtocolVersion, RequestId, ServerJsonRpcMessage, ServerResult, }, @@ -179,6 +179,11 @@ pub enum StreamableHttpError { UnexpectedEndOfStream, #[error("unexpected server response: {0}")] UnexpectedServerResponse(Cow<'static, str>), + #[error("unexpected HTTP status {status}: {body}")] + UnexpectedHttpStatus { + status: u16, + body: Cow<'static, str>, + }, #[error("Unexpected content type: {0:?}")] UnexpectedContentType(Option), #[error("Server does not support SSE")] @@ -824,6 +829,14 @@ impl Worker for StreamableHttpClientWorker { ClientJsonRpcMessage::Request(request) if matches!(&request.request, ClientRequest::InitializeRequest(_)) ); + let discover_startup_id = match &startup_request { + ClientJsonRpcMessage::Request(request) + if matches!(&request.request, ClientRequest::DiscoverRequest(_)) => + { + Some(request.id.clone()) + } + _ => None, + }; let mut saved_init_request = is_legacy_startup.then(|| startup_request.clone()); let empty_tool_cache = HashMap::new(); let (bootstrap_version, bootstrap_headers) = if is_legacy_startup { @@ -854,6 +867,27 @@ impl Worker for StreamableHttpClientWorker { WorkerQuitReason::fatal_context("process initialize response"), )? } + Err(StreamableHttpError::UnexpectedHttpStatus { status, body: _ }) + if discover_startup_id.is_some() && matches!(status, 404 | 405) => + { + // An initial discover request has no session, so 404/405 can + // indicate a legacy endpoint. Keep the worker alive long enough + // for Auto mode to retry the legacy initialize handshake. + let _ = responder.send(Ok(())); + ( + ServerJsonRpcMessage::error( + ErrorData::new( + ErrorCode(-32000), + "Discovery probe rejected by HTTP endpoint", + Some(serde_json::json!({ + crate::service::DISCOVER_PROBE_HTTP_STATUS_KEY: status + })), + ), + discover_startup_id, + ), + None, + ) + } Err(err) => { let msg = format!("{:?}", err); let _ = responder.send(Err(err)); diff --git a/crates/rmcp/tests/test_client_lifecycle_modes.rs b/crates/rmcp/tests/test_client_lifecycle_modes.rs index 375364e8..c46d00ba 100644 --- a/crates/rmcp/tests/test_client_lifecycle_modes.rs +++ b/crates/rmcp/tests/test_client_lifecycle_modes.rs @@ -7,7 +7,7 @@ use rmcp::{ Implementation, InitializeResult, ProtocolVersion, RequestId, ServerCapabilities, ServerJsonRpcMessage, ServerResult, }, - service::PeerRequestOptions, + service::{ClientInitializeError, PeerRequestOptions}, transport::{IntoTransport, Transport}, }; @@ -229,6 +229,273 @@ async fn auto_startup_falls_back_after_discover_method_not_found() { server_task.await.expect("server task"); } +async fn assert_auto_startup_falls_back( + rejection: impl FnOnce(RequestId) -> ServerJsonRpcMessage + Send + 'static, +) { + let (server_transport, client_transport) = tokio::io::duplex(4096); + let mut server = IntoTransport::::into_transport(server_transport); + let server_task = tokio::spawn(async move { + let ClientJsonRpcMessage::Request(discover) = + server.receive().await.expect("expected discover request") + else { + panic!("expected discover request"); + }; + assert!(matches!( + discover.request, + ClientRequest::DiscoverRequest(_) + )); + server + .send(rejection(discover.id)) + .await + .expect("send legacy discovery rejection"); + + let ClientJsonRpcMessage::Request(initialize) = server + .receive() + .await + .expect("expected fallback initialize request") + else { + panic!("expected fallback initialize request"); + }; + let ClientRequest::InitializeRequest(request) = initialize.request else { + panic!("expected initialize request"); + }; + assert_eq!( + request.params.protocol_version, + ProtocolVersion::V_2025_06_18 + ); + server + .send(ServerJsonRpcMessage::response( + ServerResult::InitializeResult( + InitializeResult::new(ServerCapabilities::default()), + ), + initialize.id, + )) + .await + .expect("send initialize response"); + assert!(matches!( + server.receive().await, + Some(ClientJsonRpcMessage::Notification(_)) + )); + }); + + let client = DiscoverClient + .serve_with_lifecycle( + client_transport, + ClientLifecycleMode::Auto { + preferred_versions: vec![ProtocolVersion::V_2026_07_28], + legacy_version: Some(ProtocolVersion::V_2025_06_18), + }, + ) + .await + .expect("Auto mode should fall back for a recognized legacy-only rejection"); + client.cancel().await.expect("cancel client"); + server_task.await.expect("server task"); +} + +#[tokio::test] +async fn auto_startup_falls_back_when_server_only_supports_historical_versions() { + assert_auto_startup_falls_back(|id| { + ServerJsonRpcMessage::error( + ErrorData::unsupported_protocol_version( + ProtocolVersion::V_2026_07_28, + &[ProtocolVersion::V_2025_11_25, ProtocolVersion::V_2025_06_18], + ), + Some(id), + ) + }) + .await; +} + +#[tokio::test] +async fn auto_startup_falls_back_when_discovery_only_advertises_historical_versions() { + assert_auto_startup_falls_back(|id| { + ServerJsonRpcMessage::response( + ServerResult::DiscoverResult(DiscoverResult::new( + vec![ProtocolVersion::V_2025_06_18], + ServerCapabilities::default(), + Implementation::new("legacy-server", "1.0.0"), + )), + id, + ) + }) + .await; +} + +#[tokio::test] +async fn auto_startup_falls_back_for_uncorrelated_legacy_protocol_prevalidation() { + assert_auto_startup_falls_back(|_| { + ServerJsonRpcMessage::error( + ErrorData::new( + ErrorCode(-32000), + "Bad Request: Unsupported protocol version: 2026-07-28 \ + (supported versions: 2025-11-25, 2025-06-18, 2025-03-26, \ + 2024-11-05, 2024-10-07)", + None, + ), + None, + ) + }) + .await; +} + +#[tokio::test] +async fn auto_startup_falls_back_for_uncorrelated_missing_session_prevalidation() { + assert_auto_startup_falls_back(|_| { + ServerJsonRpcMessage::error( + ErrorData::new( + ErrorCode(-32000), + "Bad Request: No valid session ID provided", + None, + ), + None, + ) + }) + .await; +} + +async fn assert_auto_startup_does_not_fall_back( + rejection: impl FnOnce(RequestId) -> ServerJsonRpcMessage + Send + 'static, +) -> ClientInitializeError { + let (server_transport, client_transport) = tokio::io::duplex(4096); + let mut server = IntoTransport::::into_transport(server_transport); + let server_task = tokio::spawn(async move { + let ClientJsonRpcMessage::Request(discover) = + server.receive().await.expect("expected discover request") + else { + panic!("expected discover request"); + }; + server + .send(rejection(discover.id)) + .await + .expect("send discovery rejection"); + assert!( + server.receive().await.is_none(), + "unsafe discovery rejection must not trigger initialize" + ); + }); + + let error = match DiscoverClient + .serve_with_lifecycle( + client_transport, + ClientLifecycleMode::Auto { + preferred_versions: vec![ProtocolVersion::V_2026_07_28], + legacy_version: Some(ProtocolVersion::V_2025_06_18), + }, + ) + .await + { + Ok(_) => panic!("unsafe discovery rejection must not trigger fallback"), + Err(error) => error, + }; + server_task.await.expect("server task"); + error +} + +#[tokio::test] +async fn auto_startup_does_not_fall_back_for_unrelated_error_response_id() { + let error = assert_auto_startup_does_not_fall_back(|_| { + ServerJsonRpcMessage::error( + ErrorData::new(ErrorCode::METHOD_NOT_FOUND, "Method not found", None), + Some(RequestId::Number(999)), + ) + }) + .await; + assert!(matches!( + error, + ClientInitializeError::ConflictInitResponseId(_, _) + )); +} + +#[tokio::test] +async fn auto_startup_does_not_fall_back_for_uncorrelated_method_not_found() { + let error = assert_auto_startup_does_not_fall_back(|_| { + ServerJsonRpcMessage::error( + ErrorData::new(ErrorCode::METHOD_NOT_FOUND, "Method not found", None), + None, + ) + }) + .await; + assert!(matches!(error, ClientInitializeError::JsonRpcError(_))); +} + +#[tokio::test] +async fn auto_startup_does_not_fall_back_for_arbitrary_uncorrelated_errors() { + let error = assert_auto_startup_does_not_fall_back(|_| { + ServerJsonRpcMessage::error( + ErrorData::new(ErrorCode(-32000), "Bad Request: database unavailable", None), + None, + ) + }) + .await; + assert!(matches!(error, ClientInitializeError::JsonRpcError(_))); +} + +#[tokio::test] +async fn auto_startup_does_not_fall_back_for_unknown_or_future_protocol_versions() { + let error = assert_auto_startup_does_not_fall_back(|_| { + ServerJsonRpcMessage::error( + ErrorData::new( + ErrorCode(-32000), + "Bad Request: Unsupported protocol version: 2026-07-28 \ + (supported versions: 2025-06-18, 2027-01-01)", + None, + ), + None, + ) + }) + .await; + assert!(matches!(error, ClientInitializeError::JsonRpcError(_))); +} + +#[tokio::test] +async fn auto_startup_does_not_fall_back_for_unknown_protocol_version_tokens() { + let error = assert_auto_startup_does_not_fall_back(|_| { + ServerJsonRpcMessage::error( + ErrorData::new( + ErrorCode(-32000), + "Bad Request: Unsupported protocol version: 2026-07-28 \ + (supported versions: 2025-06-18, next-draft)", + None, + ), + None, + ) + }) + .await; + assert!(matches!(error, ClientInitializeError::JsonRpcError(_))); +} + +#[tokio::test] +async fn auto_startup_does_not_fall_back_for_contradictory_supported_versions() { + let error = assert_auto_startup_does_not_fall_back(|_| { + ServerJsonRpcMessage::error( + ErrorData::new( + ErrorCode(-32000), + "Bad Request: Unsupported protocol version: 2026-07-28 \ + (supported versions: 2025-06-18, 2027-01-01)", + Some(serde_json::json!({ "supported": ["2025-06-18"] })), + ), + None, + ) + }) + .await; + assert!(matches!(error, ClientInitializeError::JsonRpcError(_))); +} + +#[tokio::test] +async fn auto_startup_does_not_fall_back_for_uncorrelated_protocol_version_error() { + let error = assert_auto_startup_does_not_fall_back(|_| { + ServerJsonRpcMessage::error( + ErrorData::unsupported_protocol_version( + ProtocolVersion::V_2026_07_28, + &[ProtocolVersion::V_2025_06_18], + ), + None, + ) + }) + .await; + assert!(matches!(error, ClientInitializeError::JsonRpcError(_))); +} + #[tokio::test] async fn discover_startup_retries_a_mutually_supported_version() { let unsupported: ProtocolVersion = diff --git a/crates/rmcp/tests/test_discover_http_client_startup.rs b/crates/rmcp/tests/test_discover_http_client_startup.rs index c6051a04..f176ef9a 100644 --- a/crates/rmcp/tests/test_discover_http_client_startup.rs +++ b/crates/rmcp/tests/test_discover_http_client_startup.rs @@ -5,11 +5,24 @@ feature = "transport-streamable-http-server" ))] -use std::borrow::Cow; +use std::{ + borrow::Cow, + sync::{Arc, Mutex}, +}; +use axum::{ + body::Bytes, + extract::State, + http::{StatusCode, header}, + response::{IntoResponse, Response}, + routing::post, +}; use rmcp::{ ClientLifecycleMode, ClientServiceExt, ServerHandler, - model::{ClientInfo, DiscoverResult, ErrorCode, ErrorData, ProtocolVersion}, + model::{ + ClientInfo, DiscoverResult, ErrorCode, ErrorData, InitializeResult, ProtocolVersion, + ServerCapabilities, + }, service::{MaybeSendFuture, RequestContext, RoleServer}, transport::{ StreamableHttpClientTransport, @@ -135,3 +148,274 @@ async fn auto_http_client_falls_back_to_stateful_legacy_startup() { ct.cancel(); server.await.expect("server task"); } + +#[derive(Clone, Copy)] +enum LegacyDiscoveryRejection { + UnsupportedProtocol, + MissingSession, + NotFound, + MethodNotAllowed, + Unauthorized, + Forbidden, + UnauthorizedJson, + ForbiddenJson, +} + +#[derive(Clone)] +struct LegacyPrevalidationState { + rejection: LegacyDiscoveryRejection, + methods: Arc>>, +} + +async fn legacy_prevalidation_handler( + State(state): State, + body: Bytes, +) -> Response { + let message: serde_json::Value = serde_json::from_slice(&body).expect("JSON-RPC request body"); + let method = message + .get("method") + .and_then(serde_json::Value::as_str) + .expect("JSON-RPC request method"); + state + .methods + .lock() + .expect("methods lock") + .push(method.into()); + + match method { + "server/discover" => match state.rejection { + LegacyDiscoveryRejection::UnsupportedProtocol => json_response( + StatusCode::BAD_REQUEST, + serde_json::json!({ + "jsonrpc": "2.0", + "id": null, + "error": { + "code": -32000, + "message": "Bad Request: Unsupported protocol version: 2026-07-28 \ + (supported versions: 2025-11-25, 2025-06-18, 2025-03-26, \ + 2024-11-05, 2024-10-07)", + }, + }), + ), + LegacyDiscoveryRejection::MissingSession => json_response( + StatusCode::BAD_REQUEST, + serde_json::json!({ + "jsonrpc": "2.0", + "id": null, + "error": { + "code": -32000, + "message": "Bad Request: No valid session ID provided", + }, + }), + ), + LegacyDiscoveryRejection::NotFound => { + (StatusCode::NOT_FOUND, "legacy endpoint not found").into_response() + } + LegacyDiscoveryRejection::MethodNotAllowed => { + (StatusCode::METHOD_NOT_ALLOWED, "legacy method not allowed").into_response() + } + LegacyDiscoveryRejection::Unauthorized => { + (StatusCode::UNAUTHORIZED, "authentication required").into_response() + } + LegacyDiscoveryRejection::Forbidden => { + (StatusCode::FORBIDDEN, "access forbidden").into_response() + } + LegacyDiscoveryRejection::UnauthorizedJson => json_response( + StatusCode::UNAUTHORIZED, + serde_json::json!({ + "jsonrpc": "2.0", + "id": null, + "error": { + "code": -32000, + "message": "Bad Request: No valid session ID provided", + }, + }), + ), + LegacyDiscoveryRejection::ForbiddenJson => json_response( + StatusCode::FORBIDDEN, + serde_json::json!({ + "jsonrpc": "2.0", + "id": null, + "error": { + "code": -32000, + "message": "Bad Request: Unsupported protocol version: 2026-07-28 \ + (supported versions: 2025-11-25, 2025-06-18)", + }, + }), + ), + }, + "initialize" => { + assert_eq!( + message + .get("params") + .and_then(|params| params.get("protocolVersion")), + Some(&serde_json::json!("2025-06-18")) + ); + let mut result = InitializeResult::new(ServerCapabilities::default()); + result.protocol_version = ProtocolVersion::V_2025_06_18; + json_response( + StatusCode::OK, + serde_json::json!({ + "jsonrpc": "2.0", + "id": message.get("id"), + "result": result, + }), + ) + } + "notifications/initialized" => StatusCode::ACCEPTED.into_response(), + "tools/list" => json_response( + StatusCode::OK, + serde_json::json!({ + "jsonrpc": "2.0", + "id": message.get("id"), + "result": { "tools": [] }, + }), + ), + _ => (StatusCode::BAD_REQUEST, "unexpected request").into_response(), + } +} + +fn json_response(status: StatusCode, value: serde_json::Value) -> Response { + ( + status, + [(header::CONTENT_TYPE, "application/json")], + value.to_string(), + ) + .into_response() +} + +async fn assert_http_legacy_fallback(rejection: LegacyDiscoveryRejection) { + let methods = Arc::new(Mutex::new(Vec::new())); + let router = axum::Router::new() + .route("/mcp", post(legacy_prevalidation_handler)) + .with_state(LegacyPrevalidationState { + rejection, + methods: methods.clone(), + }); + 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 cancellation = CancellationToken::new(); + let server = tokio::spawn({ + let cancellation = cancellation.clone(); + async move { + axum::serve(listener, router) + .with_graceful_shutdown(cancellation.cancelled_owned()) + .await + .expect("serve legacy HTTP endpoint"); + } + }); + + 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_06_18), + }, + ) + .await + .expect("Auto mode should recognize the deployed legacy HTTP response"); + client.list_tools(None).await.expect("list legacy tools"); + client.cancel().await.expect("cancel client"); + + assert_eq!( + *methods.lock().expect("methods lock"), + [ + "server/discover", + "initialize", + "notifications/initialized", + "tools/list", + ] + ); + cancellation.cancel(); + server.await.expect("server task"); +} + +#[tokio::test] +async fn auto_http_client_falls_back_after_unsupported_protocol_prevalidation() { + assert_http_legacy_fallback(LegacyDiscoveryRejection::UnsupportedProtocol).await; +} + +#[tokio::test] +async fn auto_http_client_falls_back_after_missing_session_prevalidation() { + assert_http_legacy_fallback(LegacyDiscoveryRejection::MissingSession).await; +} + +#[tokio::test] +async fn auto_http_client_falls_back_after_initial_http_not_found() { + assert_http_legacy_fallback(LegacyDiscoveryRejection::NotFound).await; +} + +#[tokio::test] +async fn auto_http_client_falls_back_after_initial_http_method_not_allowed() { + assert_http_legacy_fallback(LegacyDiscoveryRejection::MethodNotAllowed).await; +} + +async fn assert_http_auth_rejection_does_not_downgrade(rejection: LegacyDiscoveryRejection) { + let methods = Arc::new(Mutex::new(Vec::new())); + let router = axum::Router::new() + .route("/mcp", post(legacy_prevalidation_handler)) + .with_state(LegacyPrevalidationState { + rejection, + methods: methods.clone(), + }); + 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 cancellation = CancellationToken::new(); + let server = tokio::spawn({ + let cancellation = cancellation.clone(); + async move { + axum::serve(listener, router) + .with_graceful_shutdown(cancellation.cancelled_owned()) + .await + .expect("serve auth-rejecting HTTP endpoint"); + } + }); + + let transport = StreamableHttpClientTransport::from_config( + StreamableHttpClientTransportConfig::with_uri(format!("http://{address}/mcp")), + ); + let result = ClientInfo::default() + .serve_with_lifecycle( + transport, + ClientLifecycleMode::Auto { + preferred_versions: vec![ProtocolVersion::V_2026_07_28], + legacy_version: Some(ProtocolVersion::V_2025_06_18), + }, + ) + .await; + assert!( + result.is_err(), + "authentication failures must not downgrade" + ); + assert_eq!(*methods.lock().expect("methods lock"), ["server/discover"]); + cancellation.cancel(); + server.await.expect("server task"); +} + +#[tokio::test] +async fn auto_http_client_does_not_downgrade_after_http_401() { + assert_http_auth_rejection_does_not_downgrade(LegacyDiscoveryRejection::Unauthorized).await; +} + +#[tokio::test] +async fn auto_http_client_does_not_downgrade_after_http_403() { + assert_http_auth_rejection_does_not_downgrade(LegacyDiscoveryRejection::Forbidden).await; +} + +#[tokio::test] +async fn auto_http_client_does_not_downgrade_after_json_rpc_http_401() { + assert_http_auth_rejection_does_not_downgrade(LegacyDiscoveryRejection::UnauthorizedJson).await; +} + +#[tokio::test] +async fn auto_http_client_does_not_downgrade_after_json_rpc_http_403() { + assert_http_auth_rejection_does_not_downgrade(LegacyDiscoveryRejection::ForbiddenJson).await; +} 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 ea49a417..60919619 100644 --- a/crates/rmcp/tests/test_streamable_http_4xx_error_body.rs +++ b/crates/rmcp/tests/test_streamable_http_4xx_error_body.rs @@ -74,7 +74,7 @@ async fn http_4xx_json_rpc_error_body_is_surfaced_as_json_response() { } } -/// HTTP 4xx with non-JSON content-type must still return `UnexpectedServerResponse` +/// HTTP 4xx with non-JSON content-type must retain the status and response body. /// (no regression on the original error path). #[tokio::test] async fn http_4xx_non_json_body_returns_unexpected_server_response() { @@ -92,13 +92,13 @@ async fn http_4xx_non_json_body_returns_unexpected_server_response() { .await; match result { - Err(StreamableHttpError::UnexpectedServerResponse(_)) => {} - other => panic!("expected UnexpectedServerResponse, got: {other:?}"), + Err(StreamableHttpError::UnexpectedHttpStatus { status: 400, .. }) => {} + other => panic!("expected UnexpectedHttpStatus, got: {other:?}"), } } /// HTTP 4xx with Content-Type: application/json but a body that is NOT a valid -/// JSON-RPC message must fall back to `UnexpectedServerResponse`. +/// JSON-RPC message must retain the status and response body. #[tokio::test] async fn http_4xx_malformed_json_body_falls_back_to_unexpected_server_response() { let url = spawn_mock_server(400, "application/json", r#"{"error":"not jsonrpc"}"#).await; @@ -115,7 +115,7 @@ async fn http_4xx_malformed_json_body_falls_back_to_unexpected_server_response() .await; match result { - Err(StreamableHttpError::UnexpectedServerResponse(_)) => {} - other => panic!("expected UnexpectedServerResponse, got: {other:?}"), + Err(StreamableHttpError::UnexpectedHttpStatus { status: 400, .. }) => {} + other => panic!("expected UnexpectedHttpStatus, got: {other:?}"), } }