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..8f39107c2 100644 --- a/crates/rmcp/src/transport/common/auth/streamable_http_client.rs +++ b/crates/rmcp/src/transport/common/auth/streamable_http_client.rs @@ -71,6 +71,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/streamable_http_client.rs b/crates/rmcp/src/transport/streamable_http_client.rs index a8e59a244..6eb758ccc 100644 --- a/crates/rmcp/src/transport/streamable_http_client.rs +++ b/crates/rmcp/src/transport/streamable_http_client.rs @@ -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" )] @@ -1681,8 +1711,17 @@ impl Worker for StreamableHttpClientWorker { context.send_to_handler(message).await?; Ok(()) } - Ok(StreamableHttpPostResponse::RawJson(message, ..)) => { - context.send_to_handler(message).await?; + 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, ..)) => { @@ -2531,6 +2570,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": "" })); @@ -2558,6 +2627,34 @@ 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() { diff --git a/crates/rmcp/tests/test_custom_request.rs b/crates/rmcp/tests/test_custom_request.rs index 6ccc5992f..7f2efbcae 100644 --- a/crates/rmcp/tests/test_custom_request.rs +++ b/crates/rmcp/tests/test_custom_request.rs @@ -170,6 +170,217 @@ async fn typed_test_client() -> anyhow::Result 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?;