From 9b27f986962986b3814ccf4c65f6c754aa3ab3aa Mon Sep 17 00:00:00 2001 From: ECD5A Date: Thu, 20 Aug 2026 14:32:49 +0500 Subject: [PATCH] rmcp: return forbidden for invalid origin Signed-off-by: ECD5A --- .../transport/streamable_http_server/tower.rs | 4 ++-- crates/rmcp/tests/test_custom_headers.rs | 23 ++++++++++++++++++- 2 files changed, 24 insertions(+), 3 deletions(-) diff --git a/crates/rmcp/src/transport/streamable_http_server/tower.rs b/crates/rmcp/src/transport/streamable_http_server/tower.rs index f1fef585f..74c0aad0a 100644 --- a/crates/rmcp/src/transport/streamable_http_server/tower.rs +++ b/crates/rmcp/src/transport/streamable_http_server/tower.rs @@ -893,13 +893,13 @@ fn validate_origin_header( .inspect_err(|_| { tracing::warn!(origin = ?origin_header, "rejected request with non-UTF-8 Origin header"); }) - .map_err(|_| bad_request_response("Bad Request: Invalid Origin header encoding"))?; + .map_err(|_| forbidden_response("Forbidden: Invalid Origin header encoding"))?; let origin = parse_origin_value(origin_str).ok_or_else(|| { tracing::warn!( origin = origin_str, "rejected request with malformed Origin header", ); - bad_request_response("Bad Request: Invalid Origin header") + forbidden_response("Forbidden: Invalid Origin header") })?; if !origin_is_allowed(&origin, allowed_origins) { tracing::warn!( diff --git a/crates/rmcp/tests/test_custom_headers.rs b/crates/rmcp/tests/test_custom_headers.rs index 736dce18e..f722e3241 100644 --- a/crates/rmcp/tests/test_custom_headers.rs +++ b/crates/rmcp/tests/test_custom_headers.rs @@ -1127,7 +1127,10 @@ mod origin_validation { use std::sync::Arc; use bytes::Bytes; - use http::{Method, Request, header::CONTENT_TYPE}; + use http::{ + Method, Request, + header::{CONTENT_TYPE, HeaderValue, ORIGIN}, + }; use http_body_util::Full; use rmcp::{ handler::server::ServerHandler, @@ -1199,6 +1202,24 @@ mod origin_validation { assert_eq!(response.status(), http::StatusCode::FORBIDDEN); } + #[tokio::test] + async fn malformed_origin_is_forbidden() { + let service = service_with_allowed_origins(&["http://localhost:8080"]); + let response = service.handle(init_request(Some("not-an-origin"))).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"]); + let mut request = init_request(None); + request + .headers_mut() + .insert(ORIGIN, HeaderValue::from_bytes(&[0xff]).unwrap()); + let response = service.handle(request).await; + assert_eq!(response.status(), http::StatusCode::FORBIDDEN); + } + #[tokio::test] async fn missing_origin_passes_through() { let service = service_with_allowed_origins(&["http://localhost:8080"]);