From 66df48a585234821036081461586461a50765dbb Mon Sep 17 00:00:00 2001 From: GraciousGazelles <55840159+GraciousGazelles@users.noreply.github.com> Date: Sun, 4 Oct 2026 18:47:48 +1100 Subject: [PATCH 1/2] fix(http): validate Origin literals and cache known schemas --- .../transport/streamable_http_server/tower.rs | 195 +++++++++++++++++- crates/rmcp/tests/test_custom_headers.rs | 97 +++++++++ 2 files changed, 283 insertions(+), 9 deletions(-) diff --git a/crates/rmcp/src/transport/streamable_http_server/tower.rs b/crates/rmcp/src/transport/streamable_http_server/tower.rs index 5d7cee556..68faab27d 100644 --- a/crates/rmcp/src/transport/streamable_http_server/tower.rs +++ b/crates/rmcp/src/transport/streamable_http_server/tower.rs @@ -968,6 +968,45 @@ enum NormalizedOrigin { }, } +fn is_valid_ip_literal(value: &str) -> bool { + if value.parse::().is_ok() { + return true; + } + + let Some(version_and_address) = value.strip_prefix('v').or_else(|| value.strip_prefix('V')) + else { + return false; + }; + let Some((version, address)) = version_and_address.split_once('.') else { + return false; + }; + + !version.is_empty() + && version.bytes().all(|byte| byte.is_ascii_hexdigit()) + && !address.is_empty() + && address.bytes().all(|byte| { + byte.is_ascii_alphanumeric() + || matches!( + byte, + b'-' | b'.' + | b'_' + | b'~' + | b'!' + | b'$' + | b'&' + | b'\'' + | b'(' + | b')' + | b'*' + | b'+' + | b',' + | b';' + | b'=' + | b':' + ) + }) +} + fn parse_origin_value(value: &str) -> Option { let value = value.trim(); if value.is_empty() { @@ -976,6 +1015,43 @@ 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 + .chars() + .any(|character| matches!(character, '@' | '/' | '?' | '#')) + { + return None; + } + let valid_port = |port: &str| { + !port.is_empty() + && port.bytes().all(|byte| byte.is_ascii_digit()) + && port.parse::().is_ok() + }; + if let Some(bracketed_authority) = serialized_authority.strip_prefix('[') { + let closing_bracket = bracketed_authority.find(']')?; + let bracketed_host = &bracketed_authority[..closing_bracket]; + let suffix = &bracketed_authority[closing_bracket + 1..]; + if !is_valid_ip_literal(bracketed_host) + || (!suffix.is_empty() && !suffix.strip_prefix(':').is_some_and(valid_port)) + { + return None; + } + } else { + if serialized_authority + .chars() + .any(|character| matches!(character, '[' | ']')) + { + return None; + } + if let Some((host, port)) = serialized_authority.split_once(':') + && (host.is_empty() || port.contains(':') || !valid_port(port)) + { + return None; + } + } + if serialized_authority.is_empty() { + return None; + } let uri = http::Uri::try_from(value).ok()?; let scheme = uri.scheme_str()?.to_ascii_lowercase(); let authority = uri.authority()?; @@ -1202,6 +1278,9 @@ fn validate_origin_header( headers: &HeaderMap, config: &StreamableHttpServerConfig, ) -> HttpResult<()> { + if headers.get_all(http::header::ORIGIN).iter().take(2).count() > 1 { + return Err(forbidden_response("Forbidden: Multiple Origin headers").into()); + } if !config.validate_empty_origin_allowlist && config.allowed_origins.is_empty() { return Ok(()); } @@ -1323,10 +1402,9 @@ pub struct StreamableHttpService { /// than racing to replay the initialize handshake. `None` when no external /// session store is configured (avoids allocating the map). 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. - tool_schemas: Arc>>>>, + /// Caches known tool input schemas for SEP-2243 `Mcp-Param-*` validation. + /// Unknown names and failed service-factory lookups are never retained. + tool_schemas: Arc>>>, } impl Clone for StreamableHttpService { @@ -1580,21 +1658,23 @@ where Ok(self.stateless_sse_response(Some(first), receiver, request_ct)) } - /// Returns the cached input schema for `name`, constructing a service once - /// per name to read its `ServerHandler::get_tool` definition. Used to - /// validate SEP-2243 `Mcp-Param-*` headers against the request body. + /// Returns a cached schema for known tools, otherwise constructs a service + /// and reads its `ServerHandler::get_tool` definition. Only successful + /// schemas are retained. Used to validate SEP-2243 `Mcp-Param-*` headers. fn tool_schema(&self, name: &str) -> Option> { if let Ok(cache) = self.tool_schemas.read() && let Some(schema) = cache.get(name) { - return schema.clone(); + return Some(schema.clone()); } let schema = self .get_service() .ok() .and_then(|service| service.get_tool(name)) .map(|tool| tool.input_schema); - if let Ok(mut cache) = self.tool_schemas.write() { + if let Some(schema) = schema.as_ref() + && let Ok(mut cache) = self.tool_schemas.write() + { cache.insert(name.to_owned(), schema.clone()); } schema @@ -2460,3 +2540,100 @@ impl Stream for CancelOnDisconnect { polled } } + +#[cfg(test)] +mod tool_schema_cache_tests { + use super::*; + use crate::{ + ServerHandler, + model::{ServerCapabilities, Tool}, + }; + use std::sync::atomic::{AtomicUsize, Ordering}; + + #[derive(Clone)] + struct CacheTestHandler; + + impl ServerHandler for CacheTestHandler { + fn get_info(&self) -> ServerConfig { + ServerConfig::new(ServerCapabilities::builder().enable_tools().build()) + } + + fn get_tool(&self, name: &str) -> Option { + if name != "known" { + return None; + } + let schema = + serde_json::json!({"type":"object","properties":{"value":{"type":"string"}}}) + .as_object() + .expect("schema is an object") + .clone(); + Some(Tool::new("known", "known tool", Arc::new(schema))) + } + } + + fn service( + factory: impl Fn() -> Result + Send + Sync + 'static, + ) -> StreamableHttpService + { + StreamableHttpService::new( + factory, + Arc::new(super::super::session::local::LocalSessionManager::default()), + StreamableHttpServerConfig::default(), + ) + } + + #[test] + fn distinct_unknown_tool_names_are_not_retained() { + let factory_calls = Arc::new(AtomicUsize::new(0)); + let calls = factory_calls.clone(); + let service = service(move || { + calls.fetch_add(1, Ordering::SeqCst); + Ok(CacheTestHandler) + }); + + for index in 0..128 { + assert!(service.tool_schema(&format!("unknown-{index}")).is_none()); + } + + assert_eq!(factory_calls.load(Ordering::SeqCst), 128); + assert!(service.tool_schemas.read().unwrap().is_empty()); + } + + #[test] + fn service_factory_failures_are_not_retained() { + let factory_calls = Arc::new(AtomicUsize::new(0)); + let calls = factory_calls.clone(); + let service = service(move || { + calls.fetch_add(1, Ordering::SeqCst); + Err(std::io::Error::other("factory unavailable")) + }); + + for _ in 0..3 { + assert!(service.tool_schema("known").is_none()); + } + + assert_eq!(factory_calls.load(Ordering::SeqCst), 3); + assert!(service.tool_schemas.read().unwrap().is_empty()); + } + + #[test] + fn known_schema_is_cached_and_shared_across_clones() { + let factory_calls = Arc::new(AtomicUsize::new(0)); + let calls = factory_calls.clone(); + let service = service(move || { + calls.fetch_add(1, Ordering::SeqCst); + Ok(CacheTestHandler) + }); + let first = service.tool_schema("known").expect("known schema"); + let cloned_service = service.clone(); + let second = cloned_service + .tool_schema("known") + .expect("cached known schema"); + + assert!(Arc::ptr_eq(&first, &second)); + assert_eq!(factory_calls.load(Ordering::SeqCst), 1); + let cache = service.tool_schemas.read().unwrap(); + assert_eq!(cache.len(), 1); + assert!(cache.contains_key("known")); + } +} diff --git a/crates/rmcp/tests/test_custom_headers.rs b/crates/rmcp/tests/test_custom_headers.rs index 5d4419e5c..5c5d7fb96 100644 --- a/crates/rmcp/tests/test_custom_headers.rs +++ b/crates/rmcp/tests/test_custom_headers.rs @@ -1215,6 +1215,87 @@ mod origin_validation { assert_eq!(response.status(), http::StatusCode::FORBIDDEN); } + #[tokio::test] + async fn serialized_origin_rejects_userinfo_path_query_and_fragment() { + let service = service_with_allowed_origins(&["http://localhost:8080"]); + for malformed in [ + "http://user@localhost:8080", + "http://localhost:8080/path", + "http://localhost:8080?query", + "http://localhost:8080#fragment", + ] { + let response = service.handle(init_request(Some(malformed))).await; + assert_eq!( + response.status(), + http::StatusCode::FORBIDDEN, + "malformed Origin {malformed:?} must not match an allowed authority" + ); + } + } + + #[tokio::test] + async fn serialized_origin_rejects_malformed_ports_that_would_default_to_http() { + let service = service_with_allowed_origins(&["http://localhost:80"]); + for malformed in [ + "http://localhost:", + "http://localhost:abc", + "http://localhost:65536", + ] { + let response = service.handle(init_request(Some(malformed))).await; + assert_eq!( + response.status(), + http::StatusCode::FORBIDDEN, + "malformed Origin {malformed:?} must not match an allowed authority" + ); + } + } + + #[tokio::test] + async fn serialized_origin_rejects_malformed_authority_brackets() { + for (allowed, malformed) in [ + ("http://[::1]:80", "http://[::1]garbage"), + ("http://[::1]:80", "http://[::1]]:80"), + ("http://localhost:8080", "http://localhost:8080[::1]:8080"), + ] { + let service = service_with_allowed_origins(&[allowed]); + let response = service.handle(init_request(Some(malformed))).await; + assert_eq!( + response.status(), + http::StatusCode::FORBIDDEN, + "malformed Origin {malformed:?} must not match {allowed:?}" + ); + } + } + + #[tokio::test] + async fn bracketed_registered_names_and_ipv4_are_forbidden() { + for (allowed, malformed) in [ + ("http://localhost:8080", "http://[localhost]:8080"), + ("http://127.0.0.1:8080", "http://[127.0.0.1]:8080"), + ] { + let service = service_with_allowed_origins(&[allowed]); + let response = service.handle(init_request(Some(malformed))).await; + assert_eq!( + response.status(), + http::StatusCode::FORBIDDEN, + "invalid IP literal {malformed:?} must not match {allowed:?}" + ); + } + } + + #[tokio::test] + async fn duplicated_origin_headers_are_forbidden_before_value_selection() { + 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, + HeaderValue::from_static("http://attacker.example"), + ); + + let response = service.handle(request).await; + assert_eq!(response.status(), http::StatusCode::FORBIDDEN); + } + #[tokio::test] async fn non_utf8_origin_is_forbidden() { let service = service_with_allowed_origins(&["http://localhost:8080"]); @@ -1330,6 +1411,22 @@ mod origin_validation { assert_eq!(response.status(), http::StatusCode::OK); } + #[tokio::test] + async fn bracketed_ipv6_origin_accepts_its_explicit_default_port() { + let service = service_with_allowed_origins(&["http://[::1]:80"]); + let response = service.handle(init_request(Some("http://[::1]"))).await; + assert_eq!(response.status(), http::StatusCode::OK); + } + + #[tokio::test] + async fn bracketed_ipvfuture_origin_is_accepted() { + let service = service_with_allowed_origins(&["http://[v1.alpha:beta]"]); + let response = service + .handle(init_request(Some("http://[v1.alpha:beta]"))) + .await; + assert_eq!(response.status(), http::StatusCode::OK); + } + #[tokio::test] async fn explicit_default_port_allows_matching_explicit_origin_port() { let service = service_with_allowed_origins(&["https://client.example:443"]); From 434d5387a1b08cc4e5e4ea5e743755f2b655864b Mon Sep 17 00:00:00 2001 From: GraciousGazelles <55840159+GraciousGazelles@users.noreply.github.com> Date: Sun, 4 Oct 2026 19:36:03 +1100 Subject: [PATCH 2/2] style: group cache test imports --- crates/rmcp/src/transport/streamable_http_server/tower.rs | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/crates/rmcp/src/transport/streamable_http_server/tower.rs b/crates/rmcp/src/transport/streamable_http_server/tower.rs index 68faab27d..03f197ab0 100644 --- a/crates/rmcp/src/transport/streamable_http_server/tower.rs +++ b/crates/rmcp/src/transport/streamable_http_server/tower.rs @@ -2543,12 +2543,13 @@ impl Stream for CancelOnDisconnect { #[cfg(test)] mod tool_schema_cache_tests { + use std::sync::atomic::{AtomicUsize, Ordering}; + use super::*; use crate::{ ServerHandler, model::{ServerCapabilities, Tool}, }; - use std::sync::atomic::{AtomicUsize, Ordering}; #[derive(Clone)] struct CacheTestHandler;