Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
196 changes: 187 additions & 9 deletions crates/rmcp/src/transport/streamable_http_server/tower.rs
Original file line number Diff line number Diff line change
Expand Up @@ -968,6 +968,45 @@ enum NormalizedOrigin {
},
}

fn is_valid_ip_literal(value: &str) -> bool {
if value.parse::<std::net::Ipv6Addr>().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<NormalizedOrigin> {
let value = value.trim();
if value.is_empty() {
Expand All @@ -976,6 +1015,43 @@ fn parse_origin_value(value: &str) -> Option<NormalizedOrigin> {
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::<u16>().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()?;
Expand Down Expand Up @@ -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(());
}
Expand Down Expand Up @@ -1323,10 +1402,9 @@ pub struct StreamableHttpService<S, M> {
/// than racing to replay the initialize handshake. `None` when no external
/// session store is configured (avoids allocating the map).
pending_restores: Option<PendingRestores>,
/// 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<std::sync::RwLock<HashMap<String, Option<Arc<JsonObject>>>>>,
/// Caches known tool input schemas for SEP-2243 `Mcp-Param-*` validation.
/// Unknown names and failed service-factory lookups are never retained.
tool_schemas: Arc<std::sync::RwLock<HashMap<String, Arc<JsonObject>>>>,
}

impl<S, M> Clone for StreamableHttpService<S, M> {
Expand Down Expand Up @@ -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<Arc<JsonObject>> {
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
Expand Down Expand Up @@ -2460,3 +2540,101 @@ impl<S: Stream> Stream for CancelOnDisconnect<S> {
polled
}
}

#[cfg(test)]
mod tool_schema_cache_tests {
use std::sync::atomic::{AtomicUsize, Ordering};

use super::*;
use crate::{
ServerHandler,
model::{ServerCapabilities, Tool},
};

#[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<Tool> {
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<CacheTestHandler, std::io::Error> + Send + Sync + 'static,
) -> StreamableHttpService<CacheTestHandler, super::super::session::local::LocalSessionManager>
{
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"));
}
}
97 changes: 97 additions & 0 deletions crates/rmcp/tests/test_custom_headers.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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"]);
Expand Down Expand Up @@ -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"]);
Expand Down
Loading