use crate::error::Result;
use crate::server::http_middleware::{
adapters::{from_axum_with_limit, into_axum},
ServerHttpContext, ServerHttpMiddlewareChain, ServerHttpResponse,
};
use crate::server::tower_layers::{AllowedOrigins, DnsRebindingLayer, SecurityHeadersLayer};
use crate::server::Server;
use crate::shared::http_constants::{
APPLICATION_JSON, MCP_METHOD, MCP_NAME, MCP_PROTOCOL_VERSION, MCP_SESSION_ID, TEXT_EVENT_STREAM,
};
use crate::shared::TransportMessage;
use crate::types::{ClientRequest, Request};
use async_trait::async_trait;
use axum::{
body::Body,
extract::State,
http::{header, HeaderMap, HeaderValue, StatusCode},
response::{sse::Event, IntoResponse, Response, Sse},
routing::{delete, get, post},
Json, Router,
};
use futures_util::StreamExt;
use parking_lot::RwLock;
use serde_json::json;
use std::collections::HashMap;
use std::convert::Infallible;
use std::net::SocketAddr;
use std::sync::Arc;
#[cfg(not(target_arch = "wasm32"))]
use tokio::sync::mpsc;
use tokio_stream::wrappers::UnboundedReceiverStream;
use uuid::Uuid;
#[rustfmt::skip]
#[cfg_attr(feature = "v1-compat", path = "streamable_http_server/v1_session.rs")]
#[cfg_attr(not(feature = "v1-compat"), path = "streamable_http_server/v1_session_off.rs")]
pub(crate) mod v1;
#[async_trait]
pub trait EventStore: Send + Sync {
async fn store_event(
&self,
stream_id: &str,
event_id: &str,
message: &TransportMessage,
) -> Result<()>;
async fn replay_events_after(
&self,
last_event_id: &str,
) -> Result<Vec<(String, TransportMessage)>>;
async fn get_stream_for_event(&self, event_id: &str) -> Result<Option<String>>;
}
type EventList = Vec<(String, TransportMessage)>;
type EventsMap = HashMap<String, EventList>;
#[cfg_attr(
feature = "v1-compat",
doc = r"
The `v1-compat` half pins the CONFIG WIRING, which is gated — this example does
not compile on `--no-default-features --features full-v2`, and that is the
severance being asserted rather than a bug:
```rust
use pmcp::server::streamable_http_server::{InMemoryEventStore, StreamableHttpServerConfig};
use std::sync::Arc;
let store = Arc::new(InMemoryEventStore::default());
let config = StreamableHttpServerConfig {
event_store: Some(Arc::clone(&store)),
..Default::default()
};
assert!(config.event_store.is_some());
```
"
)]
#[derive(Debug, Default)]
pub struct InMemoryEventStore {
events: Arc<RwLock<EventsMap>>,
event_to_stream: Arc<RwLock<HashMap<String, String>>>,
event_order: Arc<RwLock<Vec<String>>>,
}
#[async_trait]
impl EventStore for InMemoryEventStore {
async fn store_event(
&self,
stream_id: &str,
event_id: &str,
message: &TransportMessage,
) -> Result<()> {
let mut events = self.events.write();
let stream_events = events.entry(stream_id.to_string()).or_default();
stream_events.push((event_id.to_string(), message.clone()));
self.event_to_stream
.write()
.insert(event_id.to_string(), stream_id.to_string());
self.event_order.write().push(event_id.to_string());
Ok(())
}
async fn replay_events_after(
&self,
last_event_id: &str,
) -> Result<Vec<(String, TransportMessage)>> {
let event_order = self.event_order.read();
let mut result = Vec::new();
let start_pos = event_order
.iter()
.position(|id| id == last_event_id)
.map_or(0, |pos| pos + 1);
let events = self.events.read();
let event_to_stream = self.event_to_stream.read();
for i in start_pos..event_order.len() {
let event_id = &event_order[i];
if let Some(stream_id) = event_to_stream.get(event_id) {
if let Some(stream_events) = events.get(stream_id) {
for (eid, msg) in stream_events {
if eid == event_id {
result.push((eid.clone(), msg.clone()));
break;
}
}
}
}
}
Ok(result)
}
async fn get_stream_for_event(&self, event_id: &str) -> Result<Option<String>> {
Ok(self.event_to_stream.read().get(event_id).cloned())
}
}
#[cfg(feature = "v1-compat")]
#[cfg_attr(docsrs, doc(cfg(feature = "v1-compat")))]
type SessionCallback = Box<dyn Fn(&str) + Send + Sync>;
#[cfg_attr(
feature = "v1-compat",
doc = r#"
Stateful MCP 2025-11-25 configuration — `v1-compat` builds only:
```rust
use pmcp::server::streamable_http_server::StreamableHttpServerConfig;
let config = StreamableHttpServerConfig {
session_id_generator: Some(Box::new(|| {
format!("session-{}", uuid::Uuid::new_v4())
})),
on_session_initialized: Some(Box::new(|session_id| {
println!("Session started: {}", session_id);
})),
on_session_closed: Some(Box::new(|session_id| {
println!("Session ended: {}", session_id);
})),
..Default::default()
};
assert!(config.session_id_generator.is_some());
assert!(config.on_session_closed.is_some());
```
"#
)]
pub struct StreamableHttpServerConfig {
#[cfg(feature = "v1-compat")]
#[cfg_attr(docsrs, doc(cfg(feature = "v1-compat")))]
pub session_id_generator: Option<Box<dyn Fn() -> String + Send + Sync>>,
pub enable_json_response: bool,
#[cfg(feature = "v1-compat")]
#[cfg_attr(docsrs, doc(cfg(feature = "v1-compat")))]
pub event_store: Option<Arc<InMemoryEventStore>>,
#[cfg(feature = "v1-compat")]
#[cfg_attr(docsrs, doc(cfg(feature = "v1-compat")))]
pub on_session_initialized: Option<SessionCallback>,
#[cfg(feature = "v1-compat")]
#[cfg_attr(docsrs, doc(cfg(feature = "v1-compat")))]
pub on_session_closed: Option<SessionCallback>,
pub http_middleware: Option<Arc<ServerHttpMiddlewareChain>>,
pub allowed_origins: Option<AllowedOrigins>,
pub max_request_bytes: usize,
}
impl std::fmt::Debug for StreamableHttpServerConfig {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let mut out = f.debug_struct("StreamableHttpServerConfig");
#[cfg(feature = "v1-compat")]
out.field("session_id_generator", &self.session_id_generator.is_some());
out.field("enable_json_response", &self.enable_json_response);
#[cfg(feature = "v1-compat")]
out.field("event_store", &self.event_store.is_some());
#[cfg(feature = "v1-compat")]
out.field(
"on_session_initialized",
&self.on_session_initialized.is_some(),
);
#[cfg(feature = "v1-compat")]
out.field("on_session_closed", &self.on_session_closed.is_some());
out.field("http_middleware", &self.http_middleware.is_some());
out.field("allowed_origins", &self.allowed_origins);
out.field("max_request_bytes", &self.max_request_bytes);
out.finish()
}
}
impl Default for StreamableHttpServerConfig {
fn default() -> Self {
Self {
#[cfg(feature = "v1-compat")]
session_id_generator: Some(Box::new(|| Uuid::new_v4().to_string())),
enable_json_response: false,
#[cfg(feature = "v1-compat")]
event_store: Some(Arc::new(InMemoryEventStore::default())),
#[cfg(feature = "v1-compat")]
on_session_initialized: None,
#[cfg(feature = "v1-compat")]
on_session_closed: None,
http_middleware: None,
allowed_origins: None,
max_request_bytes: crate::server::limits::DEFAULT_MAX_REQUEST_BYTES,
}
}
}
impl StreamableHttpServerConfig {
pub fn stateless() -> Self {
Self {
#[cfg(feature = "v1-compat")]
session_id_generator: None,
enable_json_response: true,
#[cfg(feature = "v1-compat")]
event_store: None,
#[cfg(feature = "v1-compat")]
on_session_initialized: None,
#[cfg(feature = "v1-compat")]
on_session_closed: None,
http_middleware: None,
allowed_origins: Some(AllowedOrigins::any()),
max_request_bytes: crate::server::limits::DEFAULT_MAX_REQUEST_BYTES,
}
}
}
#[derive(Clone)]
pub(crate) struct ServerState {
server: Arc<tokio::sync::Mutex<Server>>,
config: Arc<StreamableHttpServerConfig>,
allowed_origins: AllowedOrigins,
v1: v1::V1State,
}
pub(crate) fn build_mcp_router(state: ServerState) -> Router<()> {
Router::new()
.route("/", post(handle_post_request))
.route("/", get(handle_get_sse))
.route("/", delete(handle_delete_session))
.with_state(state)
}
pub(crate) fn make_server_state(
server: Arc<tokio::sync::Mutex<Server>>,
config: StreamableHttpServerConfig,
) -> ServerState {
let allowed_origins = config
.allowed_origins
.clone()
.unwrap_or_else(AllowedOrigins::localhost);
let v1 = v1::V1State::new(&config);
ServerState {
server,
config: Arc::new(config),
allowed_origins,
v1,
}
}
pub struct StreamableHttpServer {
addr: SocketAddr,
state: ServerState,
}
impl std::fmt::Debug for StreamableHttpServer {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("StreamableHttpServer")
.field("addr", &self.addr)
.field("state", &"ServerState { ... }")
.finish()
}
}
fn create_error_response(status: StatusCode, code: i32, message: &str) -> Response {
let error_body = json!({
"jsonrpc": "2.0",
"error": {
"code": code,
"message": message
},
"id": null
});
(status, Json(error_body)).into_response()
}
use crate::types::mrtr::MAX_HEADER_SENTINEL_LEN as MAX_V2_HEADER_SENTINEL_LEN;
use crate::types::mrtr::MAX_HEADER_VALUE_LEN as MAX_V2_HEADER_VALUE_LEN;
fn envelope_for_live_request(
payload: crate::types::jsonrpc::ResponsePayload<serde_json::Value, crate::types::JSONRPCError>,
live_id: crate::types::RequestId,
) -> crate::types::JSONRPCResponse {
match payload {
crate::types::jsonrpc::ResponsePayload::Result(result) => {
crate::types::JSONRPCResponse::success(live_id, result)
},
crate::types::jsonrpc::ResponsePayload::Error(error) => {
crate::types::JSONRPCResponse::error(live_id, error)
},
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum HeaderProtocolVersion {
Absent,
Malformed,
V2,
Other,
}
enum V2Classification {
Legacy,
Enforce,
Reject(i32, &'static str),
}
fn v2_status_for_code(code: i32) -> StatusCode {
use crate::types::protocol::error_codes as ec;
match code {
ec::METHOD_NOT_FOUND => StatusCode::NOT_FOUND,
ec::HEADER_MISMATCH
| ec::MISSING_REQUIRED_CLIENT_CAPABILITY
| ec::UNSUPPORTED_PROTOCOL_VERSION
| ec::PARSE_ERROR
| ec::INVALID_REQUEST
| ec::INVALID_PARAMS => StatusCode::BAD_REQUEST,
_ => StatusCode::OK,
}
}
fn status_for_error(
era: Option<crate::types::protocol::Era>,
code: i32,
v1_status: StatusCode,
) -> StatusCode {
if matches!(era, Some(crate::types::protocol::Era::V2)) {
v2_status_for_code(code)
} else {
v1_status
}
}
fn raw_request_id(body: &[u8]) -> serde_json::Value {
serde_json::from_slice::<serde_json::Value>(body)
.ok()
.and_then(|v| v.get("id").cloned())
.unwrap_or(serde_json::Value::Null)
}
fn create_error_response_with_id(
status: StatusCode,
id: serde_json::Value,
code: i32,
message: &str,
data: Option<serde_json::Value>,
) -> Response {
let mut error = serde_json::Map::new();
error.insert("code".to_string(), json!(code));
error.insert("message".to_string(), json!(message));
if let Some(data) = data {
error.insert("data".to_string(), data);
}
let mut body = serde_json::Map::new();
body.insert("jsonrpc".to_string(), json!("2.0"));
body.insert("error".to_string(), serde_json::Value::Object(error));
body.insert("id".to_string(), id);
(status, Json(serde_json::Value::Object(body))).into_response()
}
async fn map_unparsed_body_for_v2(
state: &ServerState,
raw_body: &[u8],
v1_response: Response,
) -> Response {
use crate::types::protocol::error_codes::METHOD_NOT_FOUND;
let Ok(envelope) = serde_json::from_slice::<serde_json::Value>(raw_body) else {
return v1_response;
};
let Some(method) = envelope.get("method").and_then(serde_json::Value::as_str) else {
return v1_response;
};
if envelope.get("id").is_none() {
return v1_response;
}
let raw_meta = params_meta_of(Some(&envelope));
let resolved = {
let server = state.server.lock().await;
server.resolve_raw_meta_protocol_context(raw_meta.as_ref())
};
let Ok(Some(context)) = resolved else {
return v1_response;
};
if context.era != crate::types::protocol::Era::V2 {
return v1_response;
}
create_error_response_with_id(
v2_status_for_code(METHOD_NOT_FOUND),
envelope
.get("id")
.cloned()
.unwrap_or(serde_json::Value::Null),
METHOD_NOT_FOUND,
&format!("Method not found: {method}"),
None,
)
}
fn v2_dispatch_response_status(
era: Option<crate::types::protocol::Era>,
response: &crate::types::JSONRPCResponse,
) -> Option<StatusCode> {
if !matches!(era, Some(crate::types::protocol::Era::V2)) {
return None;
}
let crate::types::jsonrpc::ResponsePayload::Error(ref error) = response.payload else {
return None;
};
Some(v2_status_for_code(error.code))
}
fn v2_gate_reject_response(
raw_body: &[u8],
era: Option<crate::types::protocol::Era>,
code: i32,
message: &str,
data: Option<serde_json::Value>,
) -> Response {
let status = status_for_error(era, code, StatusCode::BAD_REQUEST);
create_error_response_with_id(status, raw_request_id(raw_body), code, message, data)
}
enum V2GateOutcome {
Passthrough,
EnforceOk { method: String, name: String },
Reject {
code: i32,
message: String,
data: Option<serde_json::Value>,
},
}
fn decode_version_header(headers: &HeaderMap) -> HeaderProtocolVersion {
let Some(raw) = headers.get(MCP_PROTOCOL_VERSION) else {
return HeaderProtocolVersion::Absent;
};
if raw.as_bytes().len() > MAX_V2_HEADER_VALUE_LEN {
return HeaderProtocolVersion::Malformed;
}
match raw.to_str() {
Err(_) => HeaderProtocolVersion::Malformed,
Ok(s) if s == crate::types::protocol::PROTOCOL_VERSION_2026_07_28 => {
HeaderProtocolVersion::V2
},
Ok(_) => HeaderProtocolVersion::Other,
}
}
fn bounded_header_str(headers: &HeaderMap, name: &str) -> Option<String> {
let raw = headers.get(name)?;
if raw.as_bytes().len() > MAX_V2_HEADER_SENTINEL_LEN {
return None;
}
raw.to_str().ok().map(str::to_string)
}
fn classify_era_cell(header: HeaderProtocolVersion, meta_is_v2: bool) -> V2Classification {
let header_is_v2 = matches!(header, HeaderProtocolVersion::V2);
match (header_is_v2, meta_is_v2) {
(true, true) => V2Classification::Enforce,
(false, false) => V2Classification::Legacy,
(true, false) => V2Classification::Reject(
crate::types::protocol::error_codes::HEADER_MISMATCH,
"MCP-Protocol-Version header claims v2 but _meta protocolVersion disagrees",
),
(false, true) => V2Classification::Reject(
crate::types::protocol::error_codes::HEADER_MISMATCH,
"_meta claims v2 but MCP-Protocol-Version header is absent or not 2026-07-28",
),
}
}
const ERR_MISSING_V2_HEADERS: &str =
"v2 requests must carry Mcp-Method and MCP-Protocol-Version headers";
const ERR_MISSING_MCP_NAME: &str =
"Mcp-Name header is required: this method carries a routing name";
fn require_v2_headers(headers: &HeaderMap) -> std::result::Result<(String, String), &'static str> {
if headers.get(MCP_PROTOCOL_VERSION).is_none() {
return Err(ERR_MISSING_V2_HEADERS);
}
let Some(method) = bounded_header_str(headers, MCP_METHOD) else {
return Err(ERR_MISSING_V2_HEADERS);
};
if !is_name_bearing_method(&method) {
return Ok((method, String::new()));
}
match bounded_header_str(headers, MCP_NAME) {
Some(name) => Ok((method, name)),
None => Err(ERR_MISSING_MCP_NAME),
}
}
fn cross_check_method(
mcp_method: &str,
body_method: Option<&str>,
) -> std::result::Result<(), &'static str> {
match body_method {
Some(bm) if bm == mcp_method => Ok(()),
_ => Err("Mcp-Method header does not match the JSON-RPC body method"),
}
}
fn is_name_bearing_method(method: &str) -> bool {
crate::types::mrtr::name_bearing_key(method).is_some()
}
fn cross_check_name(
mcp_name: &str,
method: &str,
body_name: Option<&str>,
) -> std::result::Result<(), &'static str> {
if !is_name_bearing_method(method) {
return Ok(());
}
let Some(decoded) = crate::types::mrtr::decode_header_value(mcp_name) else {
return Err("Mcp-Name header is a malformed =?base64?...?= sentinel value");
};
match body_name {
Some(bn) if bn == decoded => Ok(()),
_ => Err("Mcp-Name header does not match the request's logical name"),
}
}
fn classify_v2_request(
headers: &HeaderMap,
meta_is_v2: bool,
body_method: Option<&str>,
body_name: Option<&str>,
) -> V2GateOutcome {
use crate::types::protocol::error_codes::HEADER_MISMATCH;
let reject = |msg: &str| V2GateOutcome::Reject {
code: HEADER_MISMATCH,
message: msg.to_string(),
data: None,
};
let header = decode_version_header(headers);
match classify_era_cell(header, meta_is_v2) {
V2Classification::Legacy => V2GateOutcome::Passthrough,
V2Classification::Reject(code, msg) => V2GateOutcome::Reject {
code,
message: msg.to_string(),
data: None,
},
V2Classification::Enforce => {
let (method, name) = match require_v2_headers(headers) {
Ok(pair) => pair,
Err(msg) => return reject(msg),
};
if let Err(msg) = cross_check_method(&method, body_method) {
return reject(msg);
}
if let Err(msg) = cross_check_name(&name, &method, body_name) {
return reject(msg);
}
V2GateOutcome::EnforceOk { method, name }
},
}
}
#[cfg(test)]
fn extract_body_method_and_name(body: &[u8]) -> (Option<String>, Option<String>) {
method_and_name_of(raw_body_json(body).as_ref())
}
fn method_and_name_of(value: Option<&serde_json::Value>) -> (Option<String>, Option<String>) {
let Some(value) = value else {
return (None, None);
};
match crate::types::mrtr::frame_routing_pair(value) {
Some((method, name)) => (Some(method.to_string()), name),
None => (None, None),
}
}
fn apply_v2_outbound_headers(headers: &mut HeaderMap, method: &str, name: &str) {
if let Ok(v) = HeaderValue::from_str(method) {
headers.insert(MCP_METHOD, v);
}
if let Ok(v) = HeaderValue::from_str(name) {
headers.insert(MCP_NAME, v);
}
if let Ok(v) = HeaderValue::from_str(crate::types::protocol::PROTOCOL_VERSION_2026_07_28) {
headers.insert(MCP_PROTOCOL_VERSION, v);
}
}
fn negotiation_error_to_gate_reject(
error: &crate::types::protocol::context::ProtocolNegotiationError,
accept_list: &[crate::types::ProtocolVersion],
) -> V2GateOutcome {
use crate::types::protocol::context::ProtocolNegotiationError;
use crate::types::protocol::error_codes::UNSUPPORTED_PROTOCOL_VERSION;
match error {
ProtocolNegotiationError::UnsupportedVersion(requested) => {
let supported: Vec<&str> = accept_list.iter().map(|v| v.as_str()).collect();
V2GateOutcome::Reject {
code: UNSUPPORTED_PROTOCOL_VERSION,
message: format!("Unsupported protocol version: {requested}"),
data: Some(json!({ "requested": requested, "supported": supported })),
}
},
ProtocolNegotiationError::MalformedMeta(_) => {
let (code, message) = crate::server::core::negotiation_error_to_rejection(error);
V2GateOutcome::Reject {
code,
message,
data: None,
}
},
}
}
#[cfg(test)]
fn raw_params_meta(body: &[u8]) -> Option<serde_json::Value> {
params_meta_of(raw_body_json(body).as_ref())
}
fn raw_body_json(body: &[u8]) -> Option<serde_json::Value> {
serde_json::from_slice::<serde_json::Value>(body).ok()
}
fn params_meta_of(value: Option<&serde_json::Value>) -> Option<serde_json::Value> {
let params = value?.get("params")?;
params
.get(crate::types::mrtr::META_KEY)
.or_else(|| params.get("meta"))
.filter(|meta| !meta.is_null())
.cloned()
}
fn params_of(value: Option<&serde_json::Value>) -> &serde_json::Value {
const NO_PARAMS: &serde_json::Value = &serde_json::Value::Null;
value.and_then(|v| v.get("params")).unwrap_or(NO_PARAMS)
}
fn attach_v2_mrtr_params(
context: Option<crate::types::protocol::ProtocolContext>,
outcome: V2GateOutcome,
body_json: Option<&serde_json::Value>,
method: Option<&str>,
) -> (
Option<crate::types::protocol::ProtocolContext>,
V2GateOutcome,
) {
if !matches!(outcome, V2GateOutcome::EnforceOk { .. }) {
return (context, outcome);
}
if !method.is_some_and(crate::types::mrtr::mrtr_eligible) {
return (context, outcome);
}
let Some(ctx) = context else {
return (None, outcome);
};
match crate::types::mrtr::extract_mrtr_params(params_of(body_json)) {
Ok(mrtr) => (Some(ctx.with_mrtr_params(mrtr)), outcome),
Err(reason) => {
tracing::warn!(
target: "mcp.http",
reason = ?reason,
"rejecting a v2 request whose MRTR params are present but unusable"
);
let message = reason.to_string();
(
Some(ctx),
V2GateOutcome::Reject {
code: crate::types::protocol::error_codes::INVALID_PARAMS,
message,
data: None,
},
)
},
}
}
async fn run_v2_header_gate(
state: &ServerState,
headers: &HeaderMap,
raw_body: &[u8],
body_method_override: Option<&str>,
) -> (
Option<crate::types::protocol::ProtocolContext>,
V2GateOutcome,
) {
{
let server = state.server.lock().await;
if !crate::types::protocol::context::is_v2_opted_in(server.supported_protocol_versions()) {
return (None, V2GateOutcome::Passthrough);
}
}
let body_json = raw_body_json(raw_body);
let raw_meta = params_meta_of(body_json.as_ref());
let resolved = {
let server = state.server.lock().await;
server
.resolve_raw_meta_protocol_context(raw_meta.as_ref())
.map_err(|err| {
negotiation_error_to_gate_reject(&err, server.supported_protocol_versions())
})
};
let context = match resolved {
Ok(ctx) => ctx,
Err(reject) => return (None, reject),
};
let Some(ref pc) = context else {
return (context.clone(), V2GateOutcome::Passthrough);
};
let meta_is_v2 = pc.era == crate::types::protocol::Era::V2;
let (extracted_method, body_name) = method_and_name_of(body_json.as_ref());
let body_method = body_method_override.or(extracted_method.as_deref());
let outcome = classify_v2_request(headers, meta_is_v2, body_method, body_name.as_deref());
attach_v2_mrtr_params(context, outcome, body_json.as_ref(), body_method)
}
enum HttpIngress {
Public(TransportMessage),
Discover { id: crate::types::RequestId },
SubscriptionsListen {
id: crate::types::RequestId,
params: Option<serde_json::Value>,
},
TasksUpdate {
id: crate::types::RequestId,
params: serde_json::Value,
},
}
impl HttpIngress {
fn is_initialize(&self) -> bool {
match self {
Self::Public(msg) => is_initialize_request(msg),
Self::Discover { .. } | Self::SubscriptionsListen { .. } | Self::TasksUpdate { .. } => {
false
},
}
}
}
fn classify_http_ingress(body: &[u8]) -> Option<HttpIngress> {
let req: crate::types::JSONRPCRequest<serde_json::Value> = serde_json::from_slice(body).ok()?;
if req.method == crate::types::subscriptions::SUBSCRIPTIONS_LISTEN_METHOD {
return Some(HttpIngress::SubscriptionsListen {
id: req.id,
params: req.params,
});
}
if req.method != crate::types::protocol::SERVER_DISCOVER_METHOD
&& req.method != crate::types::protocol::TASKS_UPDATE_METHOD
{
return None;
}
let (id, ingress) = crate::shared::protocol_helpers::parse_request_or_internal(req).ok()?;
match ingress {
crate::shared::protocol_helpers::IngressRequest::Internal(internal) => match internal {
crate::types::protocol::InternalClientRequest::ServerDiscover(_) => {
Some(HttpIngress::Discover { id })
},
crate::types::protocol::InternalClientRequest::TasksUpdate { params } => {
Some(HttpIngress::TasksUpdate { id, params })
},
},
crate::shared::protocol_helpers::IngressRequest::Public(_) => None,
}
}
impl StreamableHttpServer {
pub fn new(addr: SocketAddr, server: Arc<tokio::sync::Mutex<Server>>) -> Self {
Self::with_config(addr, server, StreamableHttpServerConfig::default())
}
pub fn with_config(
addr: SocketAddr,
server: Arc<tokio::sync::Mutex<Server>>,
config: StreamableHttpServerConfig,
) -> Self {
let state = make_server_state(server, config);
Self { addr, state }
}
pub async fn start(self) -> Result<(SocketAddr, tokio::task::JoinHandle<()>)> {
let allowed = self.state.allowed_origins.clone();
let cors = crate::server::tower_layers::build_mcp_cors_layer(&allowed);
let app = build_mcp_router(self.state)
.layer(SecurityHeadersLayer::default())
.layer(DnsRebindingLayer::new(allowed))
.layer(cors);
let listener = tokio::net::TcpListener::bind(self.addr).await?;
let local_addr = listener.local_addr()?;
let server_task = tokio::spawn(async move {
axum::serve(listener, app).await.unwrap();
});
Ok((local_addr, server_task))
}
}
pub(crate) fn method_not_allowed_for_verb(verb: &str) -> Response {
let mut response = create_error_response(
StatusCode::METHOD_NOT_ALLOWED,
crate::types::protocol::error_codes::METHOD_NOT_FOUND,
&format!("HTTP {verb} is not supported on the MCP endpoint for protocol 2026-07-28"),
);
response
.headers_mut()
.insert(header::ALLOW, HeaderValue::from_static("POST, OPTIONS"));
response
}
fn v2_method_not_allowed(headers: &HeaderMap, verb: &str, v2_opted_in: bool) -> Option<Response> {
if !v2_opted_in || !matches!(decode_version_header(headers), HeaderProtocolVersion::V2) {
return None;
}
Some(method_not_allowed_for_verb(verb))
}
async fn v2_verb_rejection(
state: &ServerState,
headers: &HeaderMap,
verb: &str,
) -> Option<Response> {
if !matches!(decode_version_header(headers), HeaderProtocolVersion::V2) {
return None;
}
let opted_in = {
let server = state.server.lock().await;
crate::types::protocol::context::is_v2_opted_in(server.supported_protocol_versions())
};
v2_method_not_allowed(headers, verb, opted_in)
}
fn validate_content_type_json(headers: &HeaderMap) -> std::result::Result<(), Response> {
let Some(content_type) = headers.get(header::CONTENT_TYPE) else {
return Err(create_error_response(
StatusCode::UNSUPPORTED_MEDIA_TYPE,
crate::types::protocol::error_codes::PARSE_ERROR,
"Content-Type header is required",
));
};
let ct = content_type.to_str().unwrap_or("");
if !ct.contains(APPLICATION_JSON) {
return Err(create_error_response(
StatusCode::UNSUPPORTED_MEDIA_TYPE,
crate::types::protocol::error_codes::PARSE_ERROR,
"Content-Type must be application/json",
));
}
Ok(())
}
fn validate_accept_post(headers: &HeaderMap) -> std::result::Result<(), Response> {
let Some(accept) = headers.get(header::ACCEPT) else {
return Err(create_error_response(
StatusCode::NOT_ACCEPTABLE,
crate::types::protocol::error_codes::PARSE_ERROR,
"Accept header is required",
));
};
let accept_str = accept.to_str().unwrap_or("");
if !accept_str.contains(APPLICATION_JSON) && !accept_str.contains(TEXT_EVENT_STREAM) {
return Err(create_error_response(
StatusCode::NOT_ACCEPTABLE,
crate::types::protocol::error_codes::PARSE_ERROR,
"Accept header must include application/json or text/event-stream",
));
}
Ok(())
}
fn validate_accept_sse(headers: &HeaderMap) -> std::result::Result<(), Response> {
let Some(accept) = headers.get(header::ACCEPT) else {
return Err(create_error_response(
StatusCode::NOT_ACCEPTABLE,
crate::types::protocol::error_codes::PARSE_ERROR,
"Accept header is required for SSE",
));
};
let accept_str = accept.to_str().unwrap_or("");
if !accept_str.contains(TEXT_EVENT_STREAM) {
return Err(create_error_response(
StatusCode::NOT_ACCEPTABLE,
crate::types::protocol::error_codes::PARSE_ERROR,
"Accept header must be text/event-stream for SSE",
));
}
Ok(())
}
fn validate_headers(headers: &HeaderMap, method: &str) -> std::result::Result<(), Response> {
match method {
"POST" => {
validate_content_type_json(headers)?;
validate_accept_post(headers)?;
},
"GET" => validate_accept_sse(headers)?,
_ => {},
}
Ok(())
}
fn serialize_response_as_json_value(
response: &TransportMessage,
) -> std::result::Result<serde_json::Value, Response> {
let json_bytes = crate::shared::StdioTransport::serialize_message(response).map_err(|e| {
create_error_response(
StatusCode::INTERNAL_SERVER_ERROR,
crate::types::protocol::error_codes::INTERNAL_ERROR,
&format!("Failed to serialize response: {}", e),
)
})?;
tracing::debug!(
target: "mcp.http",
response = %String::from_utf8_lossy(&json_bytes),
"HTTP response serialized bytes"
);
let json_value: serde_json::Value = serde_json::from_slice(&json_bytes).map_err(|e| {
create_error_response(
StatusCode::INTERNAL_SERVER_ERROR,
crate::types::protocol::error_codes::INTERNAL_ERROR,
&format!("Failed to parse JSON response: {}", e),
)
})?;
Ok(json_value)
}
fn build_json_response(response: &TransportMessage, trace_source: &'static str) -> Response {
let json_value = match serialize_response_as_json_value(response) {
Ok(v) => v,
Err(error_response) => return error_response,
};
tracing::debug!(
target: "mcp.http",
source = trace_source,
response = %serde_json::to_string(&json_value).unwrap_or_default(),
"HTTP response (JSON mode)"
);
(StatusCode::OK, Json(json_value)).into_response()
}
fn build_sse_response_from_single_message(response: TransportMessage) -> Response {
let (tx, rx) = mpsc::unbounded_channel();
tx.send(response).unwrap();
let stream = UnboundedReceiverStream::new(rx);
let sse = Sse::new(stream.map(|msg| {
let event_id = Uuid::new_v4().to_string();
let json_bytes =
crate::shared::StdioTransport::serialize_message(&msg).unwrap_or_else(|e| {
tracing::error!(target: "mcp.sse", error = %e, "Failed to serialize SSE message");
Vec::new()
});
let json_str = String::from_utf8(json_bytes).unwrap_or_else(|_| "{}".to_string());
Ok::<_, Infallible>(
Event::default()
.id(event_id)
.event("message")
.data(json_str),
)
}));
sse.into_response()
}
fn build_response(
state: &ServerState,
response: TransportMessage,
session_id: Option<&String>,
sessions_on: bool,
) -> Response {
if state.config.enable_json_response {
return build_json_response(&response, "JSON mode");
}
let Some(sid) = session_id.filter(|_| sessions_on) else {
return build_json_response(&response, "SSE no-session fallback");
};
let Some(undelivered) = v1::route_to_session_stream(&state.v1, sid, response) else {
return StatusCode::ACCEPTED.into_response();
};
build_sse_response_from_single_message(undelivered)
}
fn validate_protocol_version_supported(
protocol_version: Option<&String>,
) -> std::result::Result<(), Response> {
let Some(version) = protocol_version else {
return Ok(());
};
if crate::SUPPORTED_PROTOCOL_VERSIONS.contains(&version.as_str()) {
return Ok(());
}
Err(create_error_response(
StatusCode::BAD_REQUEST,
crate::types::protocol::error_codes::INVALID_REQUEST,
&format!("Unsupported protocol version: {}", version),
))
}
fn validate_protocol_version(
state: &ServerState,
era: Option<crate::types::protocol::Era>,
session_id: Option<&String>,
protocol_version: Option<&String>,
) -> std::result::Result<(), Response> {
validate_protocol_version_supported(protocol_version)?;
v1::validate_protocol_version_matches_session(state, era, session_id, protocol_version)
}
async fn handle_post_request(
State(state): State<ServerState>,
request: axum::extract::Request<Body>,
) -> impl IntoResponse {
if state.config.http_middleware.is_none() {
return Box::pin(handle_post_fast_path(state, request)).await;
}
Box::pin(handle_post_with_middleware(state, request)).await
}
async fn extract_and_validate_auth(
state: &ServerState,
headers: &HeaderMap,
) -> std::result::Result<Option<crate::server::auth::AuthContext>, Response> {
let server = state.server.lock().await;
if let Some(auth_provider) = server.get_auth_provider() {
let auth_header = headers
.get(http::header::AUTHORIZATION)
.and_then(|v| v.to_str().ok());
match auth_provider.validate_request(auth_header).await {
Ok(ctx) => Ok(ctx),
Err(e) => {
Err(create_error_response(
StatusCode::UNAUTHORIZED,
crate::types::protocol::error_codes::AUTHENTICATION_REQUIRED,
&format!("Authentication failed: {}", e),
))
},
}
} else {
Ok(extract_auth_from_proxy_headers(headers))
}
}
fn extract_auth_from_proxy_headers(
headers: &HeaderMap,
) -> Option<crate::server::auth::AuthContext> {
let user_id = headers
.get("x-pmcp-user-id")
.and_then(|v| v.to_str().ok())
.map(|s| s.to_string())?;
let email = headers
.get("x-pmcp-user-email")
.and_then(|v| v.to_str().ok())
.map(|s| s.to_string());
let name = headers
.get("x-pmcp-user-name")
.and_then(|v| v.to_str().ok())
.map(|s| s.to_string());
let groups = headers
.get("x-pmcp-user-groups")
.and_then(|v| v.to_str().ok())
.map(|s| s.to_string());
let tenant_id = headers
.get("x-pmcp-tenant-id")
.and_then(|v| v.to_str().ok())
.map(|s| s.to_string());
let mut claims = std::collections::HashMap::new();
if let Some(ref email) = email {
claims.insert(
"email".to_string(),
serde_json::Value::String(email.clone()),
);
}
if let Some(ref name) = name {
claims.insert("name".to_string(), serde_json::Value::String(name.clone()));
}
if let Some(ref groups) = groups {
let groups_array: Vec<serde_json::Value> = groups
.split(',')
.map(|g| serde_json::Value::String(g.trim().to_string()))
.filter(|v| v.as_str() != Some(""))
.collect();
claims.insert("groups".to_string(), serde_json::Value::Array(groups_array));
}
if let Some(ref tenant_id) = tenant_id {
claims.insert(
"tenant_id".to_string(),
serde_json::Value::String(tenant_id.clone()),
);
}
for (name, value) in headers {
let Some(suffix) = name.as_str().strip_prefix("x-pmcp-claim-custom-") else {
continue;
};
let Ok(val_str) = value.to_str() else {
continue;
};
if suffix.is_empty() || val_str.is_empty() {
continue;
}
let snake: String = suffix
.chars()
.map(|c| if c == '-' { '_' } else { c })
.collect();
claims.insert(
format!("custom:{}", snake),
serde_json::Value::String(val_str.to_string()),
);
}
tracing::debug!(
user_id = %user_id,
email = ?email,
"Extracted auth context from proxy headers"
);
Some(crate::server::auth::AuthContext {
subject: user_id,
scopes: vec![],
claims,
token: None,
client_id: None,
expires_at: None,
authenticated: true,
})
}
fn extract_session_and_protocol_headers(headers: &HeaderMap) -> (Option<String>, Option<String>) {
let session_id = v1::incoming_session_header(headers);
let protocol_version = headers
.get(MCP_PROTOCOL_VERSION)
.and_then(|v| v.to_str().ok())
.map(|s| s.to_string());
(session_id, protocol_version)
}
type V2GateResolved = (
Option<crate::types::protocol::ProtocolContext>,
Option<(String, String)>,
);
async fn resolve_v2_gate(
state: &ServerState,
headers: &HeaderMap,
raw_body: &[u8],
ingress: &HttpIngress,
) -> std::result::Result<V2GateResolved, Response> {
match ingress {
HttpIngress::Public(TransportMessage::Request { .. })
| HttpIngress::Discover { .. }
| HttpIngress::SubscriptionsListen { .. }
| HttpIngress::TasksUpdate { .. } => {
let method_override = matches!(ingress, HttpIngress::Discover { .. })
.then_some(crate::types::protocol::SERVER_DISCOVER_METHOD);
let (ctx, gate) = run_v2_header_gate(state, headers, raw_body, method_override).await;
match gate {
V2GateOutcome::Reject {
code,
message,
data,
} => {
let era = ctx.as_ref().map(|pc| pc.era);
Err(v2_gate_reject_response(raw_body, era, code, &message, data))
},
V2GateOutcome::Passthrough => Ok((ctx, None)),
V2GateOutcome::EnforceOk { method, name } => Ok((ctx, Some((method, name)))),
}
},
HttpIngress::Public(_) => Ok((None, None)),
}
}
fn is_initialize_request(message: &TransportMessage) -> bool {
matches!(
message,
TransportMessage::Request { request: Request::Client(boxed), .. }
if matches!(**boxed, ClientRequest::Initialize(_))
)
}
fn extract_negotiated_version(response: &TransportMessage) -> Option<String> {
if let TransportMessage::Response(ref json_resp) = response {
if let crate::types::jsonrpc::ResponsePayload::Result(ref value) = json_resp.payload {
if let Ok(init_result) =
serde_json::from_value::<crate::types::InitializeResult>(value.clone())
{
return Some(init_result.protocol_version.0);
}
}
}
None
}
fn compute_outbound_protocol_version(
state: &ServerState,
response_session_id: Option<&String>,
is_init_request: bool,
negotiated_version: Option<&str>,
) -> String {
if is_init_request {
return negotiated_version.map_or_else(
|| crate::DEFAULT_PROTOCOL_VERSION.to_string(),
std::string::ToString::to_string,
);
}
if let Some(sid) = response_session_id {
if let Some(negotiated_version) = v1::session_protocol_version(&state.v1, sid.as_str()) {
return negotiated_version;
}
}
crate::DEFAULT_PROTOCOL_VERSION.to_string()
}
async fn report_middleware_error(
http_middleware: &ServerHttpMiddlewareChain,
context: &ServerHttpContext,
error_kind: &str,
) {
let err = crate::Error::protocol_msg(error_kind);
let _ = http_middleware.handle_error(&err, context).await;
}
async fn run_request_middleware(
http_middleware: &ServerHttpMiddlewareChain,
server_request: &mut crate::server::http_middleware::ServerHttpRequest,
context: &ServerHttpContext,
) -> std::result::Result<(), Response> {
if let Err(e) = http_middleware
.process_request(server_request, context)
.await
{
let _ = http_middleware.handle_error(&e, context).await;
return Err(create_error_response(
StatusCode::INTERNAL_SERVER_ERROR,
crate::types::protocol::error_codes::INTERNAL_ERROR,
&format!("Middleware rejected request: {}", e),
));
}
Ok(())
}
async fn parse_transport_message_with_middleware(
body: &[u8],
http_middleware: &ServerHttpMiddlewareChain,
context: &ServerHttpContext,
) -> std::result::Result<HttpIngress, Response> {
if let Some(ingress) = classify_http_ingress(body) {
return Ok(ingress);
}
match crate::shared::StdioTransport::parse_message(body) {
Ok(msg) => Ok(HttpIngress::Public(msg)),
Err(e) => {
let mut error_response = ServerHttpResponse::new(
StatusCode::BAD_REQUEST,
HeaderMap::new(),
format!("{{\"error\":\"Invalid JSON: {}\"}}", e).into_bytes(),
);
let _ = http_middleware
.process_response(&mut error_response, context)
.await;
Err(into_axum(error_response))
},
}
}
async fn extract_auth_with_middleware(
state: &ServerState,
server_request: &crate::server::http_middleware::ServerHttpRequest,
http_middleware: &ServerHttpMiddlewareChain,
context: &ServerHttpContext,
) -> std::result::Result<Option<crate::server::auth::AuthContext>, Response> {
let server = state.server.lock().await;
let Some(auth_provider) = server.get_auth_provider() else {
return Ok(None);
};
let auth_header = server_request.get_header("authorization");
match auth_provider.validate_request(auth_header).await {
Ok(ctx) => Ok(ctx),
Err(e) => {
let auth_error = crate::Error::authentication(format!("Authentication failed: {}", e));
let _ = http_middleware.handle_error(&auth_error, context).await;
Err(create_error_response(
StatusCode::UNAUTHORIZED,
crate::types::protocol::error_codes::AUTHENTICATION_REQUIRED,
&format!("Authentication failed: {}", e),
))
},
}
}
async fn build_success_response_with_middleware(
response_msg: &TransportMessage,
response_session_id: Option<&String>,
version_to_send: &str,
sessions_on: bool,
http_middleware: &ServerHttpMiddlewareChain,
context: &ServerHttpContext,
) -> Response {
let response_body = match serde_json::to_vec(response_msg) {
Ok(b) => b,
Err(e) => {
let serialization_error =
crate::Error::internal(format!("Failed to serialize response: {}", e));
let _ = http_middleware
.handle_error(&serialization_error, context)
.await;
return create_error_response(
StatusCode::INTERNAL_SERVER_ERROR,
crate::types::protocol::error_codes::INTERNAL_ERROR,
&format!("Failed to serialize response: {}", e),
);
},
};
let mut response_headers = HeaderMap::new();
response_headers.insert(header::CONTENT_TYPE, APPLICATION_JSON.parse().unwrap());
v1::apply_session_header(&mut response_headers, response_session_id, sessions_on);
response_headers.insert(MCP_PROTOCOL_VERSION, version_to_send.parse().unwrap());
let mut server_response =
ServerHttpResponse::new(StatusCode::OK, response_headers, response_body);
if let Err(e) = http_middleware
.process_response(&mut server_response, context)
.await
{
tracing::warn!("Response middleware processing failed: {}", e);
}
into_axum(server_response)
}
async fn read_body_with_limit(
body: Body,
max_bytes: usize,
) -> std::result::Result<String, Response> {
let body_bytes = axum::body::to_bytes(body, max_bytes).await.map_err(|e| {
create_error_response(
StatusCode::PAYLOAD_TOO_LARGE,
crate::types::protocol::error_codes::INVALID_REQUEST,
&format!("Request body exceeds limit: {}", e),
)
})?;
Ok(String::from_utf8_lossy(&body_bytes).to_string())
}
fn parse_transport_message_fast(body: &[u8]) -> std::result::Result<HttpIngress, Response> {
if let Some(ingress) = classify_http_ingress(body) {
return Ok(ingress);
}
crate::shared::StdioTransport::parse_message(body)
.map(HttpIngress::Public)
.map_err(|e| {
create_error_response(
StatusCode::BAD_REQUEST,
crate::types::protocol::error_codes::PARSE_ERROR,
&format!("Invalid JSON: {}", e),
)
})
}
struct FastPathDispatch {
is_init_request: bool,
response_session_id: Option<String>,
protocol_context: Option<crate::types::protocol::ProtocolContext>,
v2_outbound: Option<(String, String)>,
sessions_on: bool,
}
async fn handle_fast_path_request(
state: &ServerState,
id: crate::types::RequestId,
request: Request,
auth_context: Option<crate::server::auth::AuthContext>,
dispatch: FastPathDispatch,
session_id: Option<&String>,
) -> Response {
let FastPathDispatch {
is_init_request,
response_session_id,
protocol_context,
v2_outbound,
sessions_on,
} = dispatch;
let era = protocol_context.as_ref().map(|pc| pc.era);
let live_id = id.clone();
let json_response =
dispatch_request_or_retire(state, id, request, auth_context, protocol_context).await;
tracing::debug!(
target: "mcp.http",
response = %serde_json::to_string(&json_response).unwrap_or_default(),
"StreamableHttpServer response"
);
let v2_status = v2_dispatch_response_status(era, &json_response);
let response_msg =
TransportMessage::Response(envelope_for_live_request(json_response.payload, live_id));
let negotiated_version = if is_init_request {
let version = extract_negotiated_version(&response_msg);
v1::update_session_after_init(state, response_session_id.as_ref(), version.clone());
version
} else {
None
};
v1::store_response_event(state, era, response_session_id.as_ref(), &response_msg).await;
let mut response = build_response(state, response_msg, session_id, sessions_on);
v1::apply_session_header(
response.headers_mut(),
response_session_id.as_ref(),
sessions_on,
);
let version_to_send = compute_outbound_protocol_version(
state,
response_session_id.as_ref(),
is_init_request,
negotiated_version.as_deref(),
);
response
.headers_mut()
.insert(MCP_PROTOCOL_VERSION, version_to_send.parse().unwrap());
if let Some((method, name)) = &v2_outbound {
apply_v2_outbound_headers(response.headers_mut(), method, name);
}
if let Some(status) = v2_status {
*response.status_mut() = status;
}
response
}
struct InternalResponseShape<'a> {
response_session_id: Option<&'a String>,
v2_outbound: Option<(String, String)>,
sessions_on: bool,
}
async fn assemble_discover_response_fast(
state: &ServerState,
id: crate::types::RequestId,
protocol_context: Option<&crate::types::protocol::ProtocolContext>,
shape: InternalResponseShape<'_>,
session_id: Option<&String>,
) -> Response {
let InternalResponseShape {
response_session_id,
v2_outbound,
sessions_on,
} = shape;
let live_id = id.clone();
let json_response = {
let server = state.server.lock().await;
server.handle_discover(id, protocol_context)
};
let era = protocol_context.map(|pc| pc.era);
let v2_status = v2_dispatch_response_status(era, &json_response);
let response_msg =
TransportMessage::Response(envelope_for_live_request(json_response.payload, live_id));
v1::store_response_event(state, era, response_session_id, &response_msg).await;
let mut response = build_response(state, response_msg, session_id, sessions_on);
v1::apply_session_header(response.headers_mut(), response_session_id, sessions_on);
let version_to_send =
compute_outbound_protocol_version(state, response_session_id, false, None);
response
.headers_mut()
.insert(MCP_PROTOCOL_VERSION, version_to_send.parse().unwrap());
if let Some((method, name)) = &v2_outbound {
apply_v2_outbound_headers(response.headers_mut(), method, name);
}
if let Some(status) = v2_status {
*response.status_mut() = status;
}
response
}
struct TasksUpdateCall<'a> {
id: crate::types::RequestId,
params: serde_json::Value,
protocol_context: Option<&'a crate::types::protocol::ProtocolContext>,
auth_context: Option<&'a crate::server::auth::AuthContext>,
}
async fn tasks_update_json_response(
state: &ServerState,
call: &TasksUpdateCall<'_>,
) -> crate::types::JSONRPCResponse {
let server = state.server.lock().await;
server
.handle_tasks_update(
call.id.clone(),
&call.params,
call.auth_context,
call.protocol_context,
)
.await
}
async fn assemble_tasks_update_fast(
state: &ServerState,
call: TasksUpdateCall<'_>,
shape: InternalResponseShape<'_>,
session_id: Option<&String>,
) -> Response {
let InternalResponseShape {
response_session_id,
v2_outbound,
sessions_on,
} = shape;
let live_id = call.id.clone();
let protocol_context = call.protocol_context;
let json_response = tasks_update_json_response(state, &call).await;
let era = protocol_context.map(|pc| pc.era);
let v2_status = v2_dispatch_response_status(era, &json_response);
let response_msg =
TransportMessage::Response(envelope_for_live_request(json_response.payload, live_id));
v1::store_response_event(state, era, response_session_id, &response_msg).await;
let mut response = build_response(state, response_msg, session_id, sessions_on);
v1::apply_session_header(response.headers_mut(), response_session_id, sessions_on);
let version_to_send =
compute_outbound_protocol_version(state, response_session_id, false, None);
response
.headers_mut()
.insert(MCP_PROTOCOL_VERSION, version_to_send.parse().unwrap());
if let Some((method, name)) = &v2_outbound {
apply_v2_outbound_headers(response.headers_mut(), method, name);
}
if let Some(status) = v2_status {
*response.status_mut() = status;
}
response
}
async fn assemble_tasks_update_with_middleware(
state: &ServerState,
call: TasksUpdateCall<'_>,
shape: InternalResponseShape<'_>,
http_middleware: &ServerHttpMiddlewareChain,
http_context: &ServerHttpContext,
) -> Response {
let InternalResponseShape {
response_session_id,
v2_outbound,
sessions_on,
} = shape;
let live_id = call.id.clone();
let protocol_context = call.protocol_context;
let json_response = tasks_update_json_response(state, &call).await;
let era = protocol_context.map(|pc| pc.era);
let v2_status = v2_dispatch_response_status(era, &json_response);
let response_msg =
TransportMessage::Response(envelope_for_live_request(json_response.payload, live_id));
v1::store_response_event(state, era, response_session_id, &response_msg).await;
let version_to_send =
compute_outbound_protocol_version(state, response_session_id, false, None);
let mut response = build_success_response_with_middleware(
&response_msg,
response_session_id,
&version_to_send,
sessions_on,
http_middleware,
http_context,
)
.await;
if let Some((method, name)) = &v2_outbound {
apply_v2_outbound_headers(response.headers_mut(), method, name);
}
if let Some(status) = v2_status {
*response.status_mut() = status;
}
response
}
const X_ACCEL_BUFFERING: &str = "x-accel-buffering";
const LISTEN_KEEP_ALIVE_INTERVAL: std::time::Duration = std::time::Duration::from_secs(15);
fn v2_retired_method_of(request: &Request) -> Option<&'static str> {
let Request::Client(client) = request else {
return None;
};
match **client {
ClientRequest::Subscribe(_) => Some("resources/subscribe"),
ClientRequest::Unsubscribe(_) => Some("resources/unsubscribe"),
_ => None,
}
}
async fn dispatch_request_or_retire(
state: &ServerState,
id: crate::types::RequestId,
request: Request,
auth_context: Option<crate::server::auth::AuthContext>,
protocol_context: Option<crate::types::protocol::ProtocolContext>,
) -> crate::types::JSONRPCResponse {
if matches!(
protocol_context.as_ref().map(|pc| pc.era),
Some(crate::types::protocol::Era::V2)
) {
if let Some(method) = v2_retired_method_of(&request) {
return crate::types::JSONRPCResponse::error(
id,
crate::types::jsonrpc::JSONRPCError {
code: crate::types::protocol::error_codes::METHOD_NOT_FOUND,
message: format!(
"Method not found: {method} (retired in MCP 2026-07-28; use {})",
crate::types::subscriptions::SUBSCRIPTIONS_LISTEN_METHOD
),
data: None,
},
);
}
}
let server = state.server.lock().await;
server
.handle_request_with_context(id, request, auth_context, protocol_context)
.await
}
struct ListenServerView {
capabilities: crate::types::ServerCapabilities,
info: crate::types::Implementation,
registry: Arc<crate::server::subscriptions::ListenRegistry>,
has_auth_provider: bool,
}
async fn listen_server_view(state: &ServerState) -> ListenServerView {
let server = state.server.lock().await;
ListenServerView {
capabilities: server.capabilities().clone(),
info: server.info().clone(),
registry: Arc::clone(server.listen_registry()),
has_auth_provider: server.get_auth_provider().is_some(),
}
}
fn listen_rejection_response(
era: Option<crate::types::protocol::Era>,
id: crate::types::RequestId,
code: i32,
message: String,
) -> Response {
let response = envelope_for_live_request(
crate::types::jsonrpc::ResponsePayload::Error(crate::types::jsonrpc::JSONRPCError {
code,
message,
data: None,
}),
id,
);
let status = v2_dispatch_response_status(era, &response);
let mut http = build_json_response(
&TransportMessage::Response(response),
"subscriptions/listen gate",
);
if let Some(status) = status {
*http.status_mut() = status;
}
http
}
fn listen_ack_frame(
agreed: &crate::types::subscriptions::SubscriptionFilter,
subscription_id: &crate::types::RequestId,
) -> String {
let params = crate::types::subscriptions::SubscriptionAcknowledgedParams::new(
agreed.clone(),
subscription_id,
);
json!({
"jsonrpc": "2.0",
"method": crate::types::subscriptions::ACKNOWLEDGED_METHOD,
"params": params,
})
.to_string()
}
fn listen_terminal_result_frame(
subscription_id: &crate::types::RequestId,
protocol_context: Option<&crate::types::protocol::ProtocolContext>,
server_info: &crate::types::Implementation,
) -> String {
let result = crate::types::subscriptions::SubscriptionsListenResult::new(subscription_id);
let mut response = envelope_for_live_request(
crate::types::jsonrpc::ResponsePayload::Result(
serde_json::to_value(result).unwrap_or_else(|_| json!({})),
),
subscription_id.clone(),
);
crate::server::core::inject_v2_result_envelope(
&mut response,
protocol_context,
server_info,
crate::server::core::ResponseDisposition::Complete,
crate::server::core::ReservedFieldOwner::None,
crate::types::caching::Cacheable::No,
);
serde_json::to_string(&response).unwrap_or_else(|_| "{}".to_string())
}
fn listen_sse_event(frame: crate::server::subscriptions::ListenFrame) -> Event {
match frame {
crate::server::subscriptions::ListenFrame::Message(payload) => {
Event::default().event("message").data(payload)
},
crate::server::subscriptions::ListenFrame::Comment(text) => Event::default().comment(text),
}
}
fn attach_listen_response_headers(response: &mut Response, v2_outbound: Option<&(String, String)>) {
let headers = response.headers_mut();
headers.insert(
header::CACHE_CONTROL,
HeaderValue::from_static("no-cache, no-transform"),
);
headers.insert(X_ACCEL_BUFFERING, HeaderValue::from_static("no"));
if let Some((method, name)) = v2_outbound {
apply_v2_outbound_headers(headers, method, name);
}
}
fn resolve_agreed_filter(
params: Option<serde_json::Value>,
view: &ListenServerView,
) -> std::result::Result<crate::types::subscriptions::SubscriptionFilter, (i32, String)> {
use crate::types::protocol::error_codes::INVALID_PARAMS;
use crate::types::subscriptions::SubscriptionsListenParams;
let Some(value) = params else {
return Err((
INVALID_PARAMS,
"Invalid subscriptions/listen params: `notifications` is required".to_string(),
));
};
let parsed = serde_json::from_value::<SubscriptionsListenParams>(value).map_err(|e| {
(
INVALID_PARAMS,
format!("Invalid subscriptions/listen params: {e}"),
)
})?;
Ok(parsed
.notifications
.intersect_with_capabilities(&view.capabilities))
}
fn resolve_listen_principal(
auth_context: Option<&crate::server::auth::AuthContext>,
has_auth_provider: bool,
) -> Option<String> {
match (auth_context, has_auth_provider) {
(Some(context), _) => Some(context.subject.clone()),
(None, true) => None,
(None, false) => Some(crate::server::subscriptions::anonymous_principal()),
}
}
async fn assemble_subscriptions_listen(
state: &ServerState,
id: crate::types::RequestId,
params: Option<serde_json::Value>,
protocol_context: Option<&crate::types::protocol::ProtocolContext>,
v2_outbound: Option<(String, String)>,
auth_context: Option<&crate::server::auth::AuthContext>,
) -> Response {
use crate::server::subscriptions::{ListenFrame, ListenKey, LISTEN_CHANNEL_CAPACITY};
use crate::types::protocol::error_codes::{AUTHENTICATION_REQUIRED, METHOD_NOT_FOUND};
use crate::types::subscriptions::SUBSCRIPTIONS_LISTEN_METHOD;
let era = protocol_context.map(|pc| pc.era);
if !matches!(era, Some(crate::types::protocol::Era::V2)) {
return listen_rejection_response(
era,
id,
METHOD_NOT_FOUND,
format!("Method not found: {SUBSCRIPTIONS_LISTEN_METHOD}"),
);
}
debug_assert!(
!v1::resumability_active(state, era),
"a v2 request already has resumability off (plan 08); the listen stream asserts that \
rather than re-deriving it"
);
let view = listen_server_view(state).await;
if !crate::types::subscriptions::advertises_subscriptions(&view.capabilities) {
return listen_rejection_response(
era,
id,
METHOD_NOT_FOUND,
format!(
"Method not found: {SUBSCRIPTIONS_LISTEN_METHOD} (this server advertises no \
subscription-delivered capability)"
),
);
}
let Some(principal) = resolve_listen_principal(auth_context, view.has_auth_provider) else {
return listen_rejection_response(
era,
id,
AUTHENTICATION_REQUIRED,
format!(
"{SUBSCRIPTIONS_LISTEN_METHOD} requires an authenticated caller on this server"
),
);
};
let agreed = match resolve_agreed_filter(params, &view) {
Ok(filter) => filter,
Err((code, message)) => return listen_rejection_response(era, id, code, message),
};
let (sender, receiver) = mpsc::channel(LISTEN_CHANNEL_CAPACITY + 1);
if sender
.try_send(ListenFrame::Message(listen_ack_frame(&agreed, &id)))
.is_err()
{
return listen_rejection_response(
era,
id,
crate::types::protocol::error_codes::INTERNAL_ERROR,
"failed to queue the subscription acknowledgement".to_string(),
);
}
let terminal = listen_terminal_result_frame(&id, protocol_context, &view.info);
let registry = view.registry;
let key = ListenKey {
principal,
request_id: id.clone(),
};
let guard = match registry.register(key, agreed, sender, terminal) {
Ok(guard) => guard,
Err(rejection) => {
return listen_rejection_response(
era,
id,
rejection.code(),
rejection.message().to_string(),
);
},
};
let frames =
futures_util::stream::unfold((receiver, guard), |(mut receiver, guard)| async move {
receiver
.recv()
.await
.map(|frame| (frame, (receiver, guard)))
});
let events = frames.map(|frame| Ok::<_, Infallible>(listen_sse_event(frame)));
let mut response = Sse::new(events)
.keep_alive(axum::response::sse::KeepAlive::new().interval(LISTEN_KEEP_ALIVE_INTERVAL))
.into_response();
attach_listen_response_headers(&mut response, v2_outbound.as_ref());
response
}
fn guard_legacy_version_fast(
state: &ServerState,
era: Option<crate::types::protocol::Era>,
is_init_request: bool,
is_v2_request: bool,
session_id: Option<&String>,
protocol_version: Option<&String>,
) -> std::result::Result<(), Response> {
if !is_init_request && !is_v2_request {
validate_protocol_version(state, era, session_id, protocol_version)?;
}
Ok(())
}
struct FastIngress {
headers: HeaderMap,
body: String,
ingress: HttpIngress,
session_id: Option<String>,
protocol_version: Option<String>,
is_init_request: bool,
}
async fn read_and_classify_fast(
state: &ServerState,
request: axum::extract::Request<Body>,
) -> std::result::Result<FastIngress, Response> {
let (parts, body) = request.into_parts();
let headers = parts.headers;
let body = read_body_with_limit(body, state.config.max_request_bytes).await?;
validate_headers(&headers, "POST")?;
let ingress = match parse_transport_message_fast(body.as_bytes()) {
Ok(i) => i,
Err(response) => {
return Err(map_unparsed_body_for_v2(state, body.as_bytes(), response).await)
},
};
let (session_id, protocol_version) = extract_session_and_protocol_headers(&headers);
let is_init_request = ingress.is_initialize();
Ok(FastIngress {
headers,
body,
ingress,
session_id,
protocol_version,
is_init_request,
})
}
async fn handle_post_fast_path(
state: ServerState,
request: axum::extract::Request<Body>,
) -> Response {
match handle_post_fast_path_inner(state, request).await {
Ok(response) | Err(response) => response,
}
}
async fn handle_post_fast_path_inner(
state: ServerState,
request: axum::extract::Request<Body>,
) -> std::result::Result<Response, Response> {
let FastIngress {
headers,
body,
ingress,
session_id,
protocol_version,
is_init_request,
} = read_and_classify_fast(&state, request).await?;
let (protocol_context, v2_outbound) =
resolve_v2_gate(&state, &headers, body.as_bytes(), &ingress).await?;
let is_v2_request = v2_outbound.is_some();
let era = protocol_context.as_ref().map(|pc| pc.era);
let sessions_on = v1::sessions_active(&state, era);
let response_session_id = v1::resolve_session_for_request(
&state,
era,
is_init_request,
session_id.clone(),
protocol_version.clone(),
)?;
guard_legacy_version_fast(
&state,
era,
is_init_request,
is_v2_request,
session_id.as_ref(),
protocol_version.as_ref(),
)?;
let auth_context = extract_and_validate_auth(&state, &headers).await?;
Ok(dispatch_message_fast(
&state,
ingress,
FastPathDispatch {
is_init_request,
response_session_id,
protocol_context,
v2_outbound,
sessions_on,
},
auth_context,
session_id.as_ref(),
)
.await)
}
fn build_middleware_context(
server_request: &crate::server::http_middleware::ServerHttpRequest,
) -> ServerHttpContext {
let session_id = server_request
.get_header(MCP_SESSION_ID)
.map(str::to_string);
let request_id = server_request
.get_header("x-request-id")
.map_or_else(|| Uuid::new_v4().to_string(), str::to_string);
ServerHttpContext {
request_id,
start_time: std::time::Instant::now(),
session_id,
}
}
async fn convert_axum_to_middleware_request(
request: axum::extract::Request<Body>,
max_request_bytes: usize,
) -> std::result::Result<crate::server::http_middleware::ServerHttpRequest, Response> {
let (parts, body) = request.into_parts();
from_axum_with_limit(parts, body, max_request_bytes)
.await
.map_err(|e| {
create_error_response(
StatusCode::PAYLOAD_TOO_LARGE,
crate::types::protocol::error_codes::INVALID_REQUEST,
&format!("Request body exceeds limit: {}", e),
)
})
}
async fn resolve_session_with_error_hook(
state: &ServerState,
era: Option<crate::types::protocol::Era>,
is_init_request: bool,
session_id: Option<String>,
protocol_version: Option<String>,
http_middleware: &ServerHttpMiddlewareChain,
http_context: &ServerHttpContext,
) -> std::result::Result<Option<String>, Response> {
match v1::resolve_session_for_request(state, era, is_init_request, session_id, protocol_version)
{
Ok(sid) => Ok(sid),
Err(error_response) => {
let kind = if is_init_request {
"Session initialization failed"
} else {
"Session validation failed"
};
report_middleware_error(http_middleware, http_context, kind).await;
Err(error_response)
},
}
}
async fn validate_protocol_version_with_error_hook(
state: &ServerState,
era: Option<crate::types::protocol::Era>,
is_init_request: bool,
session_id: Option<&String>,
protocol_version: Option<&String>,
http_middleware: &ServerHttpMiddlewareChain,
http_context: &ServerHttpContext,
) -> std::result::Result<(), Response> {
if is_init_request {
return Ok(());
}
if let Err(error_response) = validate_protocol_version(state, era, session_id, protocol_version)
{
report_middleware_error(
http_middleware,
http_context,
"Protocol version validation failed",
)
.await;
return Err(error_response);
}
Ok(())
}
async fn resolve_v2_gate_with_error_hook(
state: &ServerState,
headers: &HeaderMap,
raw_body: &[u8],
ingress: &HttpIngress,
http_middleware: &ServerHttpMiddlewareChain,
http_context: &ServerHttpContext,
) -> std::result::Result<V2GateResolved, Response> {
match resolve_v2_gate(state, headers, raw_body, ingress).await {
Ok(resolved) => Ok(resolved),
Err(error_response) => {
report_middleware_error(http_middleware, http_context, "v2 header gate rejected").await;
Err(error_response)
},
}
}
struct MiddlewareDispatch {
is_init_request: bool,
response_session_id: Option<String>,
protocol_context: Option<crate::types::protocol::ProtocolContext>,
v2_outbound: Option<(String, String)>,
sessions_on: bool,
}
async fn assemble_discover_response_with_middleware(
state: &ServerState,
id: crate::types::RequestId,
protocol_context: Option<&crate::types::protocol::ProtocolContext>,
shape: InternalResponseShape<'_>,
http_middleware: &ServerHttpMiddlewareChain,
http_context: &ServerHttpContext,
) -> Response {
let InternalResponseShape {
response_session_id,
v2_outbound,
sessions_on,
} = shape;
let live_id = id.clone();
let json_response = {
let server = state.server.lock().await;
server.handle_discover(id, protocol_context)
};
let era = protocol_context.map(|pc| pc.era);
let v2_status = v2_dispatch_response_status(era, &json_response);
let response_msg =
TransportMessage::Response(envelope_for_live_request(json_response.payload, live_id));
v1::store_response_event(state, era, response_session_id, &response_msg).await;
let version_to_send =
compute_outbound_protocol_version(state, response_session_id, false, None);
let mut response = build_success_response_with_middleware(
&response_msg,
response_session_id,
&version_to_send,
sessions_on,
http_middleware,
http_context,
)
.await;
if let Some((method, name)) = &v2_outbound {
apply_v2_outbound_headers(response.headers_mut(), method, name);
}
if let Some(status) = v2_status {
*response.status_mut() = status;
}
response
}
async fn dispatch_message_fast(
state: &ServerState,
ingress: HttpIngress,
dispatch: FastPathDispatch,
auth_context: Option<crate::server::auth::AuthContext>,
session_id: Option<&String>,
) -> Response {
match ingress {
HttpIngress::Public(TransportMessage::Request { id, request }) => {
Box::pin(handle_fast_path_request(
state,
id,
request,
auth_context,
dispatch,
session_id,
))
.await
},
HttpIngress::Discover { id, .. } => {
let FastPathDispatch {
response_session_id,
protocol_context,
v2_outbound,
sessions_on,
..
} = dispatch;
assemble_discover_response_fast(
state,
id,
protocol_context.as_ref(),
InternalResponseShape {
response_session_id: response_session_id.as_ref(),
v2_outbound,
sessions_on,
},
session_id,
)
.await
},
HttpIngress::SubscriptionsListen { id, params } => {
let FastPathDispatch {
protocol_context,
v2_outbound,
..
} = dispatch;
Box::pin(assemble_subscriptions_listen(
state,
id,
params,
protocol_context.as_ref(),
v2_outbound,
auth_context.as_ref(),
))
.await
},
HttpIngress::TasksUpdate { id, params } => {
let FastPathDispatch {
response_session_id,
protocol_context,
v2_outbound,
sessions_on,
..
} = dispatch;
Box::pin(assemble_tasks_update_fast(
state,
TasksUpdateCall {
id,
params,
protocol_context: protocol_context.as_ref(),
auth_context: auth_context.as_ref(),
},
InternalResponseShape {
response_session_id: response_session_id.as_ref(),
v2_outbound,
sessions_on,
},
session_id,
))
.await
},
HttpIngress::Public(
TransportMessage::Notification { .. } | TransportMessage::Response(_),
) => StatusCode::ACCEPTED.into_response(),
}
}
async fn dispatch_message_with_middleware(
state: &ServerState,
ingress: HttpIngress,
dispatch: MiddlewareDispatch,
auth_context: Option<crate::server::auth::AuthContext>,
http_middleware: &ServerHttpMiddlewareChain,
http_context: &ServerHttpContext,
) -> Response {
let MiddlewareDispatch {
is_init_request,
response_session_id,
protocol_context,
v2_outbound,
sessions_on,
} = dispatch;
match ingress {
HttpIngress::Discover { id, .. } => {
assemble_discover_response_with_middleware(
state,
id,
protocol_context.as_ref(),
InternalResponseShape {
response_session_id: response_session_id.as_ref(),
v2_outbound,
sessions_on,
},
http_middleware,
http_context,
)
.await
},
HttpIngress::SubscriptionsListen { id, params } => {
assemble_subscriptions_listen(
state,
id,
params,
protocol_context.as_ref(),
v2_outbound,
auth_context.as_ref(),
)
.await
},
HttpIngress::TasksUpdate { id, params } => {
assemble_tasks_update_with_middleware(
state,
TasksUpdateCall {
id,
params,
protocol_context: protocol_context.as_ref(),
auth_context: auth_context.as_ref(),
},
InternalResponseShape {
response_session_id: response_session_id.as_ref(),
v2_outbound,
sessions_on,
},
http_middleware,
http_context,
)
.await
},
HttpIngress::Public(TransportMessage::Request { id, request }) => {
let era = protocol_context.as_ref().map(|pc| pc.era);
let live_id = id.clone();
let json_response =
dispatch_request_or_retire(state, id, request, auth_context, protocol_context)
.await;
let v2_status = v2_dispatch_response_status(era, &json_response);
let response_msg = TransportMessage::Response(envelope_for_live_request(
json_response.payload,
live_id,
));
let negotiated_version = if is_init_request {
let version = extract_negotiated_version(&response_msg);
v1::update_session_after_init(state, response_session_id.as_ref(), version.clone());
version
} else {
None
};
v1::store_response_event(state, era, response_session_id.as_ref(), &response_msg).await;
let version_to_send = compute_outbound_protocol_version(
state,
response_session_id.as_ref(),
is_init_request,
negotiated_version.as_deref(),
);
let mut response = build_success_response_with_middleware(
&response_msg,
response_session_id.as_ref(),
&version_to_send,
sessions_on,
http_middleware,
http_context,
)
.await;
if let Some((method, name)) = &v2_outbound {
apply_v2_outbound_headers(response.headers_mut(), method, name);
}
if let Some(status) = v2_status {
*response.status_mut() = status;
}
response
},
HttpIngress::Public(
TransportMessage::Notification { .. } | TransportMessage::Response(_),
) => StatusCode::ACCEPTED.into_response(),
}
}
#[allow(clippy::too_many_arguments)]
async fn guard_legacy_version_with_middleware(
state: &ServerState,
era: Option<crate::types::protocol::Era>,
is_init_request: bool,
is_v2_request: bool,
session_id: Option<&String>,
protocol_version: Option<&String>,
http_middleware: &ServerHttpMiddlewareChain,
http_context: &ServerHttpContext,
) -> std::result::Result<(), Response> {
if !is_v2_request {
validate_protocol_version_with_error_hook(
state,
era,
is_init_request,
session_id,
protocol_version,
http_middleware,
http_context,
)
.await?;
}
Ok(())
}
struct MwIngress {
server_request: crate::server::http_middleware::ServerHttpRequest,
http_context: ServerHttpContext,
ingress: HttpIngress,
session_id: Option<String>,
protocol_version: Option<String>,
is_init_request: bool,
}
async fn read_and_classify_with_middleware(
state: &ServerState,
request: axum::extract::Request<Body>,
http_middleware: &ServerHttpMiddlewareChain,
) -> std::result::Result<MwIngress, Response> {
let mut server_request =
convert_axum_to_middleware_request(request, state.config.max_request_bytes).await?;
let http_context = build_middleware_context(&server_request);
run_request_middleware(http_middleware, &mut server_request, &http_context).await?;
if let Err(error_response) = validate_headers(&server_request.headers, "POST") {
report_middleware_error(http_middleware, &http_context, "Header validation failed").await;
return Err(error_response);
}
let ingress = match parse_transport_message_with_middleware(
&server_request.body,
http_middleware,
&http_context,
)
.await
{
Ok(i) => i,
Err(response) => {
return Err(map_unparsed_body_for_v2(state, &server_request.body, response).await)
},
};
let (session_id, protocol_version) =
extract_session_and_protocol_headers(&server_request.headers);
let is_init_request = ingress.is_initialize();
Ok(MwIngress {
server_request,
http_context,
ingress,
session_id,
protocol_version,
is_init_request,
})
}
async fn handle_post_with_middleware(
state: ServerState,
request: axum::extract::Request<Body>,
) -> Response {
match handle_post_with_middleware_inner(state, request).await {
Ok(response) | Err(response) => response,
}
}
async fn handle_post_with_middleware_inner(
state: ServerState,
request: axum::extract::Request<Body>,
) -> std::result::Result<Response, Response> {
let http_middleware = state
.config
.http_middleware
.as_ref()
.expect("Middleware chain must exist");
let MwIngress {
server_request,
http_context,
ingress,
session_id,
protocol_version,
is_init_request,
} = read_and_classify_with_middleware(&state, request, http_middleware).await?;
let (protocol_context, v2_outbound) = resolve_v2_gate_with_error_hook(
&state,
&server_request.headers,
&server_request.body,
&ingress,
http_middleware,
&http_context,
)
.await?;
let is_v2_request = v2_outbound.is_some();
let era = protocol_context.as_ref().map(|pc| pc.era);
let sessions_on = v1::sessions_active(&state, era);
let response_session_id = resolve_session_with_error_hook(
&state,
era,
is_init_request,
session_id.clone(),
protocol_version.clone(),
http_middleware,
&http_context,
)
.await?;
guard_legacy_version_with_middleware(
&state,
era,
is_init_request,
is_v2_request,
session_id.as_ref(),
protocol_version.as_ref(),
http_middleware,
&http_context,
)
.await?;
let auth_context =
extract_auth_with_middleware(&state, &server_request, http_middleware, &http_context)
.await?;
Ok(Box::pin(dispatch_message_with_middleware(
&state,
ingress,
MiddlewareDispatch {
is_init_request,
response_session_id,
protocol_context,
v2_outbound,
sessions_on,
},
auth_context,
http_middleware,
&http_context,
))
.await)
}
async fn handle_get_sse(State(state): State<ServerState>, headers: HeaderMap) -> impl IntoResponse {
if let Some(rejection) = v2_verb_rejection(&state, &headers, "GET").await {
return rejection;
}
v1::handle_get_sse_body(&state, &headers).await
}
async fn handle_delete_session(
State(state): State<ServerState>,
headers: HeaderMap,
) -> impl IntoResponse {
if let Some(rejection) = v2_verb_rejection(&state, &headers, "DELETE").await {
return rejection;
}
v1::handle_delete_body(&state, &headers)
}
#[cfg(test)]
mod tests {
use super::*;
use super::v1::{apply_session_header, sessions_active_for};
use crate::types::protocol::Era;
const V1_HALF_IS_COMPILED: bool = cfg!(feature = "v1-compat");
#[test]
fn sessions_active_truth_table() {
assert!(!sessions_active_for(true, Some(Era::V2)));
assert_eq!(
sessions_active_for(true, Some(Era::V1)),
V1_HALF_IS_COMPILED
);
assert_eq!(sessions_active_for(true, None), V1_HALF_IS_COMPILED);
assert!(!sessions_active_for(false, Some(Era::V2)));
assert!(!sessions_active_for(false, Some(Era::V1)));
assert!(!sessions_active_for(false, None));
}
#[test]
fn v2_always_suppresses_sessions() {
for cfg in [true, false] {
assert!(
!sessions_active_for(cfg, Some(Era::V2)),
"v2 must be session-free with cfg_has_generator = {cfg}"
);
}
}
#[test]
fn session_header_is_never_emitted_when_sessions_are_inactive() {
let sid = "sess-123".to_string();
let mut headers = HeaderMap::new();
apply_session_header(&mut headers, Some(&sid), false);
assert!(
headers.get(MCP_SESSION_ID).is_none(),
"sessions inactive → no Mcp-Session-Id"
);
let mut headers = HeaderMap::new();
apply_session_header(&mut headers, Some(&sid), true);
assert_eq!(
headers.get(MCP_SESSION_ID).and_then(|v| v.to_str().ok()),
V1_HALF_IS_COMPILED.then_some("sess-123"),
);
let mut headers = HeaderMap::new();
apply_session_header(&mut headers, None, true);
assert!(headers.get(MCP_SESSION_ID).is_none());
let bad = "bad\nvalue".to_string();
let mut headers = HeaderMap::new();
apply_session_header(&mut headers, Some(&bad), true);
assert!(headers.get(MCP_SESSION_ID).is_none());
}
proptest::proptest! {
#[test]
fn sessions_active_is_exactly_its_stated_expression(
cfg_has_generator in proptest::prelude::any::<bool>(),
era_code in 0u8..3,
) {
let era = match era_code {
0 => None,
1 => Some(Era::V1),
_ => Some(Era::V2),
};
let expected =
V1_HALF_IS_COMPILED && !matches!(era, Some(Era::V2)) && cfg_has_generator;
proptest::prop_assert_eq!(sessions_active_for(cfg_has_generator, era), expected);
}
}
#[test]
fn extract_custom_claim_header_inserted_under_cognito_key() {
let mut h = HeaderMap::new();
h.insert("x-pmcp-user-id", "user-123".parse().unwrap());
h.insert(
"x-pmcp-claim-custom-primary-creator",
"rosen".parse().unwrap(),
);
let ctx = extract_auth_from_proxy_headers(&h).expect("auth ctx");
assert_eq!(
ctx.claims.get("custom:primary_creator"),
Some(&serde_json::Value::String("rosen".into())),
);
}
#[test]
#[allow(clippy::unnecessary_get_then_check)]
fn extract_custom_claim_empty_value_dropped() {
let mut h = HeaderMap::new();
h.insert("x-pmcp-user-id", "user-123".parse().unwrap());
h.insert("x-pmcp-claim-custom-empty", "".parse().unwrap());
let ctx = extract_auth_from_proxy_headers(&h).expect("auth ctx");
assert!(ctx.claims.get("custom:empty").is_none());
}
#[test]
fn extract_custom_claim_kebab_to_snake() {
let mut h = HeaderMap::new();
h.insert("x-pmcp-user-id", "u".parse().unwrap());
h.insert(
"x-pmcp-claim-custom-promo-code",
"SUMMER25".parse().unwrap(),
);
let ctx = extract_auth_from_proxy_headers(&h).expect("auth ctx");
assert_eq!(
ctx.claims.get("custom:promo_code"),
Some(&serde_json::Value::String("SUMMER25".into())),
);
}
#[test]
fn extract_custom_claim_coexists_with_standard_headers() {
let mut h = HeaderMap::new();
h.insert("x-pmcp-user-id", "u".parse().unwrap());
h.insert("x-pmcp-user-email", "u@example.com".parse().unwrap());
h.insert("x-pmcp-user-groups", "g1,g2".parse().unwrap());
h.insert("x-pmcp-claim-custom-tier", "gold".parse().unwrap());
let ctx = extract_auth_from_proxy_headers(&h).expect("auth ctx");
assert_eq!(ctx.subject, "u");
assert_eq!(ctx.claims["email"], "u@example.com");
assert_eq!(ctx.claims["custom:tier"], "gold");
}
use crate::types::protocol::error_codes::{HEADER_MISMATCH, METHOD_NOT_FOUND};
use crate::types::protocol::PROTOCOL_VERSION_2026_07_28 as V2;
fn headers_from(pairs: &[(&str, &str)]) -> HeaderMap {
let mut h = HeaderMap::new();
for (k, v) in pairs {
let name = http::header::HeaderName::from_bytes(k.as_bytes()).unwrap();
h.insert(name, HeaderValue::from_str(v).unwrap());
}
h
}
#[test]
fn decode_version_header_classifies_each_kind() {
assert_eq!(
decode_version_header(&headers_from(&[])),
HeaderProtocolVersion::Absent
);
assert_eq!(
decode_version_header(&headers_from(&[(MCP_PROTOCOL_VERSION, V2)])),
HeaderProtocolVersion::V2
);
assert_eq!(
decode_version_header(&headers_from(&[(MCP_PROTOCOL_VERSION, "2025-11-25")])),
HeaderProtocolVersion::Other
);
let big = "x".repeat(MAX_V2_HEADER_VALUE_LEN + 1);
assert_eq!(
decode_version_header(&headers_from(&[(MCP_PROTOCOL_VERSION, &big)])),
HeaderProtocolVersion::Malformed
);
}
#[test]
fn classify_era_cell_covers_every_matrix_cell() {
assert!(matches!(
classify_era_cell(HeaderProtocolVersion::V2, true),
V2Classification::Enforce
));
assert!(matches!(
classify_era_cell(HeaderProtocolVersion::Other, false),
V2Classification::Legacy
));
assert!(matches!(
classify_era_cell(HeaderProtocolVersion::Absent, false),
V2Classification::Legacy
));
assert!(matches!(
classify_era_cell(HeaderProtocolVersion::V2, false),
V2Classification::Reject(HEADER_MISMATCH, _)
));
assert!(matches!(
classify_era_cell(HeaderProtocolVersion::Absent, true),
V2Classification::Reject(HEADER_MISMATCH, _)
));
assert!(matches!(
classify_era_cell(HeaderProtocolVersion::Malformed, true),
V2Classification::Reject(HEADER_MISMATCH, _)
));
}
#[test]
fn require_v2_headers_truth_table() {
let name_bearing = NAME_BEARING_METHODS[0];
let ok = headers_from(&[
(MCP_PROTOCOL_VERSION, V2),
(MCP_METHOD, name_bearing),
(MCP_NAME, "search"),
]);
assert_eq!(
require_v2_headers(&ok).unwrap(),
(name_bearing.to_string(), "search".to_string())
);
let missing = headers_from(&[(MCP_PROTOCOL_VERSION, V2), (MCP_METHOD, name_bearing)]);
assert_eq!(require_v2_headers(&missing), Err(ERR_MISSING_MCP_NAME));
for method in NAME_LESS_METHODS {
let h = headers_from(&[(MCP_PROTOCOL_VERSION, V2), (MCP_METHOD, method)]);
assert_eq!(
require_v2_headers(&h),
Ok((method.to_string(), String::new()))
);
let stray = headers_from(&[
(MCP_PROTOCOL_VERSION, V2),
(MCP_METHOD, method),
(MCP_NAME, "attacker-supplied"),
]);
assert_eq!(
require_v2_headers(&stray),
Ok((method.to_string(), String::new()))
);
}
for method in NAME_BEARING_METHODS {
let h = headers_from(&[(MCP_PROTOCOL_VERSION, V2), (MCP_METHOD, method)]);
assert_eq!(require_v2_headers(&h), Err(ERR_MISSING_MCP_NAME));
}
let no_method = headers_from(&[(MCP_PROTOCOL_VERSION, V2), (MCP_NAME, "search")]);
assert_eq!(require_v2_headers(&no_method), Err(ERR_MISSING_V2_HEADERS));
for method in NAME_BEARING_METHODS.iter().chain(NAME_LESS_METHODS.iter()) {
let h = headers_from(&[(MCP_METHOD, method), (MCP_NAME, "search")]);
assert_eq!(require_v2_headers(&h), Err(ERR_MISSING_V2_HEADERS));
}
assert!(
!ERR_MISSING_V2_HEADERS.contains("Mcp-Name"),
"the universally-required-headers message must not name a header that is \
only conditionally required (Phase 118 D-13)"
);
assert!(ERR_MISSING_MCP_NAME.contains("Mcp-Name"));
assert!(ERR_MISSING_MCP_NAME.contains("routing name"));
}
#[test]
fn is_name_bearing_method_matches_the_literal_contract() {
for method in NAME_BEARING_METHODS {
assert!(
is_name_bearing_method(method),
"{method} carries a routing name and MUST be name-bearing (D-18)"
);
}
for method in NAME_LESS_METHODS {
assert!(
!is_name_bearing_method(method),
"{method} carries no routing name and MUST NOT be name-bearing"
);
}
}
#[test]
fn cross_check_method_and_name_fail_closed() {
assert!(cross_check_method("tools/call", Some("tools/call")).is_ok());
assert!(cross_check_method("tools/call", Some("resources/read")).is_err());
assert!(cross_check_method("tools/call", None).is_err());
assert!(cross_check_name("search", "tools/call", Some("search")).is_ok());
assert!(cross_check_name("search", "tools/call", Some("other")).is_err());
assert!(cross_check_name("search", "tools/call", None).is_err());
assert!(cross_check_name("anything", "tools/list", None).is_ok());
}
#[test]
fn classify_v2_request_accepts_well_formed_v2() {
let h = headers_from(&[
(MCP_PROTOCOL_VERSION, V2),
(MCP_METHOD, "tools/call"),
(MCP_NAME, "search"),
]);
let out = classify_v2_request(&h, true, Some("tools/call"), Some("search"));
assert!(matches!(out, V2GateOutcome::EnforceOk { .. }));
}
#[test]
fn classify_v2_request_rejects_method_body_mismatch() {
let h = headers_from(&[
(MCP_PROTOCOL_VERSION, V2),
(MCP_METHOD, "tools/call"),
(MCP_NAME, "search"),
]);
let out = classify_v2_request(&h, true, Some("resources/read"), Some("search"));
assert!(matches!(
out,
V2GateOutcome::Reject {
code: HEADER_MISMATCH,
..
}
));
}
#[test]
fn name_less_method_with_empty_mcp_name_is_enforce_ok() {
let h = headers_from(&[
(MCP_PROTOCOL_VERSION, V2),
(MCP_METHOD, "tools/list"),
(MCP_NAME, ""),
]);
let out = classify_v2_request(&h, true, Some("tools/list"), None);
assert!(
matches!(out, V2GateOutcome::EnforceOk { .. }),
"an EMPTY Mcp-Name on a name-less v2 method must be ACCEPTED"
);
}
#[test]
fn name_less_method_with_absent_mcp_name_is_accepted() {
let h = headers_from(&[(MCP_PROTOCOL_VERSION, V2), (MCP_METHOD, "tools/list")]);
let out = classify_v2_request(&h, true, Some("tools/list"), None);
assert!(
matches!(out, V2GateOutcome::EnforceOk { .. }),
"an ABSENT Mcp-Name on a name-LESS method must be ACCEPTED (D-13)"
);
}
#[test]
fn sentinel_encoded_mcp_name_matches_a_non_ascii_body_name() {
let name = "日本語ツール";
let encoded = crate::types::mrtr::encode_header_value(name);
assert_ne!(encoded, name, "a non-ASCII name must be sentinel-encoded");
assert!(cross_check_name(&encoded, "tools/call", Some(name)).is_ok());
assert!(cross_check_name(&encoded, "tools/call", Some("other")).is_err());
let h = headers_from(&[
(MCP_PROTOCOL_VERSION, V2),
(MCP_METHOD, "tools/call"),
(MCP_NAME, &encoded),
]);
let out = classify_v2_request(&h, true, Some("tools/call"), Some(name));
assert!(matches!(out, V2GateOutcome::EnforceOk { .. }));
}
#[test]
fn malformed_mcp_name_sentinel_is_a_header_mismatch() {
for bad in ["=?base64?not-base64!!", "=?base64?%%%%?="] {
assert!(
cross_check_name(bad, "tools/call", Some("search")).is_err(),
"malformed sentinel `{bad}` must be rejected"
);
let h = headers_from(&[
(MCP_PROTOCOL_VERSION, V2),
(MCP_METHOD, "tools/call"),
(MCP_NAME, bad),
]);
let out = classify_v2_request(&h, true, Some("tools/call"), Some("search"));
assert!(matches!(
out,
V2GateOutcome::Reject {
code: HEADER_MISMATCH,
..
}
));
}
}
#[test]
fn v2_status_table_covers_every_transport_code() {
use crate::types::protocol::error_codes as ec;
assert_eq!(
v2_status_for_code(ec::METHOD_NOT_FOUND),
StatusCode::NOT_FOUND
);
for code in [
ec::HEADER_MISMATCH,
ec::MISSING_REQUIRED_CLIENT_CAPABILITY,
ec::UNSUPPORTED_PROTOCOL_VERSION,
ec::PARSE_ERROR,
ec::INVALID_REQUEST,
ec::INVALID_PARAMS,
] {
assert_eq!(
v2_status_for_code(code),
StatusCode::BAD_REQUEST,
"{code} must map to 400 on v2"
);
}
for code in [ec::INTERNAL_ERROR, ec::REQUEST_TIMEOUT, ec::V1_TASK_PENDING] {
assert_eq!(v2_status_for_code(code), StatusCode::OK);
}
}
#[test]
fn status_mapping_is_era_gated_so_v1_is_untouched() {
use crate::types::protocol::Era;
for era in [None, Some(Era::V1)] {
for code in [
METHOD_NOT_FOUND,
HEADER_MISMATCH,
crate::types::protocol::error_codes::PARSE_ERROR,
] {
assert_eq!(status_for_error(era, code, StatusCode::OK), StatusCode::OK);
}
}
assert_eq!(
status_for_error(Some(Era::V2), METHOD_NOT_FOUND, StatusCode::OK),
StatusCode::NOT_FOUND
);
}
#[test]
fn raw_request_id_survives_a_body_that_never_typed_parses() {
assert_eq!(
raw_request_id(br#"{"jsonrpc":"2.0","id":7,"method":"totally/unknown"}"#),
serde_json::json!(7)
);
assert_eq!(
raw_request_id(br#"{"jsonrpc":"2.0","id":"abc","method":"nope","params":{}}"#),
serde_json::json!("abc")
);
assert_eq!(
raw_request_id(br#"{"jsonrpc":"2.0","method":"notify"}"#),
serde_json::Value::Null
);
assert_eq!(raw_request_id(b"{not json"), serde_json::Value::Null);
assert_eq!(raw_request_id(&[0xff, 0xfe, 0x00]), serde_json::Value::Null);
}
#[test]
fn v2_dispatch_status_reads_the_code_not_the_call_site() {
use crate::types::jsonrpc::{JSONRPCError, ResponsePayload};
use crate::types::protocol::Era;
let error_response = |code: i32| crate::types::JSONRPCResponse {
jsonrpc: "2.0".to_string(),
id: crate::types::RequestId::Number(1),
payload: ResponsePayload::Error(JSONRPCError {
code,
message: "x".to_string(),
data: None,
}),
};
assert_eq!(
v2_dispatch_response_status(
Some(Era::V2),
&error_response(
crate::types::protocol::error_codes::MISSING_REQUIRED_CLIENT_CAPABILITY
)
),
Some(StatusCode::BAD_REQUEST)
);
assert_eq!(
v2_dispatch_response_status(Some(Era::V2), &error_response(METHOD_NOT_FOUND)),
Some(StatusCode::NOT_FOUND)
);
assert_eq!(
v2_dispatch_response_status(Some(Era::V1), &error_response(METHOD_NOT_FOUND)),
None
);
assert_eq!(
v2_dispatch_response_status(None, &error_response(METHOD_NOT_FOUND)),
None
);
let ok = crate::types::JSONRPCResponse {
jsonrpc: "2.0".to_string(),
id: crate::types::RequestId::Number(1),
payload: ResponsePayload::Result(serde_json::json!({})),
};
assert_eq!(v2_dispatch_response_status(Some(Era::V2), &ok), None);
}
#[test]
fn v2_method_not_allowed_only_fires_on_the_v2_version_header() {
for verb in ["GET", "DELETE"] {
let h = headers_from(&[(MCP_PROTOCOL_VERSION, V2)]);
let response = v2_method_not_allowed(&h, verb, true).expect("v2 must be 405");
assert_eq!(response.status(), StatusCode::METHOD_NOT_ALLOWED);
}
assert!(v2_method_not_allowed(&headers_from(&[]), "GET", true).is_none());
assert!(v2_method_not_allowed(
&headers_from(&[(MCP_PROTOCOL_VERSION, "2025-11-25")]),
"GET",
true
)
.is_none());
let big = "x".repeat(MAX_V2_HEADER_VALUE_LEN + 1);
assert!(v2_method_not_allowed(
&headers_from(&[(MCP_PROTOCOL_VERSION, &big)]),
"DELETE",
true
)
.is_none());
for verb in ["GET", "DELETE"] {
assert!(
v2_method_not_allowed(&headers_from(&[(MCP_PROTOCOL_VERSION, V2)]), verb, false)
.is_none(),
"{verb}: a non-opted-in server must not answer 405"
);
}
}
#[test]
fn unsupported_version_reject_carries_a_supported_array() {
use crate::types::protocol::context::ProtocolNegotiationError;
let accept = vec![ProtocolVersion("2025-11-25".to_string()), v2_version()];
let outcome = negotiation_error_to_gate_reject(
&ProtocolNegotiationError::UnsupportedVersion("1999-01-01".to_string()),
&accept,
);
let V2GateOutcome::Reject { code, data, .. } = outcome else {
panic!("an unsupported version must reject");
};
assert_eq!(
code,
crate::types::protocol::error_codes::UNSUPPORTED_PROTOCOL_VERSION
);
let data = data.expect("UNSUPPORTED_PROTOCOL_VERSION MUST carry structured data");
assert!(
data["supported"].is_array(),
"data.supported must be an ARRAY: {data}"
);
assert_eq!(data["supported"][0], "2025-11-25");
assert_eq!(data["requested"], "1999-01-01");
let outcome = negotiation_error_to_gate_reject(
&ProtocolNegotiationError::MalformedMeta("bad"),
&accept,
);
let V2GateOutcome::Reject { code, data, .. } = outcome else {
panic!("malformed _meta must reject");
};
assert_eq!(code, crate::types::protocol::error_codes::INVALID_PARAMS);
assert!(data.is_none());
}
#[test]
fn extract_body_method_and_name_reads_wire_shape() {
let body = br#"{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"search"}}"#;
let (m, n) = extract_body_method_and_name(body);
assert_eq!(m.as_deref(), Some("tools/call"));
assert_eq!(n.as_deref(), Some("search"));
assert_eq!(extract_body_method_and_name(b"not json"), (None, None));
}
#[test]
fn extract_body_method_and_name_uses_uri_for_resources_read() {
let body = br#"{"jsonrpc":"2.0","id":1,"method":"resources/read","params":{"uri":"mem://greeting"}}"#;
let (m, n) = extract_body_method_and_name(body);
assert_eq!(m.as_deref(), Some("resources/read"));
assert_eq!(
n.as_deref(),
Some("mem://greeting"),
"resources/read logical name must come from params.uri"
);
let body =
br#"{"jsonrpc":"2.0","id":1,"method":"prompts/get","params":{"name":"greeting"}}"#;
let (m, n) = extract_body_method_and_name(body);
assert_eq!(m.as_deref(), Some("prompts/get"));
assert_eq!(n.as_deref(), Some("greeting"));
let body = br#"{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"search"}}"#;
let (m, n) = extract_body_method_and_name(body);
assert_eq!(m.as_deref(), Some("tools/call"));
assert_eq!(n.as_deref(), Some("search"));
let body =
br#"{"jsonrpc":"2.0","id":1,"method":"resources/read","params":{"uri":"file:///x"}}"#;
let (_, n) = extract_body_method_and_name(body);
assert_eq!(n.as_deref(), Some("file:///x"));
}
#[test]
fn cross_check_name_accepts_resources_read_uri() {
let uri = "mem://greeting";
assert!(cross_check_name(uri, "resources/read", Some(uri)).is_ok());
assert!(cross_check_name(uri, "resources/read", Some("mem://other")).is_err());
assert!(cross_check_name(uri, "resources/read", None).is_err());
}
const NAME_BEARING_METHODS: [&str; 6] = [
"tools/call",
"prompts/get",
"resources/read",
"tasks/get",
"tasks/update",
"tasks/cancel",
];
const NAME_LESS_METHODS: [&str; 4] = [
"tools/list",
"ping",
"completion/complete",
"server/discover",
];
#[test]
fn classify_v2_request_requires_mcp_name_only_on_name_bearing_methods() {
for method in NAME_LESS_METHODS {
let h = headers_from(&[(MCP_PROTOCOL_VERSION, V2), (MCP_METHOD, method)]);
let out = classify_v2_request(&h, true, Some(method), None);
assert!(
matches!(out, V2GateOutcome::EnforceOk { .. }),
"{method} carries no routing name, so a missing Mcp-Name must be accepted"
);
}
for method in NAME_BEARING_METHODS {
let h = headers_from(&[(MCP_PROTOCOL_VERSION, V2), (MCP_METHOD, method)]);
let out = classify_v2_request(&h, true, Some(method), Some("x"));
assert!(
matches!(out, V2GateOutcome::Reject { .. }),
"{method} carries a routing name, so a missing Mcp-Name must be rejected"
);
}
let h = headers_from(&[
(MCP_PROTOCOL_VERSION, V2),
(MCP_METHOD, "tasks/get"),
(MCP_NAME, "task-a"),
]);
assert!(matches!(
classify_v2_request(&h, true, Some("tasks/get"), Some("task-b")),
V2GateOutcome::Reject { .. }
));
let h = headers_from(&[
(MCP_PROTOCOL_VERSION, V2),
(MCP_METHOD, "tasks/get"),
(MCP_NAME, "task-a"),
]);
assert!(matches!(
classify_v2_request(&h, true, Some("tasks/get"), Some("task-a")),
V2GateOutcome::EnforceOk { .. }
));
let h = headers_from(&[
(MCP_PROTOCOL_VERSION, V2),
(MCP_METHOD, "tools/list"),
(MCP_NAME, "attacker-supplied"),
]);
match classify_v2_request(&h, true, Some("tools/list"), None) {
V2GateOutcome::EnforceOk { method, name } => {
assert_eq!(method, "tools/list");
assert_eq!(name, "", "a stray Mcp-Name must be sanitized to empty");
},
V2GateOutcome::Reject { code, message, .. } => {
panic!("expected EnforceOk, got Reject({code}, {message})")
},
V2GateOutcome::Passthrough => panic!("expected EnforceOk, got Passthrough"),
}
}
#[test]
fn apply_v2_outbound_headers_sets_all_three_without_panic() {
let mut h = HeaderMap::new();
apply_v2_outbound_headers(&mut h, "tools/call", "search");
assert_eq!(h.get(MCP_METHOD).unwrap(), "tools/call");
assert_eq!(h.get(MCP_NAME).unwrap(), "search");
assert_eq!(h.get(MCP_PROTOCOL_VERSION).unwrap(), V2);
}
proptest::proptest! {
#[test]
fn v2_header_gate_proptest(
header_kind in 0u8..4,
meta_is_v2 in proptest::bool::ANY,
have_method in proptest::bool::ANY,
have_name in proptest::bool::ANY,
method_val in "[a-z/]{0,20}",
name_val in "[a-z]{0,20}",
body_method in proptest::option::of("[a-z/]{0,20}"),
body_name in proptest::option::of("[a-z]{0,20}"),
) {
let mut pairs: Vec<(&str, String)> = Vec::new();
match header_kind {
0 => {}, 1 => pairs.push((MCP_PROTOCOL_VERSION, V2.to_string())),
2 => pairs.push((MCP_PROTOCOL_VERSION, "2025-11-25".to_string())),
_ => pairs.push((MCP_PROTOCOL_VERSION, "\u{ff}bogus".to_string())),
}
if have_method {
pairs.push((MCP_METHOD, method_val.clone()));
}
if have_name {
pairs.push((MCP_NAME, name_val.clone()));
}
let mut h = HeaderMap::new();
for (k, v) in &pairs {
if let Ok(hv) = HeaderValue::from_str(v) {
let name = http::header::HeaderName::from_bytes(k.as_bytes()).unwrap();
h.insert(name, hv);
}
}
let out = classify_v2_request(&h, meta_is_v2, body_method.as_deref(), body_name.as_deref());
let header_is_v2 = decode_version_header(&h) == HeaderProtocolVersion::V2;
match out {
V2GateOutcome::Passthrough => {
proptest::prop_assert!(!header_is_v2 && !meta_is_v2);
},
V2GateOutcome::EnforceOk { ref name, .. } => {
proptest::prop_assert!(header_is_v2 && meta_is_v2);
proptest::prop_assert!(have_method);
proptest::prop_assert!(have_name || !is_name_bearing_method(&method_val));
if !is_name_bearing_method(&method_val) {
proptest::prop_assert!(name.is_empty());
}
},
V2GateOutcome::Reject { code, .. } => {
proptest::prop_assert_eq!(code, HEADER_MISMATCH);
},
}
}
}
fn any_v2_method() -> impl proptest::strategy::Strategy<Value = String> {
use proptest::strategy::Strategy as _;
proptest::prop_oneof![
proptest::sample::select(NAME_BEARING_METHODS.as_slice()).prop_map(str::to_string),
proptest::sample::select(NAME_LESS_METHODS.as_slice()).prop_map(str::to_string),
"[a-z/]{0,20}",
]
}
proptest::proptest! {
#[test]
fn require_v2_headers_is_exactly_its_truth_table(
have_version in proptest::bool::ANY,
have_method in proptest::bool::ANY,
have_name in proptest::bool::ANY,
method in any_v2_method(),
name_val in "[a-zA-Z0-9._-]{0,40}",
) {
let mut h = HeaderMap::new();
if have_version {
h.insert(
http::header::HeaderName::from_bytes(MCP_PROTOCOL_VERSION.as_bytes()).unwrap(),
HeaderValue::from_static(crate::types::protocol::PROTOCOL_VERSION_2026_07_28),
);
}
if have_method {
if let Ok(v) = HeaderValue::from_str(&method) {
h.insert(
http::header::HeaderName::from_bytes(MCP_METHOD.as_bytes()).unwrap(),
v,
);
}
}
if have_name {
h.insert(
http::header::HeaderName::from_bytes(MCP_NAME.as_bytes()).unwrap(),
HeaderValue::from_str(&name_val).unwrap(),
);
}
let out = require_v2_headers(&h);
let expected_ok =
have_version && have_method && (have_name || !is_name_bearing_method(&method));
proptest::prop_assert_eq!(out.is_ok(), expected_ok);
if let Ok((got_method, got_name)) = out {
proptest::prop_assert_eq!(&got_method, &method);
if is_name_bearing_method(&got_method) {
proptest::prop_assert_eq!(&got_name, &name_val);
} else {
proptest::prop_assert!(got_name.is_empty());
}
}
}
#[test]
fn v2_header_gate_never_panics_on_arbitrary_bytes(
version_bytes in proptest::collection::vec(proptest::num::u8::ANY, 0..40),
method_bytes in proptest::collection::vec(proptest::num::u8::ANY, 0..40),
name_bytes in proptest::collection::vec(proptest::num::u8::ANY, 0..40),
body_bytes in proptest::collection::vec(proptest::num::u8::ANY, 0..120),
meta_is_v2 in proptest::bool::ANY,
) {
let mut h = HeaderMap::new();
for (header_name, raw) in [
(MCP_PROTOCOL_VERSION, &version_bytes),
(MCP_METHOD, &method_bytes),
(MCP_NAME, &name_bytes),
] {
if let Ok(value) = HeaderValue::from_bytes(raw) {
h.insert(
http::header::HeaderName::from_bytes(header_name.as_bytes()).unwrap(),
value,
);
}
}
let _ = require_v2_headers(&h);
let (body_method, body_name) = extract_body_method_and_name(&body_bytes);
let out = classify_v2_request(
&h,
meta_is_v2,
body_method.as_deref(),
body_name.as_deref(),
);
if let V2GateOutcome::Reject { code, .. } = out {
proptest::prop_assert_eq!(code, HEADER_MISMATCH);
}
}
}
use crate::types::ProtocolVersion;
fn v2_version() -> ProtocolVersion {
ProtocolVersion(crate::types::protocol::PROTOCOL_VERSION_2026_07_28.to_string())
}
fn state_with_accept(accept: Vec<ProtocolVersion>) -> ServerState {
let server = Server::builder()
.name("raw-gate-test")
.version("1.0.0")
.with_supported_protocol_versions(accept)
.build()
.expect("server builds");
make_server_state(
Arc::new(tokio::sync::Mutex::new(server)),
StreamableHttpServerConfig::default(),
)
}
#[test]
fn server_discover_is_not_name_bearing() {
assert!(!is_name_bearing_method("server/discover"));
}
#[test]
fn classify_http_ingress_routes_server_discover() {
let body = br#"{"jsonrpc":"2.0","id":7,"method":"server/discover","params":{"_meta":{"io.modelcontextprotocol/protocolVersion":"2026-07-28"}}}"#;
let ingress = classify_http_ingress(body).expect("server/discover classifies");
match ingress {
HttpIngress::Discover { id } => {
assert_eq!(id, crate::types::RequestId::from(7i64));
assert_eq!(
raw_params_meta(body).unwrap()["io.modelcontextprotocol/protocolVersion"],
"2026-07-28"
);
},
HttpIngress::Public(_)
| HttpIngress::SubscriptionsListen { .. }
| HttpIngress::TasksUpdate { .. } => {
panic!("server/discover must classify as Discover")
},
}
let tools = br#"{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"x"}}"#;
assert!(classify_http_ingress(tools).is_none());
let notif = br#"{"jsonrpc":"2.0","method":"server/discover"}"#;
assert!(classify_http_ingress(notif).is_none());
assert!(classify_http_ingress(b"not json").is_none());
}
#[test]
fn classify_http_ingress_routes_tasks_update_with_raw_params() {
let body = br#"{"jsonrpc":"2.0","id":"u-1","method":"tasks/update","params":{"taskId":42,"junk":[1]}}"#;
let ingress = classify_http_ingress(body).expect("tasks/update classifies");
match ingress {
HttpIngress::TasksUpdate { id, params } => {
assert_eq!(id, crate::types::RequestId::from("u-1".to_string()));
assert_eq!(
params,
serde_json::json!({ "taskId": 42, "junk": [1] }),
"the params must reach the served branch UNDECODED"
);
},
HttpIngress::Public(_)
| HttpIngress::Discover { .. }
| HttpIngress::SubscriptionsListen { .. } => {
panic!("tasks/update must classify as TasksUpdate")
},
}
let notif = br#"{"jsonrpc":"2.0","method":"tasks/update","params":{}}"#;
assert!(classify_http_ingress(notif).is_none());
}
fn v2_body_bytes(method: &str, key: &str) -> Vec<u8> {
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": method,
"params": { key: { "io.modelcontextprotocol/protocolVersion": "2026-07-28" } },
})
.to_string()
.into_bytes()
}
fn v2_headers_for(method: &str) -> HeaderMap {
headers_from(&[
(MCP_PROTOCOL_VERSION, V2),
(MCP_METHOD, method),
(MCP_NAME, ""),
])
}
#[test]
fn raw_params_meta_reads_the_spec_spelling_and_the_legacy_alias() {
let expected = serde_json::json!({ "k": "v" });
assert_eq!(
raw_params_meta(
br#"{"jsonrpc":"2.0","id":1,"method":"m","params":{"_meta":{"k":"v"}}}"#
),
Some(expected.clone())
);
assert_eq!(
raw_params_meta(
br#"{"jsonrpc":"2.0","id":1,"method":"m","params":{"meta":{"k":"v"}}}"#
),
Some(expected.clone()),
"the legacy `meta` spelling is accepted, mirroring the typed serde alias"
);
assert_eq!(
raw_params_meta(
br#"{"jsonrpc":"2.0","id":1,"method":"m","params":{"_meta":{"k":"v"},"meta":{"k":"other"}}}"#
),
Some(expected)
);
assert_eq!(raw_params_meta(br#"{"jsonrpc":"2.0","params":{}}"#), None);
assert_eq!(
raw_params_meta(br#"{"jsonrpc":"2.0","params":{"_meta":null}}"#),
None
);
assert_eq!(raw_params_meta(br#"{"jsonrpc":"2.0","id":1}"#), None);
assert_eq!(raw_params_meta(b"not json"), None);
assert_eq!(raw_params_meta(&[0xff, 0xfe, 0x00]), None);
}
fn accepted_v2() -> V2GateOutcome {
V2GateOutcome::EnforceOk {
method: "tools/call".to_string(),
name: "search".to_string(),
}
}
fn v2_context() -> crate::types::protocol::ProtocolContext {
crate::types::protocol::ProtocolContext::new(crate::types::protocol::Era::V2, v2_version())
}
fn mrtr_test_method() -> &'static str {
assert!(
crate::types::mrtr::mrtr_eligible("tools/call"),
"these tests exercise the MRTR extraction, which only runs for an eligible method"
);
"tools/call"
}
fn mrtr_body(extra: &serde_json::Value) -> Vec<u8> {
let mut params = serde_json::json!({ "name": "search", "arguments": {} });
if let (Some(target), Some(source)) = (params.as_object_mut(), extra.as_object()) {
for (key, value) in source {
target.insert(key.clone(), value.clone());
}
}
serde_json::json!({
"jsonrpc": "2.0", "id": 1, "method": "tools/call", "params": params,
})
.to_string()
.into_bytes()
}
#[test]
fn attach_v2_mrtr_params_lands_the_fields_on_the_context() {
let body = mrtr_body(&serde_json::json!({
"requestState": "opaque-token",
"inputResponses": { "user_name": { "action": "accept" } },
}));
let parsed = raw_body_json(&body);
let (ctx, outcome) = attach_v2_mrtr_params(
Some(v2_context()),
accepted_v2(),
parsed.as_ref(),
Some(mrtr_test_method()),
);
assert!(matches!(outcome, V2GateOutcome::EnforceOk { .. }));
let ctx = ctx.expect("context survives");
assert_eq!(ctx.request_state_token(), Some("opaque-token"));
assert!(ctx.input_responses().is_some());
}
#[test]
fn attach_v2_mrtr_params_skips_a_non_accepted_request() {
let body = mrtr_body(&serde_json::json!({ "requestState": "opaque-token" }));
let parsed = raw_body_json(&body);
for outcome in [
V2GateOutcome::Passthrough,
V2GateOutcome::Reject {
code: crate::types::protocol::error_codes::HEADER_MISMATCH,
message: "nope".to_string(),
data: None,
},
] {
let (ctx, _) = attach_v2_mrtr_params(
Some(v2_context()),
outcome,
parsed.as_ref(),
Some(mrtr_test_method()),
);
assert!(
ctx.expect("context survives")
.request_state_token()
.is_none(),
"MRTR extraction must not run outside the accepted v2 path"
);
}
}
#[test]
fn attach_v2_mrtr_params_absent_fields_are_the_default() {
let body = mrtr_body(&serde_json::json!({}));
let parsed = raw_body_json(&body);
let (ctx, outcome) = attach_v2_mrtr_params(
Some(v2_context()),
accepted_v2(),
parsed.as_ref(),
Some(mrtr_test_method()),
);
assert!(matches!(outcome, V2GateOutcome::EnforceOk { .. }));
let ctx = ctx.expect("context survives");
assert!(ctx.request_state_token().is_none());
assert!(ctx.input_responses().is_none());
}
#[test]
fn attach_v2_mrtr_params_rejects_every_malformed_shape() {
use crate::types::mrtr::{
MAX_INPUT_RESPONSES, MAX_INPUT_RESPONSE_BYTES, MAX_INPUT_RESPONSE_DEPTH,
MAX_REQUEST_STATE_LEN,
};
let mut too_many = serde_json::Map::new();
for index in 0..=MAX_INPUT_RESPONSES {
too_many.insert(
format!("k{index}"),
serde_json::json!({ "action": "accept" }),
);
}
let mut chunky = serde_json::Map::new();
for index in 0..8 {
chunky.insert(
format!("k{index}"),
serde_json::json!({
"action": "accept",
"content": { "v": "z".repeat(MAX_INPUT_RESPONSE_BYTES - 1_000) }
}),
);
}
let mut nested = serde_json::json!("leaf");
for _ in 0..(MAX_INPUT_RESPONSE_DEPTH + 4) {
nested = serde_json::json!({ "n": nested });
}
let cases = [
serde_json::json!({ "requestState": 42 }),
serde_json::json!({ "requestState": "x".repeat(MAX_REQUEST_STATE_LEN + 1) }),
serde_json::json!({ "inputResponses": [] }),
serde_json::json!({ "inputResponses": too_many }),
serde_json::json!({ "inputResponses": {
"big": { "action": "accept",
"content": { "v": "y".repeat(MAX_INPUT_RESPONSE_BYTES + 1) } } } }),
serde_json::json!({ "inputResponses": chunky }),
serde_json::json!({ "inputResponses": {
"deep": { "action": "accept", "content": { "v": nested } } } }),
serde_json::json!({ "inputResponses": { "bad": { "totally": "wrong" } } }),
];
for case in cases {
let body = mrtr_body(&case);
let parsed = raw_body_json(&body);
let (_, outcome) = attach_v2_mrtr_params(
Some(v2_context()),
accepted_v2(),
parsed.as_ref(),
Some(mrtr_test_method()),
);
let V2GateOutcome::Reject { code, .. } = outcome else {
panic!("a present-but-unusable MRTR field must REJECT, got a pass for {case}");
};
assert_eq!(
code,
crate::types::protocol::error_codes::INVALID_PARAMS,
"malformed MRTR maps to -32602 for {case}"
);
assert_eq!(
v2_status_for_code(code),
StatusCode::BAD_REQUEST,
"a malformed MRTR field is a 400"
);
}
}
#[test]
fn attach_v2_mrtr_params_ignores_a_non_eligible_method() {
let malformed = serde_json::json!({ "inputResponses": "not-an-object" });
for method in ["tasks/update", "tasks/get", "tools/list", "server/discover"] {
assert!(
!crate::types::mrtr::mrtr_eligible(method),
"{method} must be outside MRTR_METHODS for this test to mean anything"
);
let body = serde_json::json!({
"jsonrpc": "2.0", "id": 1, "method": method,
"params": { "taskId": "t-1", "inputResponses": "not-an-object" },
})
.to_string()
.into_bytes();
let parsed = raw_body_json(&body);
let (ctx, outcome) = attach_v2_mrtr_params(
Some(v2_context()),
accepted_v2(),
parsed.as_ref(),
Some(method),
);
assert!(
matches!(outcome, V2GateOutcome::EnforceOk { .. }),
"{method} is not an MRTR method, so its params must not be judged here \
(T-114-63/T-114-64); {malformed} was rejected"
);
let ctx = ctx.expect("context survives");
assert!(
ctx.input_responses().is_none(),
"{method} must carry NO MRTR-decoded inputResponses on the context"
);
assert!(
ctx.request_state_token().is_none(),
"{method} must carry NO MRTR requestState on the context"
);
}
}
#[test]
fn attach_v2_mrtr_params_skips_an_unresolvable_method() {
let body = mrtr_body(&serde_json::json!({ "requestState": "opaque-token" }));
let parsed = raw_body_json(&body);
let (ctx, outcome) =
attach_v2_mrtr_params(Some(v2_context()), accepted_v2(), parsed.as_ref(), None);
assert!(matches!(outcome, V2GateOutcome::EnforceOk { .. }));
assert!(ctx
.expect("context survives")
.request_state_token()
.is_none());
}
#[test]
fn attach_v2_mrtr_params_rejection_never_echoes_the_offending_value() {
let secret = "x".repeat(crate::types::mrtr::MAX_REQUEST_STATE_LEN + 1);
let body = mrtr_body(&serde_json::json!({
"inputResponses": { "super-secret-key": { "totally": "wrong" } },
"requestState": secret,
}));
let parsed = raw_body_json(&body);
let (_, outcome) = attach_v2_mrtr_params(
Some(v2_context()),
accepted_v2(),
parsed.as_ref(),
Some(mrtr_test_method()),
);
let V2GateOutcome::Reject { message, .. } = outcome else {
panic!("expected a rejection");
};
assert!(
!message.contains("super-secret-key"),
"message leaked an attacker-supplied key: {message}"
);
assert!(
!message.contains(&secret),
"message leaked the attacker-supplied value"
);
}
#[tokio::test]
async fn v2_gate_non_opted_in_passes_through() {
let state = state_with_accept(vec![ProtocolVersion("2025-11-25".to_string())]);
let headers = headers_from(&[(MCP_PROTOCOL_VERSION, V2)]);
let body = v2_body_bytes("server/discover", "_meta");
let (ctx, outcome) = run_v2_header_gate(
&state,
&headers,
&body,
Some(crate::types::protocol::SERVER_DISCOVER_METHOD),
)
.await;
assert!(ctx.is_none(), "non-opted-in resolves no context");
assert!(
matches!(outcome, V2GateOutcome::Passthrough),
"non-opted-in + v2 _meta must Passthrough, not Reject"
);
}
#[tokio::test]
async fn v2_gate_accepts_every_method_from_the_raw_body() {
let state = state_with_accept(vec![
ProtocolVersion("2025-11-25".to_string()),
v2_version(),
]);
for method in [
"tools/list",
"prompts/list",
"resources/list",
"resources/templates/list",
"completion/complete",
] {
let body = v2_body_bytes(method, "_meta");
let (ctx, outcome) =
run_v2_header_gate(&state, &v2_headers_for(method), &body, None).await;
assert_eq!(
ctx.map(|c| c.era),
Some(crate::types::protocol::Era::V2),
"{method} must resolve to the v2 era from its raw params._meta"
);
assert!(
matches!(outcome, V2GateOutcome::EnforceOk { .. }),
"{method} must be accepted as a v2 request"
);
}
}
#[tokio::test]
async fn v2_gate_discover_pins_its_method() {
let state = state_with_accept(vec![
ProtocolVersion("2025-11-25".to_string()),
v2_version(),
]);
let headers = v2_headers_for(crate::types::protocol::SERVER_DISCOVER_METHOD);
let body = v2_body_bytes("tools/call", "_meta");
let (ctx, outcome) = run_v2_header_gate(
&state,
&headers,
&body,
Some(crate::types::protocol::SERVER_DISCOVER_METHOD),
)
.await;
assert_eq!(ctx.map(|c| c.era), Some(crate::types::protocol::Era::V2));
assert!(matches!(outcome, V2GateOutcome::EnforceOk { .. }));
}
#[tokio::test]
async fn v2_gate_v2_meta_without_header_rejects() {
let state = state_with_accept(vec![
ProtocolVersion("2025-11-25".to_string()),
v2_version(),
]);
let headers = headers_from(&[(MCP_METHOD, "tools/list"), (MCP_NAME, "")]);
let body = v2_body_bytes("tools/list", "_meta");
let (_ctx, outcome) = run_v2_header_gate(&state, &headers, &body, None).await;
assert!(matches!(outcome, V2GateOutcome::Reject { .. }));
}
#[cfg(feature = "v1-compat")]
mod v1_resumability {
use super::super::v1::{resumability_active_for, resumability_store};
use super::*;
use crate::shared::http_constants::LAST_EVENT_ID;
fn dual_era_state() -> ServerState {
state_with_accept(vec![
ProtocolVersion(crate::LATEST_PROTOCOL_VERSION.to_string()),
v2_version(),
])
}
fn post_request(extra: &[(&str, &str)], body: &str) -> axum::extract::Request<Body> {
let mut builder = axum::http::Request::builder()
.method("POST")
.uri("/")
.header(header::CONTENT_TYPE, APPLICATION_JSON)
.header(
header::ACCEPT,
crate::shared::http_constants::ACCEPT_STREAMABLE,
);
for (name, value) in extra {
builder = builder.header(*name, *value);
}
builder
.body(Body::from(body.to_string()))
.expect("request builds")
}
fn v2_post_headers<'a>(
method: &'a str,
extra: &[(&'a str, &'a str)],
) -> Vec<(&'a str, &'a str)> {
let mut headers = vec![
(MCP_PROTOCOL_VERSION, V2),
(MCP_METHOD, method),
(MCP_NAME, ""),
];
headers.extend_from_slice(extra);
headers
}
#[derive(Debug, Default)]
struct SpyEventStore {
stores: std::sync::atomic::AtomicUsize,
replays: std::sync::atomic::AtomicUsize,
}
impl SpyEventStore {
fn stores(&self) -> usize {
self.stores.load(std::sync::atomic::Ordering::SeqCst)
}
fn replays(&self) -> usize {
self.replays.load(std::sync::atomic::Ordering::SeqCst)
}
}
#[async_trait]
impl EventStore for SpyEventStore {
async fn store_event(
&self,
_stream_id: &str,
_event_id: &str,
_message: &TransportMessage,
) -> Result<()> {
self.stores
.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
Ok(())
}
async fn replay_events_after(
&self,
_last_event_id: &str,
) -> Result<Vec<(String, TransportMessage)>> {
self.replays
.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
Ok(Vec::new())
}
async fn get_stream_for_event(&self, _event_id: &str) -> Result<Option<String>> {
Ok(None)
}
}
fn spy_state() -> (ServerState, Arc<SpyEventStore>) {
let spy = Arc::new(SpyEventStore::default());
let mut state = dual_era_state();
state.v1.event_store = Some(spy.clone() as v1::EventStoreHandle);
(state, spy)
}
#[test]
fn resumability_active_truth_table() {
assert!(!resumability_active_for(true, Some(Era::V2)));
assert!(resumability_active_for(true, Some(Era::V1)));
assert!(resumability_active_for(true, None));
assert!(!resumability_active_for(false, Some(Era::V2)));
assert!(!resumability_active_for(false, Some(Era::V1)));
assert!(!resumability_active_for(false, None));
}
#[test]
fn v2_always_suppresses_resumability() {
for cfg in [true, false] {
assert!(
!resumability_active_for(cfg, Some(Era::V2)),
"v2 must be resumability-free with cfg_has_event_store = {cfg}"
);
}
}
#[test]
fn resumability_store_is_the_gated_borrow() {
let (state, _spy) = spy_state();
assert!(
resumability_store(&state, Some(Era::V1)).is_some(),
"v1 keeps the store"
);
assert!(
resumability_store(&state, None).is_some(),
"a non-opted-in server keeps the store"
);
assert!(
resumability_store(&state, Some(Era::V2)).is_none(),
"v2 can never reach the store"
);
}
proptest::proptest! {
#[test]
fn resumability_active_is_exactly_its_stated_expression(
cfg_has_event_store in proptest::prelude::any::<bool>(),
era_code in 0u8..3,
) {
let era = match era_code {
0 => None,
1 => Some(Era::V1),
_ => Some(Era::V2),
};
let expected = !matches!(era, Some(Era::V2)) && cfg_has_event_store;
proptest::prop_assert_eq!(
resumability_active_for(cfg_has_event_store, era),
expected
);
}
}
#[tokio::test]
async fn spy_records_store_traffic_for_a_v1_exchange() {
let (state, spy) = spy_state();
let body = serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": crate::LATEST_PROTOCOL_VERSION,
"capabilities": {},
"clientInfo": { "name": "v1", "version": "1.0.0" },
},
})
.to_string();
let response = handle_post_fast_path(state, post_request(&[], &body)).await;
assert_eq!(response.status(), StatusCode::OK, "v1 initialize is served");
assert!(
spy.stores() > 0,
"a v1 exchange MUST still write to the event store — otherwise the \
v2 zero assertions are vacuous"
);
}
#[tokio::test]
async fn spy_records_zero_event_store_traffic_for_a_v2_exchange() {
let (state, spy) = spy_state();
let response = handle_post_fast_path(
state,
post_request(
&v2_post_headers("tools/list", &[(LAST_EVENT_ID, "12345")]),
&String::from_utf8(v2_body_bytes("tools/list", "_meta")).unwrap(),
),
)
.await;
assert_eq!(
response.status(),
StatusCode::OK,
"a v2 request carrying Last-Event-ID is served NORMALLY"
);
assert_eq!(spy.stores(), 0, "a v2 exchange must write NOTHING");
assert_eq!(spy.replays(), 0, "a v2 exchange must replay NOTHING");
}
#[tokio::test]
async fn spy_records_replay_for_a_v1_get_with_last_event_id() {
let (state, spy) = spy_state();
let headers = headers_from(&[
(http::header::ACCEPT.as_str(), TEXT_EVENT_STREAM),
(LAST_EVENT_ID, "evt-1"),
]);
let response = handle_get_sse(State(state), headers).await.into_response();
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(
spy.replays(),
1,
"a v1 GET with Last-Event-ID must still replay"
);
}
#[tokio::test]
async fn spy_records_zero_replay_for_a_v2_get() {
let (state, spy) = spy_state();
let headers = headers_from(&[
(http::header::ACCEPT.as_str(), TEXT_EVENT_STREAM),
(MCP_PROTOCOL_VERSION, V2),
(LAST_EVENT_ID, "evt-1"),
]);
let response = handle_get_sse(State(state), headers).await.into_response();
assert_eq!(response.status(), StatusCode::METHOD_NOT_ALLOWED);
assert_eq!(spy.replays(), 0, "a v2 GET must never replay");
assert_eq!(spy.stores(), 0);
}
async fn open_v1_sse_stream(state: &ServerState) -> String {
let headers = headers_from(&[(http::header::ACCEPT.as_str(), TEXT_EVENT_STREAM)]);
let response = handle_get_sse(State(state.clone()), headers)
.await
.into_response();
assert_eq!(
response.status(),
StatusCode::OK,
"v1 GET opens an SSE stream"
);
response
.headers()
.get(MCP_SESSION_ID)
.and_then(|v| v.to_str().ok())
.map(str::to_string)
.expect("a v1 SSE GET mints and echoes a session id")
}
#[tokio::test]
async fn v2_response_is_never_routed_into_a_session_sse_stream() {
let state = dual_era_state();
let victim_session = open_v1_sse_stream(&state).await;
let response = handle_post_fast_path(
state.clone(),
post_request(
&v2_post_headers("tools/list", &[(MCP_SESSION_ID, victim_session.as_str())]),
&String::from_utf8(v2_body_bytes("tools/list", "_meta")).unwrap(),
),
)
.await;
assert_ne!(
response.status(),
StatusCode::ACCEPTED,
"a v2 response must NEVER be handed to a session SSE stream — \
202 Accepted means it went to the v1 caller instead of this one"
);
assert_eq!(
response.status(),
StatusCode::OK,
"the v2 caller must get its OWN response back"
);
}
}
#[test]
fn envelope_for_live_request_restamps_a_cached_payload() {
use crate::types::jsonrpc::ResponsePayload;
let cached = crate::types::JSONRPCResponse::success(
crate::types::RequestId::Number(1),
serde_json::json!({ "cached": true }),
);
let live = envelope_for_live_request(
cached.payload.clone(),
crate::types::RequestId::String("caller-2".to_string()),
);
assert_eq!(live.id, crate::types::RequestId::String("caller-2".into()));
assert_eq!(live.jsonrpc, "2.0");
match (&cached.payload, &live.payload) {
(ResponsePayload::Result(before), ResponsePayload::Result(after)) => {
assert_eq!(before, after, "the PAYLOAD survives verbatim");
},
_ => panic!("the result arm must stay a result"),
}
let cached_error = crate::types::JSONRPCResponse::error(
crate::types::RequestId::Number(1),
crate::types::JSONRPCError::new(
crate::types::protocol::error_codes::METHOD_NOT_FOUND,
"nope",
),
);
let live_error =
envelope_for_live_request(cached_error.payload, crate::types::RequestId::Number(99));
assert_eq!(live_error.id, crate::types::RequestId::Number(99));
let ResponsePayload::Error(error) = live_error.payload else {
panic!("the error arm must stay an error");
};
assert_eq!(
error.code,
crate::types::protocol::error_codes::METHOD_NOT_FOUND
);
}
proptest::proptest! {
#[test]
fn envelope_for_live_request_always_carries_the_supplied_id(
numeric in proptest::prelude::any::<bool>(),
number in proptest::prelude::any::<i64>(),
text in "[a-zA-Z0-9-]{0,32}",
is_error in proptest::prelude::any::<bool>(),
) {
let live_id = if numeric {
crate::types::RequestId::Number(number)
} else {
crate::types::RequestId::String(text)
};
let payload = if is_error {
crate::types::jsonrpc::ResponsePayload::Error(
crate::types::JSONRPCError::new(-1, "e"),
)
} else {
crate::types::jsonrpc::ResponsePayload::Result(serde_json::json!({ "k": "v" }))
};
let response = envelope_for_live_request(payload, live_id.clone());
proptest::prop_assert_eq!(response.id, live_id);
}
}
proptest::proptest! {
#[test]
fn classify_http_ingress_never_panics(
raw in proptest::collection::vec(proptest::num::u8::ANY, 0..512),
method in "[a-z/]{0,24}",
oversized in proptest::bool::ANY,
) {
let _ = classify_http_ingress(&raw);
let meta_val = if oversized { "x".repeat(20_000) } else { "2026-07-28".to_string() };
let body = serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": method,
"params": { "_meta": { "io.modelcontextprotocol/protocolVersion": meta_val } }
});
let bytes = serde_json::to_vec(&body).unwrap();
let classified = classify_http_ingress(&bytes);
if method != "server/discover" {
proptest::prop_assert!(
!matches!(classified, Some(HttpIngress::Discover { .. })),
"non-discover method {} must never classify as Discover",
method
);
}
}
}
mod subscriptions_listen {
use super::*;
use crate::types::capabilities::{
PromptCapabilities, ResourceCapabilities, ToolCapabilities,
};
use crate::types::subscriptions::{
advertises_subscriptions, SubscriptionFilter, ACKNOWLEDGED_METHOD,
SUBSCRIPTIONS_LISTEN_METHOD, SUBSCRIPTION_ID_META_KEY,
};
use crate::types::{Implementation, RequestId, ServerCapabilities};
fn only(which: Option<&str>) -> ServerCapabilities {
let mut caps = ServerCapabilities::default();
match which {
Some("tools.listChanged") => {
caps.tools = Some(ToolCapabilities {
list_changed: Some(true),
});
},
Some("prompts.listChanged") => {
caps.prompts = Some(PromptCapabilities {
list_changed: Some(true),
});
},
Some("resources.listChanged") => {
caps.resources = Some(ResourceCapabilities {
subscribe: None,
list_changed: Some(true),
});
},
Some("resources.subscribe") => {
caps.resources = Some(ResourceCapabilities {
subscribe: Some(true),
list_changed: None,
});
},
_ => {},
}
caps
}
fn projected_capabilities(caps: &ServerCapabilities) -> ServerCapabilities {
let response = crate::server::core::build_discover_response(
RequestId::Number(1),
caps,
&Implementation::new("s", "1"),
Some(&v2_context()),
);
let crate::types::jsonrpc::ResponsePayload::Result(value) = response.payload else {
panic!("a v2 discover projects a result");
};
serde_json::from_value(value["capabilities"].clone())
.expect("the projection deserializes back into ServerCapabilities")
}
#[test]
fn discover_projection_and_listen_gate_read_the_same_predicate() {
for which in [
None,
Some("tools.listChanged"),
Some("prompts.listChanged"),
Some("resources.listChanged"),
Some("resources.subscribe"),
] {
let caps = only(which);
let expected = which.is_some();
assert_eq!(
advertises_subscriptions(&caps),
expected,
"gate verdict for {which:?}"
);
assert_eq!(
advertises_subscriptions(&projected_capabilities(&caps)),
expected,
"the discover projection must agree with the gate for {which:?}"
);
}
}
#[test]
fn classify_http_ingress_routes_subscriptions_listen() {
let body = serde_json::to_vec(&json!({
"jsonrpc": "2.0",
"id": 7,
"method": SUBSCRIPTIONS_LISTEN_METHOD,
"params": { "notifications": { "toolsListChanged": true } },
}))
.unwrap();
let Some(HttpIngress::SubscriptionsListen { id, params }) =
classify_http_ingress(&body)
else {
panic!("subscriptions/listen classifies as its own ingress");
};
assert_eq!(id, RequestId::Number(7), "the ORIGINAL id is preserved");
assert_eq!(
params.expect("params carried through")["notifications"]["toolsListChanged"],
json!(true)
);
}
#[test]
fn classify_http_ingress_leaves_other_methods_alone() {
for method in ["tools/call", "resources/subscribe", "initialize"] {
let body = serde_json::to_vec(&json!({
"jsonrpc": "2.0", "id": 1, "method": method, "params": {},
}))
.unwrap();
assert!(
!matches!(
classify_http_ingress(&body),
Some(HttpIngress::SubscriptionsListen { .. })
),
"{method} must not classify as a listen ingress"
);
}
}
#[test]
fn only_the_two_retired_resource_rpcs_are_retired() {
let subscribe = Request::Client(Box::new(ClientRequest::Subscribe(
crate::types::resources::SubscribeRequest {
uri: "mem://a".to_string(),
},
)));
let unsubscribe = Request::Client(Box::new(ClientRequest::Unsubscribe(
crate::types::resources::UnsubscribeRequest {
uri: "mem://a".to_string(),
},
)));
let list = Request::Client(Box::new(ClientRequest::ListTools(
crate::types::tools::ListToolsRequest { cursor: None },
)));
assert_eq!(
v2_retired_method_of(&subscribe),
Some("resources/subscribe")
);
assert_eq!(
v2_retired_method_of(&unsubscribe),
Some("resources/unsubscribe")
);
assert_eq!(
v2_retired_method_of(&list),
None,
"no other method is retired by HTTP-04"
);
}
#[test]
fn the_ack_frame_is_the_acknowledged_notification() {
let agreed = SubscriptionFilter {
tools_list_changed: Some(true),
..SubscriptionFilter::default()
};
let frame: serde_json::Value =
serde_json::from_str(&listen_ack_frame(&agreed, &RequestId::Number(1)))
.expect("the ack frame is JSON");
assert_eq!(frame["jsonrpc"], json!("2.0"));
assert_eq!(frame["method"], json!(ACKNOWLEDGED_METHOD));
assert!(
frame.get("id").is_none(),
"the acknowledgement is a NOTIFICATION, so it carries no id"
);
assert_eq!(
frame["params"]["notifications"],
json!({ "toolsListChanged": true })
);
assert_eq!(
frame["params"]["_meta"][SUBSCRIPTION_ID_META_KEY],
json!(1),
"the subscriptionId equals the listen request's JSON-RPC id"
);
}
#[test]
fn the_terminal_result_goes_through_the_shared_v2_envelope() {
let info = Implementation::new("listen-server", "9.9");
let frame: serde_json::Value = serde_json::from_str(&listen_terminal_result_frame(
&RequestId::Number(3),
Some(&v2_context()),
&info,
))
.expect("the terminal frame is JSON");
assert_eq!(frame["id"], json!(3), "the response id is the listen id");
assert_eq!(
frame["result"]["_meta"][SUBSCRIPTION_ID_META_KEY],
json!(3),
"SubscriptionsListenResult._meta carries the REQUIRED subscriptionId"
);
assert_eq!(
frame["result"]["resultType"],
json!("complete"),
"resultType comes from the SHARED envelope helper, not a bespoke builder"
);
assert_eq!(
frame["result"]["_meta"][crate::server::core::RESERVED_SERVER_INFO_KEY]["name"],
json!("listen-server"),
"serverInfo comes from the SHARED envelope helper too"
);
}
}
}