From 2d5abc4f0be3bf22de5c96add97e58db0cb5d4cd Mon Sep 17 00:00:00 2001 From: Pratik Gandhi Date: Mon, 10 Aug 2026 16:19:06 +0100 Subject: [PATCH 1/2] Add MCP header safety guardrails Signed-off-by: Pratik Gandhi --- _context/wiki/architecture.md | 6 +- _context/wiki/config.md | 14 + _context/wiki/security.md | 7 + .../contextforge-data-plane-lib/src/common.rs | 16 ++ .../src/gateway/mcp_service/initialization.rs | 34 ++- .../src/layers/mcp_header_limits.rs | 239 ++++++++++++++++++ .../src/layers/mod.rs | 1 + crates/contextforge-data-plane-lib/src/lib.rs | 5 + 8 files changed, 320 insertions(+), 2 deletions(-) create mode 100644 crates/contextforge-data-plane-lib/src/layers/mcp_header_limits.rs diff --git a/_context/wiki/architecture.md b/_context/wiki/architecture.md index 3ef364f8..eb52328c 100644 --- a/_context/wiki/architecture.md +++ b/_context/wiki/architecture.md @@ -11,6 +11,7 @@ TCP/TLS listener -> /contextforge-rs nested router -> mcp_origin_layer → validates Host then Origin (403 when disallowed) -> CORS layer + -> mcp_header_limits_layer → MCP standard header budgets (431 when exceeded) -> virtual_host_id_layer → inserts VirtualHostId (400 on path mismatch) -> claims_layer → inserts ContextForgeClaims (401 on bad/missing JWT) -> session_id_layer → inserts SessionId if present @@ -23,6 +24,9 @@ TCP/TLS listener checks the optional Host allowlist first, then rejects any present Origin that is malformed or not allowlisted; requests without Origin continue. See [Security](security.md#mcp-origin-and-host-validation). +`mcp_header_limits_layer` rejects excessive MCP standard headers before JWT +validation, config lookup, session creation, backend fanout, or RMCP body +parsing. MCP handlers read typed extensions — they never parse headers, paths, or Redis keys directly. @@ -30,7 +34,7 @@ MCP handlers read typed extensions — they never parse headers, paths, or Redis ```text downstream request - -> Host/Origin validation → virtual host extraction → JWT validation → session extraction + -> Host/Origin validation → MCP header limits → virtual host extraction → JWT validation → session extraction -> user config lookup → MCP handler validation -> request plugin hooks -> backend MCP call (concurrent via join_all for initialize/list) diff --git a/_context/wiki/config.md b/_context/wiki/config.md index efbd62f3..e567e440 100644 --- a/_context/wiki/config.md +++ b/_context/wiki/config.md @@ -40,10 +40,19 @@ Origin and Host settings retain the explicitly configured | --- | --- | --- | --- | | `--mcp-allowed-origins ` | `CONTEXTFORGE_GATEWAY_RS_MCP_ALLOWED_ORIGINS` | None | Browser Origin allowlist. Without it, requests lacking `Origin` pass and every request carrying `Origin` receives HTTP `403`. | | `--mcp-allowed-hosts ` | `CONTEXTFORGE_GATEWAY_RS_MCP_ALLOWED_HOSTS` | None | Optional request-authority allowlist. When configured, missing, malformed, or unlisted authorities receive HTTP `403`. | +| `--mcp-standard-header-max-count ` | `CONTEXTFORGE_DATA_PLANE_MCP_STANDARD_HEADER_MAX_COUNT` | `32` | Maximum MCP standard headers accepted on one request. | +| `--mcp-standard-header-max-value-bytes ` | `CONTEXTFORGE_DATA_PLANE_MCP_STANDARD_HEADER_MAX_VALUE_BYTES` | `8192` | Maximum byte length accepted for one MCP standard header value. | +| `--mcp-standard-header-max-total-bytes ` | `CONTEXTFORGE_DATA_PLANE_MCP_STANDARD_HEADER_MAX_TOTAL_BYTES` | `65536` | Approximate maximum combined bytes across MCP standard header names and values. | Values are comma-separated. Origin entries must be fully qualified serialized origins such as `https://app.example.com`; Host entries are authorities such as `gateway.example.com` or `gateway.example.com:8443`. See [Security](security.md#mcp-origin-and-host-validation). +The MCP standard header limits apply to `Mcp-Method`, `Mcp-Name`, +`Mcp-Protocol-Version`, `Mcp-Session-Id`, and `Mcp-Param-*`. A configured value +of `0` is treated as the documented default. The byte totals are +application-level budgets based on header name and value lengths; they are not +exact wire-size accounting and do not model HTTP/2 header compression. +Non-MCP headers remain bounded by the HTTP transport. ### Redis @@ -134,6 +143,11 @@ BackendMCPGateway | Hop-by-hop | `Connection`, `Keep-Alive`, `Proxy-Authenticate`, `Proxy-Authorization`, `Proxy-Connection`, `TE`, `Trailer`, `Trailers`, `Transfer-Encoding`, `Upgrade` | | RMCP-reserved | `Mcp-Session-Id`, `Accept`, `Last-Event-Id` | | Gateway-managed | `Host` (set from backend URL host + port; never overridden by config) | +| Computed MCP standard | `Mcp-Method`, `Mcp-Name`, `Mcp-Protocol-Version`, `Mcp-Param-*` | + +`Authorization` and `Cookie` are not protected here because backend +authentication through `passthrough_headers` or `add_headers` is intentional +runtime configuration. Redis storage: `MessagePack(User::new(sub))` → `MessagePack(UserConfig)`. diff --git a/_context/wiki/security.md b/_context/wiki/security.md index 65eea3e9..354289a9 100644 --- a/_context/wiki/security.md +++ b/_context/wiki/security.md @@ -67,6 +67,13 @@ or path/query/fragment/userinfo-bearing origins are rejected. Default ports are normalized (`https://a` equals `https://a:443`). There is no same-origin fallback; configure both allowlists for public deployments. +`mcp_header_limits_layer` enforces configurable count, per-value byte, and +approximate total byte budgets for MCP standard request headers before JWT +validation or RMCP body parsing. That budget covers `Mcp-Method`, `Mcp-Name`, +`Mcp-Protocol-Version`, `Mcp-Session-Id`, and `Mcp-Param-*`. It is an +application-level guard for MCP standard headers only; non-MCP headers remain +bounded by the HTTP transport. + ## Local Bootstrap Helpers (`with_tools`) The `contextforge-data-plane-lib/with_tools` feature compiles in: diff --git a/crates/contextforge-data-plane-lib/src/common.rs b/crates/contextforge-data-plane-lib/src/common.rs index b736b35e..072b78e0 100644 --- a/crates/contextforge-data-plane-lib/src/common.rs +++ b/crates/contextforge-data-plane-lib/src/common.rs @@ -198,6 +198,18 @@ pub struct Config { #[arg(long, env = "CONTEXTFORGE_DATA_PLANE_OTEL_EXPORTER_OTLP_METRICS_ENDPOINT")] pub otlp_metrics_endpoint: Option, + /// Maximum number of MCP standard headers accepted on a single request. + #[arg(long, env = "CONTEXTFORGE_DATA_PLANE_MCP_STANDARD_HEADER_MAX_COUNT", default_value_t = DEFAULT_MCP_STANDARD_HEADER_MAX_COUNT)] + pub mcp_standard_header_max_count: usize, + + /// Maximum byte length accepted for a single MCP standard header value. + #[arg(long, env = "CONTEXTFORGE_DATA_PLANE_MCP_STANDARD_HEADER_MAX_VALUE_BYTES", default_value_t = DEFAULT_MCP_STANDARD_HEADER_MAX_VALUE_BYTES)] + pub mcp_standard_header_max_value_bytes: usize, + + /// Approximate maximum total bytes accepted across MCP standard headers. + #[arg(long, env = "CONTEXTFORGE_DATA_PLANE_MCP_STANDARD_HEADER_MAX_TOTAL_BYTES", default_value_t = DEFAULT_MCP_STANDARD_HEADER_MAX_TOTAL_BYTES)] + pub mcp_standard_header_max_total_bytes: usize, + #[arg(long, env = "CONTEXTFORGE_DATA_PLANE_NUMBER_OF_CPUS")] pub number_of_cpus: Option, @@ -275,6 +287,10 @@ pub struct Config { pub mcp_allowed_hosts: Option>, } +pub const DEFAULT_MCP_STANDARD_HEADER_MAX_COUNT: usize = 32; +pub const DEFAULT_MCP_STANDARD_HEADER_MAX_VALUE_BYTES: usize = 8 * 1024; +pub const DEFAULT_MCP_STANDARD_HEADER_MAX_TOTAL_BYTES: usize = 64 * 1024; + #[derive(Error, Debug)] pub enum ConfigValidationError { #[error("Redis Configuration Error")] diff --git a/crates/contextforge-data-plane-lib/src/gateway/mcp_service/initialization.rs b/crates/contextforge-data-plane-lib/src/gateway/mcp_service/initialization.rs index 705c6c3b..448c56fa 100644 --- a/crates/contextforge-data-plane-lib/src/gateway/mcp_service/initialization.rs +++ b/crates/contextforge-data-plane-lib/src/gateway/mcp_service/initialization.rs @@ -17,6 +17,7 @@ use crate::gateway::{ mcp_call_validator::InitializeCallValidator, session_store::{UserSession, UserSessionStore}, }; +use crate::layers::mcp_header_limits::is_mcp_computed_header_name; pub(super) async fn initialize( mcp_service: &McpService, @@ -247,6 +248,7 @@ fn apply_header_config( /// - Hop-by-hop (RFC 7230 §6.1): `Connection`, `Keep-Alive`, `Proxy-Authenticate`, `Proxy-Authorization`, `TE`, `Trailer`, `Trailers`, `Transfer-Encoding`, `Upgrade` /// - Non-standard hop-by-hop: `Proxy-Connection` (must not cross gateway boundary) /// - RMCP transport-reserved: `Mcp-Session-Id`, `Accept`, `Last-Event-Id` +/// - MCP standard computed headers: `Mcp-Method`, `Mcp-Name`, `Mcp-Protocol-Version`, `Mcp-Param-*` fn is_protected_header(name: &http::HeaderName) -> bool { const PROTECTED: &[&str] = &[ "host", @@ -270,7 +272,7 @@ fn is_protected_header(name: &http::HeaderName) -> bool { "accept", "last-event-id", ]; - PROTECTED.iter().any(|&p| name.as_str().eq_ignore_ascii_case(p)) + PROTECTED.iter().any(|&p| name.as_str().eq_ignore_ascii_case(p)) || is_mcp_computed_header_name(name) } #[cfg(test)] @@ -406,6 +408,36 @@ mod tests { assert!(headers.is_empty(), "no RMCP-reserved header must reach the upstream config"); } + #[test] + fn computed_mcp_headers_cannot_be_passed_through_added_or_removed() { + let mut headers = HashMap::new(); + headers.insert(http::HeaderName::from_static("mcp-method"), http::HeaderValue::from_static("tools/call")); + headers.insert(http::HeaderName::from_static("mcp-param-user"), http::HeaderValue::from_static("computed")); + let ds = downstream(&[ + ("Mcp-Method", "wrong/method"), + ("Mcp-Name", "wrong-tool"), + ("Mcp-Protocol-Version", "2020-01-01"), + ("Mcp-Param-User", "wrong-user"), + ]); + let cfg = backend( + &["mcp-method", "mcp-name", "mcp-protocol-version", "mcp-param-user"], + &[ + ("Mcp-Method", "added/method"), + ("Mcp-Name", "added-tool"), + ("Mcp-Protocol-Version", "2020-01-01"), + ("Mcp-Param-User", "added-user"), + ], + &["mcp-method", "mcp-param-user"], + ); + + apply_header_config(&mut headers, &cfg, Some(&ds)); + + assert_eq!(headers[&http::HeaderName::from_static("mcp-method")], "tools/call"); + assert_eq!(headers[&http::HeaderName::from_static("mcp-param-user")], "computed"); + assert!(!headers.contains_key(&http::HeaderName::from_static("mcp-name"))); + assert!(!headers.contains_key(&http::HeaderName::from_static("mcp-protocol-version"))); + } + #[test] fn body_framing_and_connection_management_headers_cannot_be_forwarded() { let mut headers = HashMap::new(); diff --git a/crates/contextforge-data-plane-lib/src/layers/mcp_header_limits.rs b/crates/contextforge-data-plane-lib/src/layers/mcp_header_limits.rs new file mode 100644 index 00000000..9c547ed2 --- /dev/null +++ b/crates/contextforge-data-plane-lib/src/layers/mcp_header_limits.rs @@ -0,0 +1,239 @@ +use axum::{body::Body, extract::State, middleware::Next, response::Response}; +use http::{HeaderName, StatusCode, header}; + +use crate::common::{ + Config, DEFAULT_MCP_STANDARD_HEADER_MAX_COUNT, DEFAULT_MCP_STANDARD_HEADER_MAX_TOTAL_BYTES, + DEFAULT_MCP_STANDARD_HEADER_MAX_VALUE_BYTES, +}; + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub(crate) struct McpStandardHeaderLimits { + pub(crate) max_count: usize, + pub(crate) max_value_bytes: usize, + pub(crate) max_total_bytes: usize, +} + +impl McpStandardHeaderLimits { + pub(crate) fn from_config(config: &Config) -> Self { + Self { + max_count: configured_or_default( + config.mcp_standard_header_max_count, + DEFAULT_MCP_STANDARD_HEADER_MAX_COUNT, + ), + max_value_bytes: configured_or_default( + config.mcp_standard_header_max_value_bytes, + DEFAULT_MCP_STANDARD_HEADER_MAX_VALUE_BYTES, + ), + max_total_bytes: configured_or_default( + config.mcp_standard_header_max_total_bytes, + DEFAULT_MCP_STANDARD_HEADER_MAX_TOTAL_BYTES, + ), + } + } +} + +fn configured_or_default(configured: usize, default: usize) -> usize { + if configured == 0 { default } else { configured } +} + +pub(crate) async fn mcp_header_limits_layer( + State(limits): State, + request: http::Request, + next: Next, +) -> Response { + if exceeds_limits(request.headers(), limits) { + return Response::builder() + .status(StatusCode::REQUEST_HEADER_FIELDS_TOO_LARGE) + .header(header::CONTENT_TYPE, "text/plain") + .body(Body::from("MCP standard header limits exceeded")) + .expect("Expecting this to work"); + } + + next.run(request).await +} + +fn exceeds_limits(headers: &http::HeaderMap, limits: McpStandardHeaderLimits) -> bool { + let mut count = 0usize; + let mut total_bytes = 0usize; + + for (name, value) in headers.iter().filter(|(name, _)| is_mcp_limited_header_name(name)) { + count = count.saturating_add(1); + if count > limits.max_count { + return true; + } + + let value_bytes = value.as_bytes().len(); + if value_bytes > limits.max_value_bytes { + return true; + } + + // Application budget only: this is not exact HTTP/1 wire size and does + // not model HTTP/2 HPACK compression. + total_bytes = total_bytes.saturating_add(name.as_str().len()).saturating_add(value_bytes); + if total_bytes > limits.max_total_bytes { + return true; + } + } + + false +} + +pub(crate) fn is_mcp_limited_header_name(name: &HeaderName) -> bool { + is_exact_mcp_header_name(name, "mcp-method") + || is_exact_mcp_header_name(name, "mcp-name") + || is_exact_mcp_header_name(name, "mcp-protocol-version") + || is_exact_mcp_header_name(name, "mcp-session-id") + || is_mcp_param_header_name(name) +} + +pub(crate) fn is_mcp_computed_header_name(name: &HeaderName) -> bool { + is_exact_mcp_header_name(name, "mcp-method") + || is_exact_mcp_header_name(name, "mcp-name") + || is_exact_mcp_header_name(name, "mcp-protocol-version") + || is_mcp_param_header_name(name) +} + +fn is_exact_mcp_header_name(name: &HeaderName, expected: &str) -> bool { + name.as_str().eq_ignore_ascii_case(expected) +} + +fn is_mcp_param_header_name(name: &HeaderName) -> bool { + const PREFIX: &str = "mcp-param-"; + name.as_str().get(..PREFIX.len()).is_some_and(|prefix| prefix.eq_ignore_ascii_case(PREFIX)) +} + +#[cfg(test)] +mod tests { + use std::sync::Arc; + + use async_trait::async_trait; + use axum::{Router, body::Body, middleware, response::Response, routing::get}; + use contextforge_data_plane_apis::{User, user_store::UserConfig}; + use http::{Request, StatusCode}; + use tower::ServiceExt; + + use crate::{ + Config, + common::{ContextForgeDataPlaneAppState, JwtTokenDecoders}, + layers::{ + claims_id::claims_layer, + mcp_header_limits::{McpStandardHeaderLimits, mcp_header_limits_layer}, + }, + user_config_store::{ConfigStoreError, UserConfigStore}, + }; + + async fn ok() -> Response { + Response::builder().status(StatusCode::OK).body(Body::empty()).expect("Expecting this to work") + } + + fn app(limits: McpStandardHeaderLimits) -> Router { + Router::new().route("/", get(ok)).layer(middleware::from_fn_with_state(limits, mcp_header_limits_layer)) + } + + fn request_with_headers(headers: &[(&str, &str)]) -> Request { + let mut builder = Request::builder().uri("/"); + for (name, value) in headers { + builder = builder.header(*name, *value); + } + builder.body(Body::empty()).expect("Expecting this to work") + } + + #[tokio::test] + async fn rejects_too_many_mcp_headers() { + let limits = McpStandardHeaderLimits { max_count: 2, max_value_bytes: 1024, max_total_bytes: 4096 }; + let response = app(limits) + .oneshot(request_with_headers(&[ + ("Mcp-Method", "tools/call"), + ("Mcp-Name", "example"), + ("Mcp-Param-User", "alice"), + ])) + .await + .expect("Expecting this to work"); + + assert_eq!(response.status(), StatusCode::REQUEST_HEADER_FIELDS_TOO_LARGE); + } + + #[tokio::test] + async fn rejects_oversized_mcp_header_value() { + let limits = McpStandardHeaderLimits { max_count: 32, max_value_bytes: 4, max_total_bytes: 4096 }; + let response = app(limits) + .oneshot(request_with_headers(&[("Mcp-Param-User", "alice")])) + .await + .expect("Expecting this to work"); + + assert_eq!(response.status(), StatusCode::REQUEST_HEADER_FIELDS_TOO_LARGE); + } + + #[tokio::test] + async fn rejects_excessive_total_mcp_header_bytes() { + let limits = McpStandardHeaderLimits { max_count: 32, max_value_bytes: 16, max_total_bytes: 24 }; + let response = app(limits) + .oneshot(request_with_headers(&[("Mcp-Method", "tools/call"), ("Mcp-Name", "example")])) + .await + .expect("Expecting this to work"); + + assert_eq!(response.status(), StatusCode::REQUEST_HEADER_FIELDS_TOO_LARGE); + } + + #[tokio::test] + async fn counts_mcp_headers_case_insensitively() { + let limits = McpStandardHeaderLimits { max_count: 1, max_value_bytes: 1024, max_total_bytes: 4096 }; + let response = app(limits) + .oneshot(request_with_headers(&[("McP-MeThOd", "tools/call"), ("mCp-PaRaM-User", "alice")])) + .await + .expect("Expecting this to work"); + + assert_eq!(response.status(), StatusCode::REQUEST_HEADER_FIELDS_TOO_LARGE); + } + + #[tokio::test] + async fn ignores_non_mcp_headers_for_mcp_specific_budget() { + let limits = McpStandardHeaderLimits { max_count: 1, max_value_bytes: 1024, max_total_bytes: 4096 }; + let response = app(limits) + .oneshot(request_with_headers(&[ + ("X-One", "1"), + ("X-Two", "2"), + ("X-Three", "3"), + ("Mcp-Method", "tools/call"), + ])) + .await + .expect("Expecting this to work"); + + assert_eq!(response.status(), StatusCode::OK); + } + + #[derive(Clone)] + struct UnusedConfigStore; + + #[async_trait] + impl UserConfigStore for UnusedConfigStore { + async fn get_config<'a>(&self, _key: &'a User) -> Result { + unreachable!("mcp header limit rejection must run before config lookup") + } + + async fn set_config<'a>(&self, _key: &'a User, _user_config: &'a UserConfig) -> Result<(), ConfigStoreError> { + unreachable!("mcp header limit rejection must run before config lookup") + } + } + + #[tokio::test] + async fn rejects_excessive_mcp_headers_before_auth() { + let limits = McpStandardHeaderLimits { max_count: 1, max_value_bytes: 1024, max_total_bytes: 4096 }; + let state = ContextForgeDataPlaneAppState { + jwt_token_decoding_keys: JwtTokenDecoders { rs: None, hmac_sha: None }, + config_store: Arc::new(UnusedConfigStore), + config: Config::default(), + }; + let app = Router::new() + .route("/", get(ok)) + .layer(middleware::from_fn_with_state(state, claims_layer)) + .layer(middleware::from_fn_with_state(limits, mcp_header_limits_layer)); + + let response = app + .oneshot(request_with_headers(&[("Mcp-Method", "tools/call"), ("Mcp-Name", "example")])) + .await + .expect("Expecting this to work"); + + assert_eq!(response.status(), StatusCode::REQUEST_HEADER_FIELDS_TOO_LARGE); + } +} diff --git a/crates/contextforge-data-plane-lib/src/layers/mod.rs b/crates/contextforge-data-plane-lib/src/layers/mod.rs index e4b52754..83af1e2f 100644 --- a/crates/contextforge-data-plane-lib/src/layers/mod.rs +++ b/crates/contextforge-data-plane-lib/src/layers/mod.rs @@ -1,4 +1,5 @@ pub mod claims_id; +pub mod mcp_header_limits; pub mod mcp_origin; pub mod session_id; pub mod user_config_store; diff --git a/crates/contextforge-data-plane-lib/src/lib.rs b/crates/contextforge-data-plane-lib/src/lib.rs index 4a7f9b32..fe3d1ca4 100644 --- a/crates/contextforge-data-plane-lib/src/lib.rs +++ b/crates/contextforge-data-plane-lib/src/lib.rs @@ -41,6 +41,7 @@ use crate::{ gateway::LocalUserSessionStore, layers::{ claims_id::claims_layer, + mcp_header_limits::{McpStandardHeaderLimits, mcp_header_limits_layer}, mcp_origin::mcp_origin_layer, session_id::{SessionIdState, session_id_layer}, user_config_store::user_config_store_layer, @@ -137,6 +138,7 @@ impl Gateway { config_store: Arc::clone(&user_config_store), config: config.clone(), }; + let mcp_standard_header_limits = McpStandardHeaderLimits::from_config(config); let app = axum::Router::new() .nest_service("/servers/{virtual_host_name}/mcp", mcp_service) @@ -145,6 +147,9 @@ impl Gateway { .layer(middleware::from_fn_with_state(session_id_state, session_id_layer)) .layer(middleware::from_fn_with_state(mcp_add_state.clone(), claims_layer)) .layer(middleware::from_fn(virtual_host_id_layer)) + // Keep this outside auth/config/RMCP work so oversized MCP headers + // are rejected before JWT validation or body parsing. + .layer(middleware::from_fn_with_state(mcp_standard_header_limits, mcp_header_limits_layer)) .layer(cors_layer) // mcp_origin_layer is the outermost wrapper: fires before JWT auth, // session creation, and backend fan-out. From 0369a2dbe981bd2bb231755cc1c6f9a85de451fa Mon Sep 17 00:00:00 2001 From: Pratik Gandhi Date: Tue, 11 Aug 2026 13:49:38 +0100 Subject: [PATCH 2/2] Address MCP header guard review Signed-off-by: Pratik Gandhi --- .../src/gateway/mcp_service/initialization.rs | 4 +- .../src/layers/mcp_header_limits.rs | 88 +++++++-------- crates/contextforge-data-plane-lib/src/lib.rs | 103 ++++++++++++++---- .../src/mcp_standard_headers.rs | 25 +++++ 4 files changed, 147 insertions(+), 73 deletions(-) create mode 100644 crates/contextforge-data-plane-lib/src/mcp_standard_headers.rs diff --git a/crates/contextforge-data-plane-lib/src/gateway/mcp_service/initialization.rs b/crates/contextforge-data-plane-lib/src/gateway/mcp_service/initialization.rs index 448c56fa..9bb3892f 100644 --- a/crates/contextforge-data-plane-lib/src/gateway/mcp_service/initialization.rs +++ b/crates/contextforge-data-plane-lib/src/gateway/mcp_service/initialization.rs @@ -17,7 +17,7 @@ use crate::gateway::{ mcp_call_validator::InitializeCallValidator, session_store::{UserSession, UserSessionStore}, }; -use crate::layers::mcp_header_limits::is_mcp_computed_header_name; +use crate::mcp_standard_headers; pub(super) async fn initialize( mcp_service: &McpService, @@ -272,7 +272,7 @@ fn is_protected_header(name: &http::HeaderName) -> bool { "accept", "last-event-id", ]; - PROTECTED.iter().any(|&p| name.as_str().eq_ignore_ascii_case(p)) || is_mcp_computed_header_name(name) + PROTECTED.iter().any(|&p| name.as_str().eq_ignore_ascii_case(p)) || mcp_standard_headers::is_computed(name) } #[cfg(test)] diff --git a/crates/contextforge-data-plane-lib/src/layers/mcp_header_limits.rs b/crates/contextforge-data-plane-lib/src/layers/mcp_header_limits.rs index 9c547ed2..0fc65f23 100644 --- a/crates/contextforge-data-plane-lib/src/layers/mcp_header_limits.rs +++ b/crates/contextforge-data-plane-lib/src/layers/mcp_header_limits.rs @@ -1,30 +1,29 @@ use axum::{body::Body, extract::State, middleware::Next, response::Response}; -use http::{HeaderName, StatusCode, header}; +use http::{StatusCode, header}; +use tracing::debug; use crate::common::{ Config, DEFAULT_MCP_STANDARD_HEADER_MAX_COUNT, DEFAULT_MCP_STANDARD_HEADER_MAX_TOTAL_BYTES, DEFAULT_MCP_STANDARD_HEADER_MAX_VALUE_BYTES, }; +use crate::mcp_standard_headers; #[derive(Clone, Copy, Debug, PartialEq, Eq)] pub(crate) struct McpStandardHeaderLimits { - pub(crate) max_count: usize, - pub(crate) max_value_bytes: usize, - pub(crate) max_total_bytes: usize, + pub(crate) count: usize, + pub(crate) value_bytes: usize, + pub(crate) total_bytes: usize, } impl McpStandardHeaderLimits { pub(crate) fn from_config(config: &Config) -> Self { Self { - max_count: configured_or_default( - config.mcp_standard_header_max_count, - DEFAULT_MCP_STANDARD_HEADER_MAX_COUNT, - ), - max_value_bytes: configured_or_default( + count: configured_or_default(config.mcp_standard_header_max_count, DEFAULT_MCP_STANDARD_HEADER_MAX_COUNT), + value_bytes: configured_or_default( config.mcp_standard_header_max_value_bytes, DEFAULT_MCP_STANDARD_HEADER_MAX_VALUE_BYTES, ), - max_total_bytes: configured_or_default( + total_bytes: configured_or_default( config.mcp_standard_header_max_total_bytes, DEFAULT_MCP_STANDARD_HEADER_MAX_TOTAL_BYTES, ), @@ -36,12 +35,25 @@ fn configured_or_default(configured: usize, default: usize) -> usize { if configured == 0 { default } else { configured } } +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +struct McpStandardHeaderUsage { + count: usize, + value_bytes: usize, + total_bytes: usize, +} + pub(crate) async fn mcp_header_limits_layer( State(limits): State, request: http::Request, next: Next, ) -> Response { - if exceeds_limits(request.headers(), limits) { + if let Some(usage) = exceeded_limits(request.headers(), limits) { + let count = usage.count; + let value_bytes = usage.value_bytes; + let total_bytes = usage.total_bytes; + debug!( + "mcp_header_limits_layer - rejecting request count = {count} value_bytes = {value_bytes} total_bytes = {total_bytes}" + ); return Response::builder() .status(StatusCode::REQUEST_HEADER_FIELDS_TOO_LARGE) .header(header::CONTENT_TYPE, "text/plain") @@ -52,54 +64,30 @@ pub(crate) async fn mcp_header_limits_layer( next.run(request).await } -fn exceeds_limits(headers: &http::HeaderMap, limits: McpStandardHeaderLimits) -> bool { +fn exceeded_limits(headers: &http::HeaderMap, limits: McpStandardHeaderLimits) -> Option { let mut count = 0usize; let mut total_bytes = 0usize; - for (name, value) in headers.iter().filter(|(name, _)| is_mcp_limited_header_name(name)) { + for (name, value) in headers.iter().filter(|(name, _)| mcp_standard_headers::is_limited(name)) { count = count.saturating_add(1); - if count > limits.max_count { - return true; + if count > limits.count { + return Some(McpStandardHeaderUsage { count, value_bytes: 0, total_bytes }); } let value_bytes = value.as_bytes().len(); - if value_bytes > limits.max_value_bytes { - return true; + if value_bytes > limits.value_bytes { + return Some(McpStandardHeaderUsage { count, value_bytes, total_bytes }); } // Application budget only: this is not exact HTTP/1 wire size and does // not model HTTP/2 HPACK compression. total_bytes = total_bytes.saturating_add(name.as_str().len()).saturating_add(value_bytes); - if total_bytes > limits.max_total_bytes { - return true; + if total_bytes > limits.total_bytes { + return Some(McpStandardHeaderUsage { count, value_bytes, total_bytes }); } } - false -} - -pub(crate) fn is_mcp_limited_header_name(name: &HeaderName) -> bool { - is_exact_mcp_header_name(name, "mcp-method") - || is_exact_mcp_header_name(name, "mcp-name") - || is_exact_mcp_header_name(name, "mcp-protocol-version") - || is_exact_mcp_header_name(name, "mcp-session-id") - || is_mcp_param_header_name(name) -} - -pub(crate) fn is_mcp_computed_header_name(name: &HeaderName) -> bool { - is_exact_mcp_header_name(name, "mcp-method") - || is_exact_mcp_header_name(name, "mcp-name") - || is_exact_mcp_header_name(name, "mcp-protocol-version") - || is_mcp_param_header_name(name) -} - -fn is_exact_mcp_header_name(name: &HeaderName, expected: &str) -> bool { - name.as_str().eq_ignore_ascii_case(expected) -} - -fn is_mcp_param_header_name(name: &HeaderName) -> bool { - const PREFIX: &str = "mcp-param-"; - name.as_str().get(..PREFIX.len()).is_some_and(|prefix| prefix.eq_ignore_ascii_case(PREFIX)) + None } #[cfg(test)] @@ -140,7 +128,7 @@ mod tests { #[tokio::test] async fn rejects_too_many_mcp_headers() { - let limits = McpStandardHeaderLimits { max_count: 2, max_value_bytes: 1024, max_total_bytes: 4096 }; + let limits = McpStandardHeaderLimits { count: 2, value_bytes: 1024, total_bytes: 4096 }; let response = app(limits) .oneshot(request_with_headers(&[ ("Mcp-Method", "tools/call"), @@ -155,7 +143,7 @@ mod tests { #[tokio::test] async fn rejects_oversized_mcp_header_value() { - let limits = McpStandardHeaderLimits { max_count: 32, max_value_bytes: 4, max_total_bytes: 4096 }; + let limits = McpStandardHeaderLimits { count: 32, value_bytes: 4, total_bytes: 4096 }; let response = app(limits) .oneshot(request_with_headers(&[("Mcp-Param-User", "alice")])) .await @@ -166,7 +154,7 @@ mod tests { #[tokio::test] async fn rejects_excessive_total_mcp_header_bytes() { - let limits = McpStandardHeaderLimits { max_count: 32, max_value_bytes: 16, max_total_bytes: 24 }; + let limits = McpStandardHeaderLimits { count: 32, value_bytes: 16, total_bytes: 24 }; let response = app(limits) .oneshot(request_with_headers(&[("Mcp-Method", "tools/call"), ("Mcp-Name", "example")])) .await @@ -177,7 +165,7 @@ mod tests { #[tokio::test] async fn counts_mcp_headers_case_insensitively() { - let limits = McpStandardHeaderLimits { max_count: 1, max_value_bytes: 1024, max_total_bytes: 4096 }; + let limits = McpStandardHeaderLimits { count: 1, value_bytes: 1024, total_bytes: 4096 }; let response = app(limits) .oneshot(request_with_headers(&[("McP-MeThOd", "tools/call"), ("mCp-PaRaM-User", "alice")])) .await @@ -188,7 +176,7 @@ mod tests { #[tokio::test] async fn ignores_non_mcp_headers_for_mcp_specific_budget() { - let limits = McpStandardHeaderLimits { max_count: 1, max_value_bytes: 1024, max_total_bytes: 4096 }; + let limits = McpStandardHeaderLimits { count: 1, value_bytes: 1024, total_bytes: 4096 }; let response = app(limits) .oneshot(request_with_headers(&[ ("X-One", "1"), @@ -218,7 +206,7 @@ mod tests { #[tokio::test] async fn rejects_excessive_mcp_headers_before_auth() { - let limits = McpStandardHeaderLimits { max_count: 1, max_value_bytes: 1024, max_total_bytes: 4096 }; + let limits = McpStandardHeaderLimits { count: 1, value_bytes: 1024, total_bytes: 4096 }; let state = ContextForgeDataPlaneAppState { jwt_token_decoding_keys: JwtTokenDecoders { rs: None, hmac_sha: None }, config_store: Arc::new(UnusedConfigStore), diff --git a/crates/contextforge-data-plane-lib/src/lib.rs b/crates/contextforge-data-plane-lib/src/lib.rs index fe3d1ca4..4464b02f 100644 --- a/crates/contextforge-data-plane-lib/src/lib.rs +++ b/crates/contextforge-data-plane-lib/src/lib.rs @@ -14,6 +14,7 @@ mod common; mod const_values; mod gateway; mod layers; +mod mcp_standard_headers; mod telemetry; mod transports; @@ -68,10 +69,28 @@ pub struct Gateway { impl Gateway { pub async fn run_gateway(self) -> Result<()> { - let config = &self.config; - let session_manager = self.session_manager; - let user_config_store = match self.user_config_store_type { - UserConfigStoreType::Redis => Arc::new(get_config_store(config).await?), + let config = self.config.clone(); + let app = self.build_app().await?; + + let mut handlers = vec![]; + + if let Some(tcp) = Option::::try_from(&config)? { + handlers.push(tcp.handle_tcp(app.clone()).boxed()); + } + + if let Some(tls) = Option::::try_from(&config)? { + handlers.push(tls.handle_tls(app.clone()).boxed()); + } + + let _ = futures::future::join_all(handlers).await; + + Ok(()) + } + + async fn build_app(self) -> Result { + let Gateway { config, session_manager, user_config_store_type, plugin_runtime } = self; + let user_config_store = match user_config_store_type { + UserConfigStoreType::Redis => Arc::new(get_config_store(&config).await?), UserConfigStoreType::Test(store) => store, }; let user_config_store = user_config_store as Arc; @@ -82,7 +101,6 @@ impl Gateway { user_session_store: Arc::new(user_session_store.clone()), backend_transports: backend_transports.clone(), }; - let mcp_plugin_runtime = self.plugin_runtime; // mcp_origin_layer is the sole enforcement point; disable RMCP's built-in checks. // Pass the host list to RMCP as well when configured (defense-in-depth). @@ -94,7 +112,7 @@ impl Gateway { StreamableHttpServerConfig::default().disable_allowed_hosts().disable_allowed_origins() }; - let reqwest_backend_client = reqwest::Client::try_from(config)?; + let reqwest_backend_client = reqwest::Client::try_from(&config)?; // Create streamable HTTP service let mcp_service: StreamableHttpService, LocalSessionManager> = @@ -104,7 +122,7 @@ impl Gateway { .with_user_session_store(user_session_store.clone()) .with_http_client(reqwest_backend_client.clone()) .with_transports(backend_transports.clone()) - .with_plugin_runtime(mcp_plugin_runtime.clone()) + .with_plugin_runtime(plugin_runtime.clone()) .build()) }, session_manager, @@ -138,7 +156,7 @@ impl Gateway { config_store: Arc::clone(&user_config_store), config: config.clone(), }; - let mcp_standard_header_limits = McpStandardHeaderLimits::from_config(config); + let mcp_standard_header_limits = McpStandardHeaderLimits::from_config(&config); let app = axum::Router::new() .nest_service("/servers/{virtual_host_name}/mcp", mcp_service) @@ -164,19 +182,7 @@ impl Gateway { .layer(TraceLayer::new_for_http().make_span_with(telemetry::ExtractingMakeSpan)) .layer(HttpMetricsLayerBuilder::new().build()); - let mut handlers = vec![]; - - if let Some(tcp) = Option::::try_from(config)? { - handlers.push(tcp.handle_tcp(app.clone()).boxed()); - } - - if let Some(tls) = Option::::try_from(config)? { - handlers.push(tls.handle_tls(app.clone()).boxed()); - } - - let _ = futures::future::join_all(handlers).await; - - Ok(()) + Ok(app) } } @@ -185,3 +191,58 @@ pub async fn get_config_store(config: &Config) -> Result { let cache_expiry = std::time::Duration::from_secs(config.user_config_cache_expiry_seconds); RedisUserConfigStore::new(&RedisClient::try_from(redis_config)?, cache_expiry).await } + +#[cfg(test)] +mod tests { + use std::sync::Arc; + + use async_trait::async_trait; + use axum::body::Body; + use contextforge_data_plane_apis::{User, user_store::UserConfig}; + use http::{Request, StatusCode}; + use rmcp::transport::streamable_http_server::session::local::LocalSessionManager; + use tower::ServiceExt; + + use crate::{ + Config, Gateway, UserConfigStoreType, + user_config_store::{ConfigStoreError, UserConfigStore}, + }; + + #[derive(Clone)] + struct UnusedConfigStore; + + #[async_trait] + impl UserConfigStore for UnusedConfigStore { + async fn get_config<'a>(&self, _key: &'a User) -> Result { + unreachable!("mcp header limit rejection must run before config lookup") + } + + async fn set_config<'a>(&self, _key: &'a User, _user_config: &'a UserConfig) -> Result<(), ConfigStoreError> { + unreachable!("mcp header limit rejection must run before config write") + } + } + + #[tokio::test] + async fn production_router_rejects_excessive_mcp_headers_before_auth() { + let config = Config { mcp_standard_header_max_count: 1, ..Config::default() }; + let app = Gateway::builder() + .with_config(config) + .with_session_manager(Arc::new(LocalSessionManager::default())) + .with_user_config_store_type(UserConfigStoreType::Test(Arc::new(UnusedConfigStore))) + .build() + .build_app() + .await + .expect("Expecting this to work"); + let request = Request::builder() + .method("POST") + .uri("/contextforge-rs/servers/test-vhost/mcp") + .header("Mcp-Method", "tools/call") + .header("Mcp-Name", "example") + .body(Body::empty()) + .expect("Expecting this to work"); + + let response = app.oneshot(request).await.expect("Expecting this to work"); + + assert_eq!(response.status(), StatusCode::REQUEST_HEADER_FIELDS_TOO_LARGE); + } +} diff --git a/crates/contextforge-data-plane-lib/src/mcp_standard_headers.rs b/crates/contextforge-data-plane-lib/src/mcp_standard_headers.rs new file mode 100644 index 00000000..7b3c867a --- /dev/null +++ b/crates/contextforge-data-plane-lib/src/mcp_standard_headers.rs @@ -0,0 +1,25 @@ +use http::HeaderName; + +pub(crate) fn is_limited(name: &HeaderName) -> bool { + is_exact(name, "mcp-method") + || is_exact(name, "mcp-name") + || is_exact(name, "mcp-protocol-version") + || is_exact(name, "mcp-session-id") + || is_param(name) +} + +pub(crate) fn is_computed(name: &HeaderName) -> bool { + is_exact(name, "mcp-method") + || is_exact(name, "mcp-name") + || is_exact(name, "mcp-protocol-version") + || is_param(name) +} + +fn is_exact(name: &HeaderName, expected: &str) -> bool { + name.as_str().eq_ignore_ascii_case(expected) +} + +fn is_param(name: &HeaderName) -> bool { + const PREFIX: &str = "mcp-param-"; + name.as_str().get(..PREFIX.len()).is_some_and(|prefix| prefix.eq_ignore_ascii_case(PREFIX)) +}