From 151bfb8da696784ad04ddb042f89fef53a5732f5 Mon Sep 17 00:00:00 2001 From: Dale Seo Date: Wed, 28 Jan 2026 11:20:30 -0500 Subject: [PATCH 1/3] fix: add Host header validation to prevent DNS rebinding attacks --- .../apollo-mcp-server/src/host_validation.rs | 435 ++++++++++++++++++ crates/apollo-mcp-server/src/lib.rs | 1 + crates/apollo-mcp-server/src/server.rs | 6 + .../src/server/states/starting.rs | 26 +- 4 files changed, 458 insertions(+), 10 deletions(-) create mode 100644 crates/apollo-mcp-server/src/host_validation.rs diff --git a/crates/apollo-mcp-server/src/host_validation.rs b/crates/apollo-mcp-server/src/host_validation.rs new file mode 100644 index 0000000000..8f1dda92a3 --- /dev/null +++ b/crates/apollo-mcp-server/src/host_validation.rs @@ -0,0 +1,435 @@ +use std::borrow::Cow; +use std::net::IpAddr; +use std::sync::Arc; + +use axum::{ + body::Body, + extract::{Request, State}, + http::{HeaderValue, StatusCode, header::HOST}, + middleware::Next, + response::{IntoResponse, Response}, +}; +use schemars::JsonSchema; +use serde::Deserialize; +use tracing::warn; + +/// Configuration for Host header validation to prevent DNS rebinding attacks. +#[derive(Debug, Clone, Deserialize, JsonSchema)] +#[serde(default)] +pub struct HostValidationConfig { + /// Enable Host header validation (enabled by default for security) + pub enabled: bool, + + /// Additional allowed hosts beyond localhost, 127.0.0.1, ::1, and 0.0.0.0. + pub allowed_hosts: Vec, +} + +impl Default for HostValidationConfig { + fn default() -> Self { + Self { + enabled: true, + allowed_hosts: Vec::new(), + } + } +} + +impl HostValidationConfig { + /// Creates a configuration with Host header validation disabled. + #[must_use] + pub fn disabled() -> Self { + Self { + enabled: false, + allowed_hosts: Vec::new(), + } + } +} + +/// State for the Host header validation middleware. +#[derive(Clone)] +pub struct HostValidationState { + /// The validation configuration (wrapped in Arc to avoid cloning Vec on each request). + pub config: Arc, + /// The port the server is listening on, used to validate localhost requests. + pub server_port: u16, +} + +impl HostValidationState { + fn is_host_allowed(&self, host: &str) -> bool { + if !self.config.enabled { + return true; + } + + let hostname = host + .rsplit_once(':') + .map(|(h, _)| h) + .unwrap_or(host) + .trim_start_matches('[') + .trim_end_matches(']'); + + // Check if hostname is localhost: literal "localhost", loopback (127.0.0.1, ::1), or unspecified (0.0.0.0, ::) + let is_localhost = hostname.eq_ignore_ascii_case("localhost") + || hostname + .parse::() + .map(|ip| ip.is_loopback() || ip.is_unspecified()) + .unwrap_or(false); + + // Localhost: validate port against actual server port + if is_localhost { + if let Some(port_str) = host.rsplit_once(':').map(|(_, p)| p) { + if let Ok(port) = port_str.parse::() { + return port == self.server_port; + } + return false; + } + return true; + } + + // Custom hosts: validate port against config (if specified). + // No port in config means any port is allowed for flexibility with proxies. + for allowed in &self.config.allowed_hosts { + let allowed_hostname = allowed.rsplit_once(':').map(|(h, _)| h).unwrap_or(allowed); + + if hostname.eq_ignore_ascii_case(allowed_hostname) { + if let Some(allowed_port_str) = allowed.rsplit_once(':').map(|(_, p)| p) { + if let Some(host_port_str) = host.rsplit_once(':').map(|(_, p)| p) { + return allowed_port_str == host_port_str; + } + return false; + } + return true; + } + } + + false + } +} + +/// Middleware that validates the Host header to prevent DNS rebinding attacks. +pub async fn validate_host( + State(state): State, + request: Request, + next: Next, +) -> Response { + if !state.config.enabled { + return next.run(request).await; + } + + // Extract host from Host header (HTTP/1.1) or URI authority (HTTP/2). + // Use Cow to avoid allocation when Host header is present (common case). + let host: Option> = request + .headers() + .get(HOST) + .and_then(|v| v.to_str().ok()) + .map(Cow::Borrowed) + .or_else(|| { + request.uri().host().map(|h| { + // Include port from URI if present (requires allocation) + match request.uri().port_u16() { + Some(port) => Cow::Owned(format!("{}:{}", h, port)), + None => Cow::Borrowed(h), + } + }) + }); + + match host { + Some(host) => { + if state.is_host_allowed(&host) { + next.run(request).await + } else { + warn!( + host = %host, + "Rejected request with invalid Host header (possible DNS rebinding attack)" + ); + forbidden_response() + } + } + None => { + warn!("Rejected request without Host header"); + forbidden_response() + } + } +} + +fn forbidden_response() -> Response { + ( + StatusCode::FORBIDDEN, + [( + http::header::CONTENT_TYPE, + HeaderValue::from_static("text/plain"), + )], + Body::from("Forbidden: Invalid Host header"), + ) + .into_response() +} + +#[cfg(test)] +mod tests { + use super::*; + use axum::{Router, routing::get}; + use http::{Method, Request, StatusCode}; + use tower::util::ServiceExt; + + fn test_router(config: HostValidationConfig, port: u16) -> Router { + Router::new().route("/test", get(|| async { "ok" })).layer( + axum::middleware::from_fn_with_state( + HostValidationState { + config: Arc::new(config), + server_port: port, + }, + validate_host, + ), + ) + } + + #[tokio::test] + async fn test_allows_localhost() { + let config = HostValidationConfig::default(); + let app = test_router(config, 8000); + + let request = Request::builder() + .method(Method::GET) + .uri("/test") + .header("Host", "localhost:8000") + .body(Body::empty()) + .unwrap(); + + let response = app.oneshot(request).await.unwrap(); + assert_eq!(response.status(), StatusCode::OK); + } + + #[tokio::test] + async fn test_allows_localhost_without_port() { + let config = HostValidationConfig::default(); + let app = test_router(config, 8000); + + let request = Request::builder() + .method(Method::GET) + .uri("/test") + .header("Host", "localhost") + .body(Body::empty()) + .unwrap(); + + let response = app.oneshot(request).await.unwrap(); + assert_eq!(response.status(), StatusCode::OK); + } + + #[tokio::test] + async fn test_allows_127_0_0_1() { + let config = HostValidationConfig::default(); + let app = test_router(config, 8000); + + let request = Request::builder() + .method(Method::GET) + .uri("/test") + .header("Host", "127.0.0.1:8000") + .body(Body::empty()) + .unwrap(); + + let response = app.oneshot(request).await.unwrap(); + assert_eq!(response.status(), StatusCode::OK); + } + + #[tokio::test] + async fn test_allows_ipv6_localhost() { + let config = HostValidationConfig::default(); + let app = test_router(config, 8000); + + let request = Request::builder() + .method(Method::GET) + .uri("/test") + .header("Host", "[::1]:8000") + .body(Body::empty()) + .unwrap(); + + let response = app.oneshot(request).await.unwrap(); + assert_eq!(response.status(), StatusCode::OK); + } + + #[tokio::test] + async fn test_allows_0_0_0_0() { + let config = HostValidationConfig::default(); + let app = test_router(config, 8000); + + let request = Request::builder() + .method(Method::GET) + .uri("/test") + .header("Host", "0.0.0.0:8000") + .body(Body::empty()) + .unwrap(); + + let response = app.oneshot(request).await.unwrap(); + assert_eq!(response.status(), StatusCode::OK); + } + + #[tokio::test] + async fn test_rejects_attacker_host() { + let config = HostValidationConfig::default(); + let app = test_router(config, 8000); + + let request = Request::builder() + .method(Method::GET) + .uri("/test") + .header("Host", "attacker.com") + .body(Body::empty()) + .unwrap(); + + let response = app.oneshot(request).await.unwrap(); + assert_eq!(response.status(), StatusCode::FORBIDDEN); + } + + #[tokio::test] + async fn test_rejects_attacker_host_with_port() { + let config = HostValidationConfig::default(); + let app = test_router(config, 8000); + + let request = Request::builder() + .method(Method::GET) + .uri("/test") + .header("Host", "attacker.com:8000") + .body(Body::empty()) + .unwrap(); + + let response = app.oneshot(request).await.unwrap(); + assert_eq!(response.status(), StatusCode::FORBIDDEN); + } + + #[tokio::test] + async fn test_rejects_wrong_port() { + let config = HostValidationConfig::default(); + let app = test_router(config, 8000); + + let request = Request::builder() + .method(Method::GET) + .uri("/test") + .header("Host", "localhost:9999") + .body(Body::empty()) + .unwrap(); + + let response = app.oneshot(request).await.unwrap(); + assert_eq!(response.status(), StatusCode::FORBIDDEN); + } + + #[tokio::test] + async fn test_disabled_allows_any_host() { + let config = HostValidationConfig::disabled(); + let app = test_router(config, 8000); + + let request = Request::builder() + .method(Method::GET) + .uri("/test") + .header("Host", "attacker.com") + .body(Body::empty()) + .unwrap(); + + let response = app.oneshot(request).await.unwrap(); + assert_eq!(response.status(), StatusCode::OK); + } + + #[tokio::test] + async fn test_custom_allowed_host() { + let config = HostValidationConfig { + enabled: true, + allowed_hosts: vec!["mcp.test.com".to_string()], + }; + let app = test_router(config, 8000); + + let request = Request::builder() + .method(Method::GET) + .uri("/test") + .header("Host", "mcp.test.com") + .body(Body::empty()) + .unwrap(); + + let response = app.oneshot(request).await.unwrap(); + assert_eq!(response.status(), StatusCode::OK); + } + + #[tokio::test] + async fn test_custom_allowed_host_with_port() { + let config = HostValidationConfig { + enabled: true, + allowed_hosts: vec!["mcp.test.com:8000".to_string()], + }; + let app = test_router(config, 8000); + + let request = Request::builder() + .method(Method::GET) + .uri("/test") + .header("Host", "mcp.test.com:8000") + .body(Body::empty()) + .unwrap(); + + let response = app.oneshot(request).await.unwrap(); + assert_eq!(response.status(), StatusCode::OK); + } + + #[tokio::test] + async fn test_custom_allowed_host_wrong_port() { + let config = HostValidationConfig { + enabled: true, + allowed_hosts: vec!["mcp.test.com:8000".to_string()], + }; + let app = test_router(config, 8000); + + let request = Request::builder() + .method(Method::GET) + .uri("/test") + .header("Host", "mcp.test.com:9000") + .body(Body::empty()) + .unwrap(); + + let response = app.oneshot(request).await.unwrap(); + assert_eq!(response.status(), StatusCode::FORBIDDEN); + } + + #[tokio::test] + async fn test_case_insensitive_hostname() { + let config = HostValidationConfig::default(); + let app = test_router(config, 8000); + + let request = Request::builder() + .method(Method::GET) + .uri("/test") + .header("Host", "LOCALHOST:8000") + .body(Body::empty()) + .unwrap(); + + let response = app.oneshot(request).await.unwrap(); + assert_eq!(response.status(), StatusCode::OK); + } + + #[test] + fn test_is_host_allowed() { + let state = HostValidationState { + config: Arc::new(HostValidationConfig::default()), + server_port: 8000, + }; + + assert!(state.is_host_allowed("localhost")); + assert!(state.is_host_allowed("localhost:8000")); + assert!(state.is_host_allowed("127.0.0.1:8000")); + assert!(state.is_host_allowed("[::1]:8000")); + + assert!(!state.is_host_allowed("localhost:9999")); + + assert!(!state.is_host_allowed("attacker.com")); + assert!(!state.is_host_allowed("attacker.com:8000")); + } + + #[test] + fn test_default_config_is_enabled() { + let config = HostValidationConfig::default(); + assert!(config.enabled); + assert!(config.allowed_hosts.is_empty()); + } + + #[test] + fn test_disabled_config() { + let state = HostValidationState { + config: Arc::new(HostValidationConfig::disabled()), + server_port: 8000, + }; + assert!(!state.config.enabled); + assert!(state.is_host_allowed("attacker.com")); + } +} diff --git a/crates/apollo-mcp-server/src/lib.rs b/crates/apollo-mcp-server/src/lib.rs index 327397cca8..96ce77aefe 100644 --- a/crates/apollo-mcp-server/src/lib.rs +++ b/crates/apollo-mcp-server/src/lib.rs @@ -11,6 +11,7 @@ mod explorer; mod graphql; pub mod headers; pub mod health; +pub mod host_validation; mod introspection; pub(crate) mod json_schema; pub(crate) mod meter; diff --git a/crates/apollo-mcp-server/src/server.rs b/crates/apollo-mcp-server/src/server.rs index 385676a568..b313228037 100644 --- a/crates/apollo-mcp-server/src/server.rs +++ b/crates/apollo-mcp-server/src/server.rs @@ -14,6 +14,7 @@ use crate::errors::ServerError; use crate::event::Event as ServerEvent; use crate::headers::ForwardHeaders; use crate::health::HealthCheckConfig; +use crate::host_validation::HostValidationConfig; use crate::operations::{MutationMode, OperationSource}; use crate::server_info::ServerInfoConfig; @@ -88,8 +89,13 @@ pub enum Transport { #[serde(default = "Transport::default_port")] port: u16, + /// Enable stateful mode for session management #[serde(default = "Transport::default_stateful_mode")] stateful_mode: bool, + + /// Host header validation configuration for DNS rebinding protection. + #[serde(default)] + host_validation: HostValidationConfig, }, } diff --git a/crates/apollo-mcp-server/src/server/states/starting.rs b/crates/apollo-mcp-server/src/server/states/starting.rs index c61136e7f8..2701cc75a7 100644 --- a/crates/apollo-mcp-server/src/server/states/starting.rs +++ b/crates/apollo-mcp-server/src/server/states/starting.rs @@ -11,6 +11,7 @@ use tokio::sync::RwLock; use tokio_util::sync::CancellationToken; use tracing::{debug, error, info}; +use crate::host_validation::{HostValidationState, validate_host}; use crate::server::states::telemetry::otel_context_middleware; use crate::{ errors::ServerError, @@ -135,15 +136,9 @@ impl Starting { // Create health check if enabled (only for StreamableHttp transport) let health_check = match (&self.config.transport, self.config.health_check.enabled) { - ( - Transport::StreamableHttp { - auth: _, - address: _, - port: _, - stateful_mode: _, - }, - true, - ) => Some(HealthCheck::new(self.config.health_check.clone())), + (Transport::StreamableHttp { .. }, true) => { + Some(HealthCheck::new(self.config.health_check.clone())) + } _ => None, // No health check for SSE, Stdio, or when disabled }; @@ -214,6 +209,7 @@ impl Starting { address, port, stateful_mode, + host_validation, } => { info!(port = ?port, address = ?address, "Starting MCP server in Streamable HTTP mode"); let running = running.clone(); @@ -234,7 +230,15 @@ impl Starting { // include trace context as header into the response .layer(OtelInResponseLayer) // start OpenTelemetry trace on incoming request - .layer(axum::middleware::from_fn(otel_context_middleware)); + .layer(axum::middleware::from_fn(otel_context_middleware)) + // Host header validation to prevent DNS rebinding attacks + .layer(axum::middleware::from_fn_with_state( + HostValidationState { + config: Arc::new(host_validation), + server_port: port, + }, + validate_host, + )); // Add health check endpoint if configured if let Some(health_check) = health_check.filter(|h| h.config().enabled) { @@ -289,6 +293,7 @@ mod tests { use url::Url; use crate::health::HealthCheckConfig; + use crate::host_validation::HostValidationConfig; use super::*; @@ -301,6 +306,7 @@ mod tests { address: "127.0.0.1".parse().unwrap(), port: 7799, stateful_mode: false, + host_validation: HostValidationConfig::default(), }, endpoint: Url::parse("http://localhost:4000").expect("valid url"), mutation_mode: MutationMode::All, From 077181bba91a64c89f179284ea107f7f3bd1bc1c Mon Sep 17 00:00:00 2001 From: Dale Seo Date: Wed, 28 Jan 2026 13:51:01 -0500 Subject: [PATCH 2/3] chore:changeset --- .changeset/host_validation.md | 16 ++++++++++++++++ 1 file changed, 16 insertions(+) create mode 100644 .changeset/host_validation.md diff --git a/.changeset/host_validation.md b/.changeset/host_validation.md new file mode 100644 index 0000000000..7ccdb081fe --- /dev/null +++ b/.changeset/host_validation.md @@ -0,0 +1,16 @@ +--- +default: minor +--- + +Add Host header validation to prevent DNS rebinding attacks. Requests with invalid Host headers are now rejected with 403 Forbidden. Enabled by default for StreamableHttp transport. + +```yaml +transport: + type: streamable_http + host_validation: + enabled: true # default + allowed_hosts: + - mcp.dev.example.com + - mcp.staging.example.com + - mcp.example.com +``` From 5a356563a3eb4984b796f099db92a3530fd7ac4f Mon Sep 17 00:00:00 2001 From: Dale Seo Date: Wed, 28 Jan 2026 13:55:45 -0500 Subject: [PATCH 3/3] docs: add config options for host validation --- docs/source/config-file.mdx | 38 +++++++++++++++++++++++++++++++------ 1 file changed, 32 insertions(+), 6 deletions(-) diff --git a/docs/source/config-file.mdx b/docs/source/config-file.mdx index d836271af7..1f170fb11c 100644 --- a/docs/source/config-file.mdx +++ b/docs/source/config-file.mdx @@ -225,13 +225,14 @@ The available fields depend on the value of the nested `type` key. The default t ##### Transport Type Specific options -Some transport types support further configuration. For `streamable_http`, you can set the `address`, `port`, and `stateful_mode`. +Some transport types support further configuration. For `streamable_http`, you can set `address`, `port`, `stateful_mode`, and `host_validation`. -| Option | Type | Default | Description | -| :-------------- | :------- | :---------- | :------------------------------------------------------------- | -| `address` | `IpAddr` | `127.0.0.1` | The IP address to bind to | -| `port` | `u16` | `8000` | The port to bind to | -| `stateful_mode` | `bool` | `true` | Flag to enable or disable stateful mode and session management | +| Option | Type | Default | Description | +| :---------------- | :--------------- | :---------- | :------------------------------------------------------------- | +| `address` | `IpAddr` | `127.0.0.1` | The IP address to bind to | +| `port` | `u16` | `8000` | The port to bind to | +| `stateful_mode` | `bool` | `true` | Flag to enable or disable stateful mode and session management | +| `host_validation` | `HostValidation` | | Host header validation configuration | @@ -239,6 +240,31 @@ For Apollo MCP Server `≤v1.0.0`, the default `port` value is `5000`. In `v1.1. +### Host Validation + +These fields are under the `host_validation` key within the `transport` configuration. Host validation prevents DNS rebinding attacks by rejecting requests with unexpected `Host` headers. + +| Option | Type | Default | Description | +| :-------------- | :------------- | :------ | :--------------------------------------------------------------- | +| `enabled` | `bool` | `true` | Enable Host header validation to prevent DNS rebinding attacks | +| `allowed_hosts` | `List` | `[]` | Additional allowed hostnames beyond localhost (can include port) | + + + +Host validation is only available when using the `streamable_http` transport. Localhost addresses (`localhost`, `127.0.0.1`, `::1`, `0.0.0.0`) are always allowed when validation is enabled. + + + +For production deployments behind a reverse proxy, add your server's hostname: + +```yaml title="mcp.yaml" +transport: + type: streamable_http + host_validation: + allowed_hosts: + - mcp.example.com +``` + ### Auth These fields are under the top-level `transport` key, nested under the `auth` key. Learn more about [authorization and authentication](/apollo-mcp-server/auth).