use std::net::IpAddr;
use super::*;
type AssociatedCall = Pin<Box<dyn Future<Output = Result<JsonRpcResponse>> + Send + 'static>>;
pub(super) fn is_localhost_origin(origin: &str) -> bool {
strip_scheme_ci(origin, "http://")
.or_else(|| strip_scheme_ci(origin, "https://"))
.is_some_and(is_localhost_host)
}
fn strip_scheme_ci<'a>(s: &'a str, scheme: &str) -> Option<&'a str> {
let prefix = scheme.as_bytes();
if s.len() >= prefix.len() && s.as_bytes()[..prefix.len()].eq_ignore_ascii_case(prefix) {
Some(&s[prefix.len()..])
} else {
None
}
}
pub(super) fn is_localhost_host(host: &str) -> bool {
if let Ok(ip) = host.parse::<IpAddr>() {
return ip.is_loopback();
}
let host_only = if host.starts_with('[') {
let Some(close) = host.find(']') else {
return false;
};
let after_bracket = &host[close + 1..];
let port_ok = after_bracket.is_empty()
|| after_bracket
.strip_prefix(':')
.is_some_and(|port| !port.is_empty() && port.parse::<u16>().is_ok());
if !port_ok {
return false;
}
&host[1..close]
} else {
host.split(':').next().unwrap_or(host)
};
if host_only.eq_ignore_ascii_case("localhost") || host_only.eq_ignore_ascii_case("localhost.") {
return true;
}
host_only.parse::<IpAddr>().is_ok_and(|ip| ip.is_loopback())
}
pub(super) fn effective_host<'a>(
headers: &'a HeaderMap,
uri: &'a axum::http::Uri,
) -> Option<&'a str> {
if let Some(value) = headers.get(header::HOST)
&& let Ok(s) = value.to_str()
{
return Some(s);
}
uri.authority().map(|a| a.as_str())
}
pub(super) fn validate_host(
headers: &HeaderMap,
uri: &axum::http::Uri,
state: &AppState,
) -> Option<Response> {
if !state.validate_host {
return None;
}
let Some(host) = effective_host(headers, uri) else {
if state.allowed_hosts.is_empty() {
return None;
}
tracing::warn!("Rejecting request: missing Host header and no :authority fallback");
return Some((StatusCode::BAD_REQUEST, "Missing Host header").into_response());
};
if is_localhost_host(host) {
return None;
}
if state.allowed_hosts.is_empty() {
return None;
}
if state.allowed_hosts.iter().any(|h| h == host) {
return None;
}
tracing::warn!(host = %host, "Rejecting request: Host not in allowlist");
Some((StatusCode::BAD_REQUEST, "Host not allowed").into_response())
}
fn validate_origin(headers: &HeaderMap, state: &AppState) -> Option<Response> {
if !state.validate_origin {
return None;
}
if let Some(origin) = headers.get(header::ORIGIN) {
let origin_str = origin.to_str().unwrap_or("");
if is_localhost_origin(origin_str) {
return None;
}
if state.allowed_origins.is_empty() {
tracing::warn!(
origin = %origin_str,
"Rejecting request: cross-origin not allowed (no allowlist configured)"
);
return Some(
(StatusCode::FORBIDDEN, "Cross-origin requests not allowed").into_response(),
);
}
if !state
.allowed_origins
.iter()
.any(|o| o == origin_str || o == "*")
{
tracing::warn!(origin = %origin_str, "Rejecting request: Origin not in allowlist");
return Some((StatusCode::FORBIDDEN, "Origin not allowed").into_response());
}
}
None
}
fn get_session_id(headers: &HeaderMap) -> Option<String> {
headers
.get(MCP_SESSION_ID_HEADER)
.and_then(|v| v.to_str().ok())
.map(|s| s.to_string())
}
fn get_protocol_version(headers: &HeaderMap) -> Option<String> {
headers
.get(MCP_PROTOCOL_VERSION_HEADER)
.and_then(|v| v.to_str().ok())
.map(|s| s.to_string())
}
fn get_last_event_id(headers: &HeaderMap) -> Option<u64> {
headers
.get(LAST_EVENT_ID_HEADER)
.and_then(|v| v.to_str().ok())
.and_then(|s| s.parse::<u64>().ok())
}
fn is_initialize_request(body: &serde_json::Value) -> bool {
body.get("method")
.and_then(|m| m.as_str())
.map(|m| m == "initialize")
.unwrap_or(false)
}
fn is_response(parsed: &serde_json::Value) -> bool {
crate::framing::is_response_frame(parsed)
}
fn request_tool_input_schema(
service_source: &ServiceSource,
parsed: &serde_json::Value,
) -> Option<serde_json::Value> {
if parsed.get("method").and_then(serde_json::Value::as_str) != Some("tools/call") {
return None;
}
let name = parsed
.get("params")
.and_then(serde_json::Value::as_object)
.and_then(|params| params.get("name"))
.and_then(serde_json::Value::as_str)?;
match service_source {
ServiceSource::Router { router, .. } => router.tool_input_schema(name),
ServiceSource::Service(_) => None,
}
}
fn claims_modern_protocol(headers: &HeaderMap, parsed: &serde_json::Value) -> bool {
get_protocol_version(headers).as_deref() == Some(PROTOCOL_VERSION_2026_07_28)
|| parsed
.get("params")
.and_then(serde_json::Value::as_object)
.and_then(|params| params.get("_meta"))
.and_then(serde_json::Value::as_object)
.is_some_and(|meta| meta.contains_key("io.modelcontextprotocol/protocolVersion"))
}
fn validate_modern_request_meta(
parsed: &serde_json::Value,
) -> std::result::Result<String, JsonRpcError> {
let params = parsed
.get("params")
.and_then(serde_json::Value::as_object)
.ok_or_else(|| {
JsonRpcError::invalid_params("Modern requests require a params object containing _meta")
})?;
let meta_value = params
.get("_meta")
.ok_or_else(|| JsonRpcError::invalid_params("Modern requests require a _meta object"))?;
crate::protocol::validate_meta_object(meta_value)
.map_err(|error| JsonRpcError::invalid_params(error.to_string()))?;
let meta = meta_value
.as_object()
.expect("validate_meta_object accepted a JSON object");
let protocol_version = meta
.get("io.modelcontextprotocol/protocolVersion")
.and_then(serde_json::Value::as_str)
.ok_or_else(|| {
JsonRpcError::invalid_params(
"Missing or invalid _meta.io.modelcontextprotocol/protocolVersion",
)
})?;
let client_capabilities = meta
.get("io.modelcontextprotocol/clientCapabilities")
.ok_or_else(|| {
JsonRpcError::invalid_params("Missing _meta.io.modelcontextprotocol/clientCapabilities")
})?;
if !client_capabilities.is_object()
|| serde_json::from_value::<ClientCapabilities>(client_capabilities.clone()).is_err()
{
return Err(JsonRpcError::invalid_params(
"Invalid _meta.io.modelcontextprotocol/clientCapabilities",
));
}
Ok(protocol_version.to_string())
}
fn is_removed_modern_method(method: &str) -> bool {
matches!(
method,
"initialize"
| "notifications/initialized"
| "ping"
| "logging/setLevel"
| "resources/subscribe"
| "resources/unsubscribe"
| "notifications/roots/list_changed"
)
}
pub(super) fn extract_request_id(parsed: &serde_json::Value) -> Option<RequestId> {
parsed.get("id").and_then(|id| {
if let Some(n) = id.as_i64() {
Some(RequestId::Number(n))
} else {
id.as_str().map(|s| RequestId::String(s.to_string()))
}
})
}
pub(super) async fn handle_post(
State(state): State<Arc<AppState>>,
request: axum::extract::Request,
) -> Response {
let (parts, body_bytes) = request.into_parts();
let headers = parts.headers;
let uri = parts.uri.clone();
if let Some(resp) = validate_host(&headers, &uri, &state) {
return resp;
}
if let Some(resp) = validate_origin(&headers, &state) {
return resp;
}
if let Some(declared) = headers
.get(header::CONTENT_LENGTH)
.and_then(|v| v.to_str().ok())
.and_then(|v| v.parse::<usize>().ok())
&& declared > state.max_body_size
{
return body_too_large_response(state.max_body_size);
}
let body = match axum::body::to_bytes(body_bytes, state.max_body_size).await {
Ok(bytes) => match String::from_utf8(bytes.to_vec()) {
Ok(s) => s,
Err(e) => {
return json_rpc_error_response(
None,
JsonRpcError::parse_error(format!("Invalid UTF-8: {}", e)),
);
}
},
Err(e) if is_length_limit_error(&e) => {
return body_too_large_response(state.max_body_size);
}
Err(e) => {
return json_rpc_error_response(
None,
JsonRpcError::parse_error(format!("Failed to read body: {}", e)),
);
}
};
let http_extensions = parts.extensions;
let parsed: serde_json::Value =
match serde_json::from_str(crate::framing::clean_input_line(&body)) {
Ok(v) => v,
Err(e) => {
return json_rpc_error_response(
None,
JsonRpcError::parse_error(format!("Invalid JSON: {}", e)),
);
}
};
if parsed.is_array()
&& let Some(version) = get_protocol_version(&headers)
{
let revision = match version.parse::<McpProtocolRevision>() {
Ok(revision) => revision,
Err(_) => {
return json_rpc_error_response(
None,
JsonRpcError::unsupported_protocol_version(
version,
state.protocol_support.versions().iter().map(String::as_str),
),
);
}
};
if let Err(error) = inspect_runtime_value(
&parsed,
revision,
&state.protocol_support,
McpDirection::ClientToServer,
) {
let status = if revision == McpProtocolRevision::V2026_07_28 {
StatusCode::BAD_REQUEST
} else {
StatusCode::OK
};
return json_rpc_error_response_with_status(None, error, status);
}
}
let is_init = is_initialize_request(&parsed);
let request_method = parsed
.get("method")
.and_then(|method| method.as_str())
.unwrap_or_default()
.to_string();
let tool_input_schema = request_tool_input_schema(&state.service_source, &parsed);
let modern_request = claims_modern_protocol(&headers, &parsed);
if modern_request {
let id = extract_request_id(&parsed);
let body_version = match validate_modern_request_meta(&parsed) {
Ok(version) => version,
Err(error) => {
return json_rpc_error_response_with_status(id, error, StatusCode::BAD_REQUEST);
}
};
let Some(header_version) = get_protocol_version(&headers) else {
return json_rpc_error_response_with_status(
id,
JsonRpcError::header_mismatch("MCP-Protocol-Version header is required"),
StatusCode::BAD_REQUEST,
);
};
if header_version != body_version {
return json_rpc_error_response_with_status(
id,
JsonRpcError::header_mismatch(format!(
"MCP-Protocol-Version header value {header_version:?} does not match \
request _meta protocol version {body_version:?}"
)),
StatusCode::BAD_REQUEST,
);
}
if !state.protocol_support.contains(&body_version) {
return json_rpc_error_response_with_status(
id,
JsonRpcError::unsupported_protocol_version(
body_version,
state.protocol_support.versions().iter().map(String::as_str),
),
StatusCode::BAD_REQUEST,
);
}
let revision = match body_version.parse::<McpProtocolRevision>() {
Ok(revision) => revision,
Err(_) => {
return json_rpc_error_response_with_status(
id,
JsonRpcError::unsupported_protocol_version(
body_version,
state.protocol_support.versions().iter().map(String::as_str),
),
StatusCode::BAD_REQUEST,
);
}
};
if let Err(error) = inspect_runtime_value(
&parsed,
revision,
&state.protocol_support,
McpDirection::ClientToServer,
) {
return json_rpc_error_response_with_status(id, error, StatusCode::BAD_REQUEST);
}
let sep_2243_mode = crate::transport::http_headers::mode_for_version(&body_version);
if let Err(error) = crate::transport::http_headers::validate_with_tool_schema(
&headers,
&parsed,
sep_2243_mode,
tool_input_schema.as_ref(),
) {
tracing::warn!(
mode = ?sep_2243_mode,
version = %body_version,
error = %error.message,
"Rejecting modern request: HTTP header validation failed",
);
return json_rpc_error_response_with_status(id, error, StatusCode::BAD_REQUEST);
}
if is_removed_modern_method(&request_method) {
return json_rpc_error_response_with_status(
id,
JsonRpcError::method_not_found(&request_method),
StatusCode::NOT_FOUND,
);
}
}
#[cfg(feature = "stateless")]
{
let version_in_play: Option<String> = if is_init && !modern_request {
parsed
.get("params")
.and_then(|p| p.get("protocolVersion"))
.and_then(|v| v.as_str())
.map(|s| s.to_string())
} else {
get_protocol_version(&headers)
};
if let Some(ref version) = version_in_play
&& is_stateless_protocol_version(version)
&& state.protocol_support.contains(version)
&& parsed.get("method").and_then(|m| m.as_str()) != Some("subscriptions/listen")
{
if !is_init && (parsed.get("id").is_none() || is_response(&parsed)) {
return StatusCode::ACCEPTED.into_response();
}
let sep_2243_mode = crate::transport::http_headers::mode_for_version(version);
if let Err(err) = crate::transport::http_headers::validate_with_tool_schema(
&headers,
&parsed,
sep_2243_mode,
tool_input_schema.as_ref(),
) {
tracing::warn!(
mode = ?sep_2243_mode,
version = %version,
error = %err.message,
"Rejecting stateless request: SEP-2243 header validation failed",
);
let id = extract_request_id(&parsed);
let mut resp = json_rpc_error_response(id, err);
*resp.status_mut() = StatusCode::BAD_REQUEST;
return resp;
}
let id = extract_request_id(&parsed);
let request: JsonRpcRequest = match serde_json::from_value(parsed) {
Ok(r) => r,
Err(_) => {
return json_rpc_error_response(
id,
JsonRpcError::parse_error("not a valid JSON-RPC request"),
);
}
};
let server_identity = match &state.service_source {
ServiceSource::Router { router, .. } if state.stamp_server_info => {
Some(router.implementation())
}
_ => None,
};
let (notif_tx, mut notif_rx) = crate::context::notification_channel(64);
let mut service = match &state.service_source {
ServiceSource::Router { router, factory } => {
let ephemeral = router
.with_fresh_session()
.with_request_notification_sender(notif_tx);
ephemeral.session().mark_preinitialized();
JsonRpcService::new(factory(ephemeral))
}
ServiceSource::Service(mutex) => JsonRpcService::new(mutex.lock().unwrap().clone()),
};
let mut ext = crate::router::Extensions::new();
ext.insert(state.protocol_support.clone());
#[cfg(feature = "oauth")]
if let Some(claims) = http_extensions.get::<crate::oauth::token::TokenClaims>() {
ext.insert(claims.clone());
}
stash_per_request_meta(&request, &mut ext);
crate::transport::extension_bridge::apply_extension_bridges(
&state.extension_bridges,
&http_extensions,
&mut ext,
);
let cancel_token = crate::context::CancellationToken::new();
let mut cancel_guard = CancelOnDisconnect::arm(cancel_token.clone());
ext.insert(cancel_token);
service = service.with_extensions(ext);
let mut call: std::pin::Pin<
Box<dyn std::future::Future<Output = crate::error::Result<JsonRpcResponse>> + Send>,
> = Box::pin(async move {
let mut service = service;
service.call_single(request).await
});
enum FirstOutbound {
Response(crate::error::Result<JsonRpcResponse>),
Notification(crate::context::ServerNotification),
}
let first = loop {
let outbound = tokio::select! {
biased;
maybe = notif_rx.recv() => match maybe {
Some(n) => FirstOutbound::Notification(n),
None => FirstOutbound::Response((&mut call).await),
},
result = &mut call => FirstOutbound::Response(result),
};
match outbound {
FirstOutbound::Notification(notification)
if state.modern_subscriptions.publish(¬ification) =>
{
continue;
}
outbound => break outbound,
}
};
match first {
FirstOutbound::Response(result) => {
while let Ok(notification) = notif_rx.try_recv() {
if state.modern_subscriptions.publish(¬ification) {
continue;
}
let ready_call: std::pin::Pin<
Box<
dyn std::future::Future<
Output = crate::error::Result<JsonRpcResponse>,
> + Send,
>,
> = Box::pin(async move { result });
let mut resp = stateless_sse_with_notifications(
notification,
ready_call,
notif_rx,
StatelessSseContext {
version: version.clone(),
method: request_method.clone(),
cancel_guard,
server_identity,
subscriptions: state.modern_subscriptions.clone(),
},
);
resp.headers_mut().insert(
MCP_PROTOCOL_VERSION_HEADER,
HeaderValue::from_str(version).unwrap(),
);
return resp;
}
cancel_guard.disarm();
let mut response = match result {
Ok(resp) => resp,
Err(e) => {
return json_rpc_error_response(
None,
JsonRpcError::internal_error(e.to_string()),
);
}
};
if is_init
&& let JsonRpcResponse::Result(ref mut result) = response
&& let Some(pv) = result.result.get_mut("protocolVersion")
{
*pv = serde_json::Value::String(version.clone());
}
apply_protocol_result_fields(&mut response, &request_method, version);
if let Some(ref identity) = server_identity {
stamp_server_info(&mut response, identity);
}
let status = modern_response_status(&response);
let mut resp = if state.sse_responses {
sse_json_response(&response)
} else {
axum::Json(response).into_response()
};
*resp.status_mut() = status;
resp.headers_mut().insert(
MCP_PROTOCOL_VERSION_HEADER,
HeaderValue::from_str(version).unwrap(),
);
return resp;
}
FirstOutbound::Notification(first_notif) => {
let mut resp = stateless_sse_with_notifications(
first_notif,
call,
notif_rx,
StatelessSseContext {
version: version.clone(),
method: request_method.clone(),
cancel_guard,
server_identity,
subscriptions: state.modern_subscriptions.clone(),
},
);
resp.headers_mut().insert(
MCP_PROTOCOL_VERSION_HEADER,
HeaderValue::from_str(version).unwrap(),
);
return resp;
}
}
}
}
#[cfg(feature = "stateless")]
if !is_init && state.stateless_config.is_some() && get_session_id(&headers).is_none() {
let version_from_header = get_protocol_version(&headers);
let params = parsed.get("params").unwrap_or(&parsed);
let version_from_meta = crate::stateless::StatelessRequestMeta::from_params(params)
.and_then(|m| m.protocol_version);
if let Some(version) = version_from_header.or(version_from_meta) {
if let Err(err) = crate::stateless::validate_protocol_version(&version) {
return json_rpc_error_response(None, err);
}
if parsed.get("id").is_none() || is_response(&parsed) {
return StatusCode::ACCEPTED.into_response();
}
let id = extract_request_id(&parsed);
let request: JsonRpcRequest = match serde_json::from_value(parsed) {
Ok(r) => r,
Err(_) => {
return json_rpc_error_response(
id,
JsonRpcError::parse_error("not a valid JSON-RPC request"),
);
}
};
let mut service = match &state.service_source {
ServiceSource::Router { router, factory } => {
let ephemeral = router.with_fresh_session();
ephemeral.session().mark_preinitialized();
JsonRpcService::new(factory(ephemeral))
}
ServiceSource::Service(mutex) => JsonRpcService::new(mutex.lock().unwrap().clone()),
};
let mut ext = crate::router::Extensions::new();
ext.insert(state.protocol_support.clone());
#[cfg(feature = "oauth")]
if let Some(claims) = http_extensions.get::<crate::oauth::token::TokenClaims>() {
ext.insert(claims.clone());
}
#[cfg(feature = "stateless")]
stash_per_request_meta(&request, &mut ext);
crate::transport::extension_bridge::apply_extension_bridges(
&state.extension_bridges,
&http_extensions,
&mut ext,
);
if !ext.is_empty() {
service = service.with_extensions(ext);
}
let mut response = match service.call_single(request).await {
Ok(resp) => resp,
Err(e) => {
return json_rpc_error_response(
None,
JsonRpcError::internal_error(e.to_string()),
);
}
};
apply_protocol_result_fields(&mut response, &request_method, &version);
let mut resp = if state.sse_responses {
sse_json_response(&response)
} else {
axum::Json(response).into_response()
};
resp.headers_mut().insert(
MCP_PROTOCOL_VERSION_HEADER,
HeaderValue::from_str(&version).unwrap(),
);
return resp;
}
}
#[cfg(feature = "stateless")]
if modern_request && request_method == "subscriptions/listen" {
return handle_modern_subscriptions_listen_sse(state, &parsed, &http_extensions).await;
}
if !is_init
&& let Some(version) = get_protocol_version(&headers)
&& !state.protocol_support.contains(&version)
{
return json_rpc_error_response(
extract_request_id(&parsed),
JsonRpcError::unsupported_protocol_version(
version,
state.protocol_support.versions().iter().map(String::as_str),
),
);
}
let uses_transient_session = !is_init
&& !modern_request
&& get_session_id(&headers).is_none()
&& state.optional_sessions;
let session = if is_init {
let create_result = match &state.service_source {
ServiceSource::Router { router, factory } => {
state
.sessions
.create(router.with_fresh_session(), factory.clone())
.await
}
ServiceSource::Service(mutex) => {
let service = mutex.lock().unwrap().clone();
state.sessions.create_from_service(service).await
}
};
match create_result {
Some(s) => s,
None => {
return (
StatusCode::SERVICE_UNAVAILABLE,
"Maximum session limit reached",
)
.into_response();
}
}
} else if !modern_request && let Some(session_id) = get_session_id(&headers) {
match state.sessions.get(&session_id).await {
Some(s) => s,
None => {
return json_rpc_error_response(
None,
JsonRpcError::session_not_found_with_id(&session_id),
);
}
}
} else if state.optional_sessions {
let create_result = match &state.service_source {
ServiceSource::Router { router, factory } => {
state
.sessions
.create_initialized(router.with_fresh_session(), factory.clone())
.await
}
ServiceSource::Service(mutex) => {
let service = mutex.lock().unwrap().clone();
state
.sessions
.create_initialized_from_service(service)
.await
}
};
match create_result {
Some(s) => s,
None => {
return (
StatusCode::SERVICE_UNAVAILABLE,
"Maximum session limit reached",
)
.into_response();
}
}
} else {
return json_rpc_error_response(None, JsonRpcError::session_required());
};
let session_protocol_version = if uses_transient_session {
let version = state
.protocol_support
.versions()
.iter()
.find(|version| {
crate::protocol::SUPPORTED_PROTOCOL_VERSIONS.contains(&version.as_str())
})
.map_or_else(
|| state.protocol_support.preferred().to_string(),
Clone::clone,
);
*session.protocol_version.write().await = version.clone();
version
} else {
session.protocol_version.read().await.clone()
};
let session_revision = match session_protocol_version.parse::<McpProtocolRevision>() {
Ok(revision) => revision,
Err(_) => {
return json_rpc_error_response(
extract_request_id(&parsed),
JsonRpcError::unsupported_protocol_version(
session_protocol_version,
state.protocol_support.versions().iter().map(String::as_str),
),
);
}
};
if !is_init
&& let Err(error) = inspect_runtime_value(
&parsed,
session_revision,
&state.protocol_support,
McpDirection::ClientToServer,
)
{
return json_rpc_error_response(extract_request_id(&parsed), error);
}
if parsed.is_array() {
if state.strict_initialization
&& !session
.initialized_notification_received
.load(Ordering::Acquire)
{
return json_rpc_error_response(
None,
JsonRpcError::invalid_request(
"Client must send notifications/initialized before making requests",
),
);
}
let message: JsonRpcMessage = match serde_json::from_value(parsed) {
Ok(message) => message,
Err(error) => {
return json_rpc_error_response(
None,
JsonRpcError::invalid_request(format!("Invalid request batch: {error}")),
);
}
};
let mut extensions = crate::router::Extensions::new();
extensions.insert(state.protocol_support.clone());
extensions.insert(session_revision);
#[cfg(feature = "oauth")]
if let Some(claims) = http_extensions.get::<crate::oauth::token::TokenClaims>() {
extensions.insert(claims.clone());
}
crate::transport::extension_bridge::apply_extension_bridges(
&state.extension_bridges,
&http_extensions,
&mut extensions,
);
let mut service = JsonRpcService::new(session.make_service())
.with_extensions(extensions)
.protocol_support(state.protocol_support.clone())
.with_negotiated_protocol_version(&session_protocol_version);
let response = match service.call_message(message).await {
Ok(response) => response,
Err(error) => {
return json_rpc_error_response(
None,
JsonRpcError::internal_error(error.to_string()),
);
}
};
let mut response = axum::Json(response).into_response();
response.headers_mut().insert(
MCP_PROTOCOL_VERSION_HEADER,
HeaderValue::from_str(&session_protocol_version).unwrap(),
);
return response;
}
{
let method_str = parsed.get("method").and_then(|m| m.as_str()).unwrap_or("");
if method_str == "subscriptions/listen" {
let req_id = extract_request_id(&parsed);
let effective_version = if let Some(v) = get_protocol_version(&headers) {
v
} else {
session.protocol_version.read().await.clone()
};
if version_supports_subscriptions_listen(&effective_version, &state.protocol_support) {
return handle_subscriptions_listen_sse(session).await;
} else {
return json_rpc_error_response(
req_id,
JsonRpcError::method_not_found("subscriptions/listen"),
);
}
}
}
let sep_2243_version = if is_init {
match parsed
.get("params")
.and_then(|p| p.get("protocolVersion"))
.and_then(|v| v.as_str())
{
Some(v) => v.to_string(),
None => session.protocol_version.read().await.clone(),
}
} else {
session.protocol_version.read().await.clone()
};
let sep_2243_mode = crate::transport::http_headers::mode_for_version(&sep_2243_version);
if let Err(err) = crate::transport::http_headers::validate_with_tool_schema(
&headers,
&parsed,
sep_2243_mode,
tool_input_schema.as_ref(),
) {
tracing::warn!(
mode = ?sep_2243_mode,
version = %sep_2243_version,
error = %err.message,
"Rejecting request: SEP-2243 header validation failed",
);
let id = extract_request_id(&parsed);
let mut resp = json_rpc_error_response(id, err);
*resp.status_mut() = StatusCode::BAD_REQUEST;
return resp;
}
if is_response(&parsed) {
if let Some(id) = extract_request_id(&parsed) {
let result = if let Some(error) = parsed.get("error") {
let code = error.get("code").and_then(|c| c.as_i64()).unwrap_or(-1);
let message = error
.get("message")
.and_then(|m| m.as_str())
.unwrap_or("Unknown error");
Err(Error::Internal(format!(
"Client error ({}): {}",
code, message
)))
} else if let Some(result) = parsed.get("result") {
Ok(result.clone())
} else {
Err(Error::Internal(
"Response has neither result nor error".to_string(),
))
};
if session.complete_pending_request(&id, result).await {
tracing::debug!(request_id = ?id, "Completed pending request");
} else {
tracing::warn!(request_id = ?id, "Received response for unknown request");
}
}
return StatusCode::ACCEPTED.into_response();
}
if parsed.get("id").is_none() {
if let Ok(notification) = serde_json::from_value::<JsonRpcNotification>(parsed)
&& let Ok(mcp_notification) = McpNotification::from_jsonrpc(¬ification)
{
if matches!(&mcp_notification, McpNotification::Initialized) {
session
.initialized_notification_received
.store(true, Ordering::Release);
tracing::debug!(session_id = %session.id, "Received notifications/initialized");
}
session.handle_notification(mcp_notification);
}
return StatusCode::ACCEPTED.into_response();
}
if !is_init
&& state.strict_initialization
&& !session
.initialized_notification_received
.load(Ordering::Acquire)
{
let id = extract_request_id(&parsed);
tracing::warn!(
session_id = %session.id,
"Rejecting request: notifications/initialized not yet received"
);
return json_rpc_error_response(
id,
JsonRpcError::invalid_request(
"Client must send notifications/initialized before making requests",
),
);
}
let init_client_metadata: Option<(Option<Implementation>, Option<ClientCapabilities>)> =
if is_init {
let params = parsed.get("params");
let client_info = params
.and_then(|p| p.get("clientInfo"))
.and_then(|v| serde_json::from_value::<Implementation>(v.clone()).ok());
let client_capabilities = params
.and_then(|p| p.get("capabilities"))
.and_then(|v| serde_json::from_value::<ClientCapabilities>(v.clone()).ok());
Some((client_info, client_capabilities))
} else {
None
};
let id = extract_request_id(&parsed);
let request: JsonRpcRequest = match serde_json::from_value(parsed) {
Ok(r) => r,
Err(_) => {
return json_rpc_error_response(
id,
JsonRpcError::parse_error("not a valid JSON-RPC request"),
);
}
};
let mut service = JsonRpcService::new(session.make_service());
#[allow(unused_mut)]
let mut ext = crate::router::Extensions::new();
ext.insert(state.protocol_support.clone());
ext.insert(session_revision);
#[cfg(feature = "oauth")]
if let Some(claims) = http_extensions.get::<crate::oauth::token::TokenClaims>() {
ext.insert(claims.clone());
}
crate::transport::extension_bridge::apply_extension_bridges(
&state.extension_bridges,
&http_extensions,
&mut ext,
);
#[cfg(feature = "stateless")]
stash_per_request_meta(&request, &mut ext);
let mut associated_request_rx = if !is_init {
session.request_id_allocator.as_ref().map(|next_id| {
let (request_tx, request_rx) = outgoing_request_channel(32);
let requester: ClientRequesterHandle = Arc::new(
ChannelClientRequester::with_id_allocator(request_tx, next_id.clone()),
);
ext.insert(requester);
request_rx
})
} else {
None
};
if !ext.is_empty() {
service = service.with_extensions(ext);
}
let request_id = request.id.clone();
let mut call: AssociatedCall = Box::pin(async move { service.call_single(request).await });
let mut response = if let Some(mut request_rx) = associated_request_rx.take() {
tokio::select! {
result = &mut call => match result {
Ok(response) => response,
Err(error) => {
return json_rpc_error_response(
Some(request_id),
JsonRpcError::internal_error(error.to_string()),
);
}
},
outgoing = request_rx.recv() => {
match outgoing {
Some(outgoing) => {
let negotiated_version = session.protocol_version.read().await.clone();
return associated_request_sse_response(
session,
call,
request_rx,
outgoing,
request_id,
request_method,
negotiated_version,
);
}
None => match call.await {
Ok(response) => response,
Err(error) => {
return json_rpc_error_response(
Some(request_id),
JsonRpcError::internal_error(error.to_string()),
);
}
},
}
}
}
} else {
match call.await {
Ok(response) => response,
Err(error) => {
return json_rpc_error_response(
Some(request_id),
JsonRpcError::internal_error(error.to_string()),
);
}
}
};
if is_init && let JsonRpcResponse::Result(ref result) = response {
if let Some(version) = result
.result
.get("protocolVersion")
.and_then(|v| v.as_str())
{
*session.protocol_version.write().await = version.to_string();
}
if let Some((client_info, client_capabilities)) = init_client_metadata {
*session.client_info.write().await = client_info;
*session.client_capabilities.write().await = client_capabilities;
}
state.sessions.save_record(&session).await;
}
let negotiated_version = session.protocol_version.read().await.clone();
let response_version = if request_method == "server/discover"
&& state.protocol_support.contains(PROTOCOL_VERSION_2026_07_28)
{
PROTOCOL_VERSION_2026_07_28
} else {
&negotiated_version
};
apply_protocol_result_fields(&mut response, &request_method, response_version);
let mut resp = if state.sse_responses {
sse_json_response(&response)
} else {
axum::Json(response).into_response()
};
if is_init {
resp.headers_mut().insert(
MCP_SESSION_ID_HEADER,
HeaderValue::from_str(&session.id).unwrap(),
);
}
resp.headers_mut().insert(
MCP_PROTOCOL_VERSION_HEADER,
HeaderValue::from_str(&negotiated_version).unwrap(),
);
resp
}
fn associated_request_sse_response(
session: Arc<Session>,
mut call: AssociatedCall,
mut request_rx: OutgoingRequestReceiver,
first_outgoing: OutgoingRequest,
original_request_id: RequestId,
request_method: String,
negotiated_version: String,
) -> Response {
let (event_tx, event_rx) =
tokio::sync::mpsc::channel::<std::result::Result<Event, Infallible>>(32);
let call_version = negotiated_version.clone();
tokio::spawn(async move {
let mut pending_ids = Vec::new();
if !send_associated_request(&session, &event_tx, first_outgoing, &mut pending_ids).await {
session
.fail_pending_requests(
&pending_ids,
"originating POST disconnected before the client request was delivered",
)
.await;
return;
}
let mut requests_open = true;
loop {
tokio::select! {
_ = event_tx.closed() => {
session
.fail_pending_requests(
&pending_ids,
"originating POST response stream disconnected",
)
.await;
return;
}
result = &mut call => {
session
.fail_pending_requests(
&pending_ids,
"originating POST completed before the client request response arrived",
)
.await;
let mut response = match result {
Ok(response) => response,
Err(error) => JsonRpcResponse::error(
Some(original_request_id),
JsonRpcError::internal_error(error.to_string()),
),
};
apply_protocol_result_fields(
&mut response,
&request_method,
&call_version,
);
match serde_json::to_string(&response) {
Ok(data) => {
let _ = event_tx
.send(Ok(
Event::default()
.event(SSE_MESSAGE_EVENT)
.data(data),
))
.await;
}
Err(error) => {
tracing::error!(
error = %error,
"Failed to serialize associated POST response",
);
}
}
return;
}
outgoing = request_rx.recv(), if requests_open => {
match outgoing {
Some(outgoing) => {
if !send_associated_request(
&session,
&event_tx,
outgoing,
&mut pending_ids,
)
.await
{
session
.fail_pending_requests(
&pending_ids,
"originating POST disconnected before the client request was delivered",
)
.await;
return;
}
}
None => requests_open = false,
}
}
}
}
});
let stream = tokio_stream::wrappers::ReceiverStream::new(event_rx);
let mut response = Sse::new(stream)
.keep_alive(
axum::response::sse::KeepAlive::new()
.interval(Duration::from_secs(30))
.text("ping"),
)
.into_response();
response.headers_mut().insert(
MCP_PROTOCOL_VERSION_HEADER,
HeaderValue::from_str(&negotiated_version).unwrap(),
);
response
}
async fn send_associated_request(
session: &Session,
event_tx: &tokio::sync::mpsc::Sender<std::result::Result<Event, Infallible>>,
outgoing: OutgoingRequest,
pending_ids: &mut Vec<RequestId>,
) -> bool {
let id = outgoing.id.clone();
let request = JsonRpcRequest {
jsonrpc: "2.0".to_string(),
id: id.clone(),
method: outgoing.method,
params: Some(outgoing.params),
};
let data = match serde_json::to_string(&request) {
Ok(data) => data,
Err(error) => {
let _ = outgoing.response_tx.send(Err(Error::Internal(format!(
"Failed to serialize associated client request: {error}"
))));
return true;
}
};
session
.add_pending_request(id.clone(), outgoing.response_tx)
.await;
pending_ids.push(id);
event_tx
.send(Ok(Event::default().event(SSE_MESSAGE_EVENT).data(data)))
.await
.is_ok()
}
fn version_supports_subscriptions_listen(
version: &str,
protocol_support: &ProtocolSupport,
) -> bool {
version == PROTOCOL_VERSION_2026_07_28 && protocol_support.contains(version)
}
async fn handle_subscriptions_listen_sse(session: Arc<Session>) -> Response {
let rx = session.notifications_tx.subscribe();
let session_clone = session.clone();
let stream = BroadcastStream::new(rx)
.then(move |result: std::result::Result<String, _>| {
let session = session_clone.clone();
async move {
match result {
Ok(msg) => {
let event_id = session.next_event_id();
session.buffer_event(event_id, msg.clone()).await;
Some(Ok::<_, Infallible>(
Event::default()
.id(event_id.to_string())
.event(SSE_MESSAGE_EVENT)
.data(msg),
))
}
Err(_) => None,
}
}
})
.filter_map(|x| x);
Sse::new(stream)
.keep_alive(
axum::response::sse::KeepAlive::new()
.interval(Duration::from_secs(30))
.text("ping"),
)
.into_response()
}
pub(super) async fn handle_get(
State(state): State<Arc<AppState>>,
request: axum::extract::Request,
) -> Response {
let (parts, _body) = request.into_parts();
let headers = parts.headers;
let uri = parts.uri.clone();
if let Some(resp) = validate_host(&headers, &uri, &state) {
return resp;
}
if let Some(resp) = validate_origin(&headers, &state) {
return resp;
}
let accept = headers
.get(header::ACCEPT)
.and_then(|v| v.to_str().ok())
.unwrap_or("");
if !accept.contains("text/event-stream") {
return (
StatusCode::NOT_ACCEPTABLE,
"Accept header must include text/event-stream",
)
.into_response();
}
let session_id = match get_session_id(&headers) {
Some(id) => id,
None => {
return json_rpc_error_response(None, JsonRpcError::session_required());
}
};
let session = match state.sessions.get(&session_id).await {
Some(s) => s,
None => {
return json_rpc_error_response(
None,
JsonRpcError::session_not_found_with_id(&session_id),
);
}
};
let last_event_id = get_last_event_id(&headers);
let rx = session.notifications_tx.subscribe();
let session_clone = session.clone();
let replay_events: Vec<_> = if let Some(after_id) = last_event_id {
let events = session.get_events_after(after_id).await;
tracing::debug!(
after_id = after_id,
replay_count = events.len(),
"Replaying buffered events for stream resumption"
);
events
.into_iter()
.map(|e| {
Ok::<_, Infallible>(
Event::default()
.id(e.id.to_string())
.event(SSE_MESSAGE_EVENT)
.data(e.data),
)
})
.collect()
} else {
Vec::new()
};
let replay_stream = tokio_stream::iter(replay_events);
let live_stream = BroadcastStream::new(rx)
.then(move |result: std::result::Result<String, _>| {
let session = session_clone.clone();
async move {
match result {
Ok(msg) => {
let event_id = session.next_event_id();
session.buffer_event(event_id, msg.clone()).await;
Some(Ok::<_, Infallible>(
Event::default()
.id(event_id.to_string())
.event(SSE_MESSAGE_EVENT)
.data(msg),
))
}
Err(_) => None,
}
}
})
.filter_map(|x| x);
let stream = replay_stream.chain(live_stream);
Sse::new(stream)
.keep_alive(
axum::response::sse::KeepAlive::new()
.interval(Duration::from_secs(30))
.text("ping"),
)
.into_response()
}
pub(super) async fn handle_delete(
State(state): State<Arc<AppState>>,
request: axum::extract::Request,
) -> Response {
let (parts, _body) = request.into_parts();
let headers = parts.headers;
let uri = parts.uri.clone();
if let Some(resp) = validate_host(&headers, &uri, &state) {
return resp;
}
if let Some(resp) = validate_origin(&headers, &state) {
return resp;
}
let session_id = match get_session_id(&headers) {
Some(id) => id,
None => {
return json_rpc_error_response(None, JsonRpcError::session_required());
}
};
if state.sessions.remove(&session_id).await {
tracing::info!(session_id = %session_id, "Session terminated");
StatusCode::OK.into_response()
} else {
tracing::debug!(session_id = %session_id, "Session already removed or never existed");
StatusCode::OK.into_response()
}
}
pub(super) async fn handle_health() -> Response {
StatusCode::OK.into_response()
}
fn sse_json_response(response: impl serde::Serialize) -> Response {
let json = match serde_json::to_string(&response) {
Ok(s) => s,
Err(e) => {
tracing::error!(error = %e, "Failed to serialize response for SSE wrapping");
return StatusCode::INTERNAL_SERVER_ERROR.into_response();
}
};
let sse_body = format!("event: message\ndata: {json}\n\n");
(
StatusCode::OK,
[
(header::CONTENT_TYPE, "text/event-stream"),
(header::CACHE_CONTROL, "no-cache"),
],
sse_body,
)
.into_response()
}
pub(super) fn json_rpc_error_response(
id: Option<crate::protocol::RequestId>,
error: JsonRpcError,
) -> Response {
let response = JsonRpcResponse::error(id, error);
axum::Json(response).into_response()
}
pub(super) fn json_rpc_error_response_with_status(
id: Option<crate::protocol::RequestId>,
error: JsonRpcError,
status: StatusCode,
) -> Response {
let mut response = json_rpc_error_response(id, error);
*response.status_mut() = status;
response
}
fn body_too_large_response(limit: usize) -> Response {
let mut resp = json_rpc_error_response(
None,
JsonRpcError::invalid_request(format!(
"Request body exceeds the maximum size of {} bytes",
limit
)),
);
*resp.status_mut() = StatusCode::PAYLOAD_TOO_LARGE;
resp
}
fn is_length_limit_error(err: &axum::Error) -> bool {
let mut source: Option<&(dyn std::error::Error + 'static)> = Some(err);
while let Some(e) = source {
if e.is::<http_body_util::LengthLimitError>() {
return true;
}
source = e.source();
}
false
}