diff --git a/crates/rmcp/src/transport/streamable_http_server/tower.rs b/crates/rmcp/src/transport/streamable_http_server/tower.rs index 1b5382015..3d478e2b1 100644 --- a/crates/rmcp/src/transport/streamable_http_server/tower.rs +++ b/crates/rmcp/src/transport/streamable_http_server/tower.rs @@ -113,7 +113,10 @@ pub struct StreamableHttpServerConfig { /// Defaults to an empty list, which disables Origin validation for backward /// compatibility. A non-empty list enables validation. Requests carrying /// an `Origin` header must match per RFC 6454 `(scheme, host, port)`; - /// missing-`Origin` requests still pass. Entries must include a scheme; + /// missing-`Origin` requests still pass. An entry that omits the port + /// permits any port; an entry with an explicit port matches only that + /// port, resolving an origin's omitted port to the scheme default + /// (443 for `https`, 80 for `http`). Entries must include a scheme; /// `"null"` matches the browser's `Origin: null`. /// /// Call [`StreamableHttpServerConfig::enforce_origin_validation`] to enable @@ -853,6 +856,15 @@ fn parse_origin_value(value: &str) -> Option { }) } +/// The port a scheme implies when an origin serialization omits it (RFC 6454 §4). +fn default_port(scheme: &str) -> Option { + match scheme { + "http" => Some(80), + "https" => Some(443), + _ => None, + } +} + fn origin_is_allowed(origin: &NormalizedOrigin, allowed_origins: &[String]) -> bool { allowed_origins .iter() @@ -870,7 +882,17 @@ 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 + && match a_port { + // An omitted configured port permits any port. + None => true, + // RFC 6454 §6.2 omits the default port when serializing an + // origin, so an absent incoming port means the scheme default. + Some(a_port) => o_port.or_else(|| default_port(o_scheme)) == Some(*a_port), + } + } _ => false, }) } diff --git a/crates/rmcp/tests/test_custom_headers.rs b/crates/rmcp/tests/test_custom_headers.rs index b01223a6f..8f973f421 100644 --- a/crates/rmcp/tests/test_custom_headers.rs +++ b/crates/rmcp/tests/test_custom_headers.rs @@ -1311,4 +1311,58 @@ mod origin_validation { let response = service.handle(init_request(Some("null"))).await; assert_eq!(response.status(), http::StatusCode::FORBIDDEN); } + + #[tokio::test] + async fn explicit_https_default_port_allows_port_less_origin() { + let service = service_with_allowed_origins(&["https://client.example:443"]); + let response = service + .handle(init_request(Some("https://client.example"))) + .await; + assert_eq!(response.status(), http::StatusCode::OK); + } + + #[tokio::test] + async fn explicit_http_default_port_allows_port_less_origin() { + let service = service_with_allowed_origins(&["http://client.example:80"]); + let response = service + .handle(init_request(Some("http://client.example"))) + .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"]); + let response = service + .handle(init_request(Some("https://client.example:443"))) + .await; + assert_eq!(response.status(), http::StatusCode::OK); + } + + #[tokio::test] + async fn explicit_default_port_forbids_non_default_origin_port() { + let service = service_with_allowed_origins(&["https://client.example:443"]); + let response = service + .handle(init_request(Some("https://client.example:8443"))) + .await; + assert_eq!(response.status(), http::StatusCode::FORBIDDEN); + } + + #[tokio::test] + async fn explicit_non_default_port_forbids_port_less_origin() { + let service = service_with_allowed_origins(&["https://client.example:8443"]); + let response = service + .handle(init_request(Some("https://client.example"))) + .await; + assert_eq!(response.status(), http::StatusCode::FORBIDDEN); + } + + #[tokio::test] + async fn omitted_configured_port_allows_any_origin_port() { + let service = service_with_allowed_origins(&["https://client.example"]); + let response = service + .handle(init_request(Some("https://client.example:8443"))) + .await; + assert_eq!(response.status(), http::StatusCode::OK); + } }