use std::collections::HashMap;
use std::convert::Infallible;
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use std::sync::atomic::{AtomicI64, AtomicU64, Ordering};
use std::time::{Duration, Instant};
use axum::{
Router,
extract::State,
http::{HeaderMap, HeaderValue, StatusCode, header},
response::{IntoResponse, Response, Sse, sse::Event},
routing::{delete, get, post},
};
#[cfg(feature = "stateless")]
use tokio::sync::mpsc;
use tokio::sync::{Mutex, RwLock, broadcast, oneshot};
use tokio_stream::StreamExt;
use tokio_stream::wrappers::BroadcastStream;
#[cfg(feature = "stateless")]
use crate::context::ServerNotification;
use crate::context::{
ChannelClientRequester, ClientRequesterHandle, NotificationReceiver, OutgoingRequest,
OutgoingRequestReceiver, notification_channel, outgoing_request_channel,
};
use crate::error::{Error, JsonRpcError, Result};
#[cfg(feature = "stateless")]
use crate::error::{ErrorCode, McpErrorCode};
use crate::jsonrpc::{JsonRpcService, apply_protocol_result_fields};
use crate::protocol::{
ClientCapabilities, Implementation, JsonRpcNotification, JsonRpcRequest, JsonRpcResponse,
LATEST_PROTOCOL_VERSION, McpNotification, PROTOCOL_VERSION_2026_07_28, RequestId,
};
#[cfg(feature = "stateless")]
use crate::protocol::{SubscriptionFilter, SubscriptionsListenParams};
use crate::router::{McpRouter, RouterRequest, RouterResponse};
use crate::transport::service::{
CatchError, InjectAnnotations, McpBoxService, ServiceFactory, identity_factory,
};
#[cfg(feature = "stateless")]
use crate::transport::subscriptions::{
accepted_subscription_filter, subscription_complete_response, subscription_matches,
tagged_subscription_notification,
};
use crate::{ProtocolSupport, ProtocolSupportError};
use tower::util::BoxCloneService;
#[cfg(feature = "stateless")]
fn stash_per_request_meta(req: &JsonRpcRequest, ext: &mut crate::router::Extensions) {
if let Some(params) = req.params.as_ref()
&& let Some(meta) = crate::stateless::StatelessRequestMeta::from_params(params)
{
ext.insert(meta);
}
}
pub const MCP_SESSION_ID_HEADER: &str = "mcp-session-id";
pub const MCP_PROTOCOL_VERSION_HEADER: &str = "mcp-protocol-version";
pub const MCP_METHOD_HEADER: &str = "mcp-method";
pub const MCP_NAME_HEADER: &str = "mcp-name";
pub const MCP_PARAM_HEADER_PREFIX: &str = "mcp-param-";
pub const DEFAULT_MAX_BODY_SIZE: usize = 4 * 1024 * 1024;
const SSE_MESSAGE_EVENT: &str = "message";
const LAST_EVENT_ID_HEADER: &str = "last-event-id";
struct PendingRequest {
response_tx: oneshot::Sender<Result<serde_json::Value>>,
}
type AssociatedCall = Pin<Box<dyn Future<Output = Result<JsonRpcResponse>> + Send + 'static>>;
enum SessionServiceSource {
Router {
router: McpRouter,
factory: ServiceFactory,
},
Boxed(std::sync::Mutex<McpBoxService>),
}
struct Session {
id: String,
service_source: SessionServiceSource,
notifications_tx: broadcast::Sender<String>,
created_at: Instant,
last_accessed: RwLock<Instant>,
pending_requests: Mutex<HashMap<RequestId, PendingRequest>>,
request_id_allocator: Option<Arc<AtomicI64>>,
protocol_version: RwLock<String>,
client_info: RwLock<Option<Implementation>>,
client_capabilities: RwLock<Option<ClientCapabilities>>,
event_counter: AtomicU64,
event_store: Arc<dyn crate::event_store::EventStore>,
initialized_notification_received: std::sync::atomic::AtomicBool,
}
impl Session {
fn new(
router: McpRouter,
sampling_enabled: bool,
service_factory: ServiceFactory,
event_store: Arc<dyn crate::event_store::EventStore>,
) -> Self {
let (notifications_tx, _) = broadcast::channel(100);
let (notif_sender, mut notif_receiver) = notification_channel(256);
let router = router.with_notification_sender(notif_sender);
let broadcast_tx = notifications_tx.clone();
tokio::spawn(async move {
while let Some(notification) = notif_receiver.recv().await {
if let Some(json) = crate::transport::stdio::serialize_notification(¬ification) {
let _ = broadcast_tx.send(json);
}
}
});
let request_id_allocator = if sampling_enabled {
Some(Arc::new(AtomicI64::new(1)))
} else {
None
};
let now = Instant::now();
Self {
id: uuid::Uuid::new_v4().to_string(),
service_source: SessionServiceSource::Router {
router,
factory: service_factory,
},
notifications_tx,
created_at: now,
last_accessed: RwLock::new(now),
pending_requests: Mutex::new(HashMap::new()),
request_id_allocator,
protocol_version: RwLock::new(LATEST_PROTOCOL_VERSION.to_string()),
client_info: RwLock::new(None),
client_capabilities: RwLock::new(None),
event_counter: AtomicU64::new(0),
event_store,
initialized_notification_received: std::sync::atomic::AtomicBool::new(false),
}
}
fn from_service(
service: McpBoxService,
event_store: Arc<dyn crate::event_store::EventStore>,
) -> Self {
let (notifications_tx, _) = broadcast::channel(100);
let now = Instant::now();
Self {
id: uuid::Uuid::new_v4().to_string(),
service_source: SessionServiceSource::Boxed(std::sync::Mutex::new(service)),
notifications_tx,
created_at: now,
last_accessed: RwLock::new(now),
pending_requests: Mutex::new(HashMap::new()),
request_id_allocator: None,
protocol_version: RwLock::new(LATEST_PROTOCOL_VERSION.to_string()),
client_info: RwLock::new(None),
client_capabilities: RwLock::new(None),
event_counter: AtomicU64::new(0),
event_store,
initialized_notification_received: std::sync::atomic::AtomicBool::new(false),
}
}
fn restored(
record: &crate::session_store::SessionRecord,
router: McpRouter,
sampling_enabled: bool,
service_factory: ServiceFactory,
event_store: Arc<dyn crate::event_store::EventStore>,
) -> Self {
router.session().mark_initialized();
let (notifications_tx, _) = broadcast::channel(100);
let (notif_sender, mut notif_receiver) = notification_channel(256);
let router = router.with_notification_sender(notif_sender);
let broadcast_tx = notifications_tx.clone();
tokio::spawn(async move {
while let Some(notification) = notif_receiver.recv().await {
if let Some(json) = crate::transport::stdio::serialize_notification(¬ification) {
let _ = broadcast_tx.send(json);
}
}
});
let request_id_allocator = if sampling_enabled {
Some(Arc::new(AtomicI64::new(1)))
} else {
None
};
let now = Instant::now();
Self {
id: record.id.clone(),
service_source: SessionServiceSource::Router {
router,
factory: service_factory,
},
notifications_tx,
created_at: now,
last_accessed: RwLock::new(now),
pending_requests: Mutex::new(HashMap::new()),
request_id_allocator,
protocol_version: RwLock::new(record.protocol_version.clone()),
client_info: RwLock::new(record.client_info.clone()),
client_capabilities: RwLock::new(record.client_capabilities.clone()),
event_counter: AtomicU64::new(0),
event_store,
initialized_notification_received: std::sync::atomic::AtomicBool::new(true),
}
}
fn from_service_restored(
service: McpBoxService,
record: &crate::session_store::SessionRecord,
event_store: Arc<dyn crate::event_store::EventStore>,
) -> Self {
let (notifications_tx, _) = broadcast::channel(100);
let now = Instant::now();
Self {
id: record.id.clone(),
service_source: SessionServiceSource::Boxed(std::sync::Mutex::new(service)),
notifications_tx,
created_at: now,
last_accessed: RwLock::new(now),
pending_requests: Mutex::new(HashMap::new()),
request_id_allocator: None,
protocol_version: RwLock::new(record.protocol_version.clone()),
client_info: RwLock::new(record.client_info.clone()),
client_capabilities: RwLock::new(record.client_capabilities.clone()),
event_counter: AtomicU64::new(0),
event_store,
initialized_notification_received: std::sync::atomic::AtomicBool::new(true),
}
}
fn make_service(&self) -> McpBoxService {
match &self.service_source {
SessionServiceSource::Router { router, factory } => (factory)(router.clone()),
SessionServiceSource::Boxed(mutex) => mutex.lock().unwrap().clone(),
}
}
fn handle_notification(&self, notification: McpNotification) {
match &self.service_source {
SessionServiceSource::Router { router, .. } => {
router.handle_notification(notification);
}
SessionServiceSource::Boxed(_) => {
tracing::debug!(
notification = ?notification,
"Notification received on service-based session (not forwarded)"
);
}
}
}
fn next_event_id(&self) -> u64 {
self.event_counter.fetch_add(1, Ordering::SeqCst)
}
async fn buffer_event(&self, id: u64, data: String) {
let record = crate::event_store::EventRecord::new(id, data);
if let Err(e) = self.event_store.append(&self.id, record).await {
tracing::warn!(session_id = %self.id, event_id = id, error = %e, "Failed to append event to event store");
}
}
async fn get_events_after(&self, after_id: u64) -> Vec<crate::event_store::EventRecord> {
match self.event_store.replay_after(&self.id, after_id).await {
Ok(events) => events,
Err(e) => {
tracing::warn!(session_id = %self.id, error = %e, "Failed to replay events from event store");
Vec::new()
}
}
}
async fn touch(&self) {
*self.last_accessed.write().await = Instant::now();
}
async fn is_expired(&self, ttl: Duration) -> bool {
self.last_accessed.read().await.elapsed() > ttl
}
async fn add_pending_request(
&self,
id: RequestId,
response_tx: oneshot::Sender<Result<serde_json::Value>>,
) {
let mut pending = self.pending_requests.lock().await;
pending.insert(id, PendingRequest { response_tx });
}
async fn complete_pending_request(
&self,
id: &RequestId,
result: Result<serde_json::Value>,
) -> bool {
let pending = {
let mut pending_requests = self.pending_requests.lock().await;
pending_requests.remove(id)
};
match pending {
Some(pending) => {
let _ = pending.response_tx.send(result);
true
}
None => false,
}
}
async fn fail_pending_requests(&self, ids: &[RequestId], message: &str) {
let removed = {
let mut pending = self.pending_requests.lock().await;
ids.iter()
.filter_map(|id| pending.remove(id))
.collect::<Vec<_>>()
};
for pending in removed {
let _ = pending
.response_tx
.send(Err(Error::Transport(message.to_string())));
}
}
}
pub const DEFAULT_SESSION_TTL: Duration = Duration::from_secs(30 * 60);
const DEFAULT_CLEANUP_INTERVAL: Duration = Duration::from_secs(60);
#[derive(Debug, Clone)]
pub struct SessionConfig {
pub ttl: Duration,
pub max_sessions: Option<usize>,
pub cleanup_interval: Duration,
pub strict_initialization: bool,
}
impl Default for SessionConfig {
fn default() -> Self {
Self {
ttl: DEFAULT_SESSION_TTL,
max_sessions: None,
cleanup_interval: DEFAULT_CLEANUP_INTERVAL,
strict_initialization: true,
}
}
}
impl SessionConfig {
pub fn with_ttl(ttl: Duration) -> Self {
Self {
ttl,
..Default::default()
}
}
pub fn max_sessions(mut self, max: usize) -> Self {
self.max_sessions = Some(max);
self
}
pub fn cleanup_interval(mut self, interval: Duration) -> Self {
self.cleanup_interval = interval;
self
}
pub fn strict_initialization(mut self, enabled: bool) -> Self {
self.strict_initialization = enabled;
self
}
}
struct SessionRegistry {
sessions: RwLock<HashMap<String, Arc<Session>>>,
config: SessionConfig,
sampling_enabled: bool,
persistent: Arc<dyn crate::session_store::SessionStore>,
events: Arc<dyn crate::event_store::EventStore>,
service_source: ServiceSource,
auto_reinit: bool,
}
impl SessionRegistry {
fn new(
config: SessionConfig,
sampling_enabled: bool,
persistent: Arc<dyn crate::session_store::SessionStore>,
events: Arc<dyn crate::event_store::EventStore>,
service_source: ServiceSource,
auto_reinit: bool,
) -> Self {
Self {
sessions: RwLock::new(HashMap::new()),
config,
sampling_enabled,
persistent,
events,
service_source,
auto_reinit,
}
}
async fn record_for(&self, session: &Session) -> crate::session_store::SessionRecord {
let protocol_version = session.protocol_version.read().await.clone();
let last_accessed = session.last_accessed.read().await;
let mut record = crate::session_store::SessionRecord::new(
session.id.clone(),
protocol_version,
self.config.ttl,
);
record.client_info = session.client_info.read().await.clone();
record.client_capabilities = session.client_capabilities.read().await.clone();
let now = std::time::SystemTime::now();
let created_ago = session.created_at.elapsed();
let last_accessed_ago = last_accessed.elapsed();
record.created_at = now.checked_sub(created_ago).unwrap_or(now);
record.last_accessed = now.checked_sub(last_accessed_ago).unwrap_or(now);
record.expires_at = record.last_accessed + self.config.ttl;
record
}
async fn persist_new(&self, session: &Session) {
let record = self.record_for(session).await;
if let Err(e) = self.persistent.create(&mut record.clone()).await {
tracing::warn!(session_id = %session.id, error = %e, "Failed to persist session record");
}
}
async fn save_record(&self, session: &Session) {
let record = self.record_for(session).await;
if let Err(e) = self.persistent.save(&record).await {
tracing::warn!(session_id = %session.id, error = %e, "Failed to save session record");
}
}
async fn create(
&self,
router: McpRouter,
service_factory: ServiceFactory,
) -> Option<Arc<Session>> {
let session = {
let mut sessions = self.sessions.write().await;
if let Some(max) = self.config.max_sessions
&& sessions.len() >= max
{
tracing::warn!(
max_sessions = max,
current = sessions.len(),
"Session limit reached, rejecting new session"
);
return None;
}
let session = Arc::new(Session::new(
router,
self.sampling_enabled,
service_factory,
self.events.clone(),
));
sessions.insert(session.id.clone(), session.clone());
tracing::debug!(session_id = %session.id, sampling = self.sampling_enabled, "Created new session");
session
};
self.persist_new(&session).await;
Some(session)
}
async fn create_from_service(&self, service: McpBoxService) -> Option<Arc<Session>> {
let session = {
let mut sessions = self.sessions.write().await;
if let Some(max) = self.config.max_sessions
&& sessions.len() >= max
{
tracing::warn!(
max_sessions = max,
current = sessions.len(),
"Session limit reached, rejecting new session"
);
return None;
}
let session = Arc::new(Session::from_service(service, self.events.clone()));
sessions.insert(session.id.clone(), session.clone());
tracing::debug!(session_id = %session.id, "Created new session from service");
session
};
self.persist_new(&session).await;
Some(session)
}
async fn create_initialized(
&self,
router: McpRouter,
service_factory: ServiceFactory,
) -> Option<Arc<Session>> {
router.session().mark_initialized();
let session = {
let mut sessions = self.sessions.write().await;
if let Some(max) = self.config.max_sessions
&& sessions.len() >= max
{
return None;
}
let session = Arc::new(Session::new(
router,
self.sampling_enabled,
service_factory,
self.events.clone(),
));
session
.initialized_notification_received
.store(true, Ordering::Release);
sessions.insert(session.id.clone(), session.clone());
tracing::debug!(session_id = %session.id, "Created pre-initialized session (optional_sessions)");
session
};
self.persist_new(&session).await;
Some(session)
}
async fn create_initialized_from_service(
&self,
service: McpBoxService,
) -> Option<Arc<Session>> {
let session = {
let mut sessions = self.sessions.write().await;
if let Some(max) = self.config.max_sessions
&& sessions.len() >= max
{
return None;
}
let session = Arc::new(Session::from_service(service, self.events.clone()));
session
.initialized_notification_received
.store(true, Ordering::Release);
sessions.insert(session.id.clone(), session.clone());
tracing::debug!(session_id = %session.id, "Created pre-initialized session from service (optional_sessions)");
session
};
self.persist_new(&session).await;
Some(session)
}
async fn get(&self, id: &str) -> Option<Arc<Session>> {
{
let sessions = self.sessions.read().await;
if let Some(s) = sessions.get(id).cloned() {
s.touch().await;
return Some(s);
}
}
match self.persistent.load(id).await {
Ok(Some(record)) => {
tracing::info!(session_id = %id, "Restoring session from persistent store");
if let Some(session) = self.restore_from_record(record).await {
return Some(session);
}
}
Ok(None) => {}
Err(e) => {
tracing::warn!(session_id = %id, error = %e, "Failed to load session record");
}
}
if self.auto_reinit {
tracing::info!(session_id = %id, "Auto-reinitializing unknown session");
return self.auto_reinitialize(id).await;
}
None
}
async fn restore_from_record(
&self,
record: crate::session_store::SessionRecord,
) -> Option<Arc<Session>> {
let session = {
let mut sessions = self.sessions.write().await;
if let Some(max) = self.config.max_sessions
&& sessions.len() >= max
{
tracing::warn!(
max_sessions = max,
"Session limit reached, cannot restore session"
);
return None;
}
if let Some(existing) = sessions.get(&record.id).cloned() {
existing.touch().await;
return Some(existing);
}
let session: Arc<Session> = match &self.service_source {
ServiceSource::Router { router, factory } => Arc::new(Session::restored(
&record,
router.with_fresh_session(),
self.sampling_enabled,
factory.clone(),
self.events.clone(),
)),
ServiceSource::Service(svc) => {
let service = svc.lock().unwrap().clone();
Arc::new(Session::from_service_restored(
service,
&record,
self.events.clone(),
))
}
};
sessions.insert(record.id.clone(), session.clone());
tracing::debug!(session_id = %session.id, "Restored session into local registry");
session
};
if let Ok(events) = self.events.replay_after(&record.id, 0).await
&& let Some(max_id) = events.iter().map(|e| e.id).max()
{
session
.event_counter
.store(max_id + 1, std::sync::atomic::Ordering::SeqCst);
}
let mut refreshed = record;
refreshed.touch(self.config.ttl);
if let Err(e) = self.persistent.save(&refreshed).await {
tracing::warn!(session_id = %refreshed.id, error = %e, "Failed to refresh restored session record");
}
Some(session)
}
async fn auto_reinitialize(&self, id: &str) -> Option<Arc<Session>> {
let mut record = crate::session_store::SessionRecord::new(
id.to_string(),
LATEST_PROTOCOL_VERSION.to_string(),
self.config.ttl,
);
record.client_info = Some(crate::protocol::Implementation {
name: "auto-recovered".into(),
version: "unknown".into(),
title: None,
description: None,
icons: None,
website_url: None,
meta: None,
});
record.client_capabilities = Some(crate::protocol::ClientCapabilities::default());
if let Err(e) = self.persistent.create(&mut record).await {
tracing::warn!(session_id = %id, error = %e, "Failed to persist auto-reinitialized session");
}
self.restore_from_record(record).await
}
async fn remove(&self, id: &str) -> bool {
let removed = {
let mut sessions = self.sessions.write().await;
sessions.remove(id).is_some()
};
if removed {
tracing::debug!(session_id = %id, "Removed session");
if let Err(e) = self.persistent.delete(id).await {
tracing::warn!(session_id = %id, error = %e, "Failed to delete session record");
}
if let Err(e) = self.events.purge_session(id).await {
tracing::warn!(session_id = %id, error = %e, "Failed to purge session events");
}
}
removed
}
async fn broadcast_to_all(&self, json: &str) {
let sessions = self.sessions.read().await;
for session in sessions.values() {
let _ = session.notifications_tx.send(json.to_string());
}
}
async fn cleanup_expired(&self) -> usize {
let expired = {
let mut sessions = self.sessions.write().await;
let ttl = self.config.ttl;
let mut expired = Vec::new();
for (id, session) in sessions.iter() {
if session.is_expired(ttl).await {
expired.push(id.clone());
}
}
for id in &expired {
sessions.remove(id);
tracing::debug!(session_id = %id, "Expired session removed");
}
if !expired.is_empty() {
tracing::info!(
expired_count = expired.len(),
remaining = sessions.len(),
"Session cleanup completed"
);
}
expired
};
for id in &expired {
if let Err(e) = self.persistent.delete(id).await {
tracing::warn!(session_id = %id, error = %e, "Failed to delete expired session record");
}
if let Err(e) = self.events.purge_session(id).await {
tracing::warn!(session_id = %id, error = %e, "Failed to purge expired session events");
}
}
expired.len()
}
}
#[derive(Debug, Clone)]
pub struct SessionInfo {
pub id: String,
pub created_at: Duration,
pub last_activity: Duration,
}
#[derive(Clone)]
pub struct SessionHandle {
store: Arc<SessionRegistry>,
#[cfg(feature = "stateless")]
modern_subscriptions: Arc<ModernSubscriptionRegistry>,
}
impl SessionHandle {
pub async fn session_count(&self) -> usize {
self.store.sessions.read().await.len()
}
pub async fn list_sessions(&self) -> Vec<SessionInfo> {
let sessions = self.store.sessions.read().await;
let mut infos = Vec::with_capacity(sessions.len());
for session in sessions.values() {
let last_accessed = session.last_accessed.read().await;
infos.push(SessionInfo {
id: session.id.clone(),
created_at: session.created_at.elapsed(),
last_activity: last_accessed.elapsed(),
});
}
infos
}
pub async fn terminate_session(&self, id: &str) -> bool {
self.store.remove(id).await
}
#[cfg(feature = "stateless")]
pub fn subscription_count(&self) -> usize {
self.modern_subscriptions.len()
}
#[cfg(feature = "stateless")]
pub fn close_subscriptions(&self) -> usize {
self.modern_subscriptions.close_all()
}
}
#[cfg(feature = "stateless")]
impl AppState {
fn tasks_extension_enabled(&self) -> bool {
match &self.service_source {
ServiceSource::Router { router, .. } => router.final_tasks_enabled(),
ServiceSource::Service(_) => false,
}
}
}
#[derive(Clone)]
enum ServiceSource {
Router {
router: McpRouter,
factory: ServiceFactory,
},
Service(Arc<std::sync::Mutex<McpBoxService>>),
}
struct AppState {
service_source: ServiceSource,
protocol_support: ProtocolSupport,
sessions: Arc<SessionRegistry>,
validate_origin: bool,
allowed_origins: Vec<String>,
validate_host: bool,
allowed_hosts: Vec<String>,
optional_sessions: bool,
strict_initialization: bool,
#[cfg(feature = "stateless")]
stateless_config: Option<crate::stateless::StatelessConfig>,
#[cfg(feature = "stateless")]
stamp_server_info: bool,
#[cfg(feature = "stateless")]
modern_subscriptions: Arc<ModernSubscriptionRegistry>,
sse_responses: bool,
max_body_size: usize,
}
#[cfg(feature = "stateless")]
struct ModernSubscription {
subscription_id: RequestId,
filter: SubscriptionFilter,
tx: mpsc::UnboundedSender<String>,
}
#[cfg(feature = "stateless")]
struct ModernSubscriptionRegistry {
next_key: AtomicU64,
subscriptions: std::sync::Mutex<HashMap<u64, ModernSubscription>>,
server_info: Option<Implementation>,
}
#[cfg(feature = "stateless")]
impl ModernSubscriptionRegistry {
fn new(server_info: Option<Implementation>) -> Self {
Self {
next_key: AtomicU64::new(0),
subscriptions: std::sync::Mutex::new(HashMap::new()),
server_info,
}
}
fn register(
self: &Arc<Self>,
subscription_id: RequestId,
filter: SubscriptionFilter,
) -> (mpsc::UnboundedReceiver<String>, ModernSubscriptionGuard) {
let key = self.next_key.fetch_add(1, Ordering::Relaxed);
let (tx, rx) = mpsc::unbounded_channel();
self.subscriptions.lock().unwrap().insert(
key,
ModernSubscription {
subscription_id,
filter,
tx,
},
);
(
rx,
ModernSubscriptionGuard {
key,
registry: self.clone(),
},
)
}
fn publish(&self, notification: &ServerNotification) -> bool {
let subscription_scoped = matches!(
notification,
ServerNotification::ResourceUpdated { .. }
| ServerNotification::ResourcesListChanged
| ServerNotification::ToolsListChanged
| ServerNotification::PromptsListChanged
| ServerNotification::FinalTaskStatusChanged(_)
);
if !subscription_scoped {
return false;
}
let mut subscriptions = self.subscriptions.lock().unwrap();
tracing::trace!(
active_subscriptions = subscriptions.len(),
notification = ?notification,
"Routing final-protocol subscription notification"
);
subscriptions.retain(|_, subscription| {
if subscription_matches(notification, &subscription.filter)
&& let Some(json) =
tagged_subscription_notification(notification, &subscription.subscription_id)
{
return subscription.tx.send(json).is_ok();
}
!subscription.tx.is_closed()
});
true
}
fn len(&self) -> usize {
self.subscriptions.lock().unwrap().len()
}
fn close_all(&self) -> usize {
let subscriptions = {
let mut active = self.subscriptions.lock().unwrap();
active
.drain()
.map(|(_, subscription)| subscription)
.collect::<Vec<_>>()
};
let count = subscriptions.len();
for subscription in subscriptions {
let response = subscription_complete_response(
subscription.subscription_id,
self.server_info.clone(),
);
if let Ok(json) = serde_json::to_string(&response) {
let _ = subscription.tx.send(json);
}
}
count
}
}
#[cfg(feature = "stateless")]
impl Default for ModernSubscriptionRegistry {
fn default() -> Self {
Self::new(None)
}
}
#[cfg(feature = "stateless")]
struct ModernSubscriptionGuard {
key: u64,
registry: Arc<ModernSubscriptionRegistry>,
}
#[cfg(feature = "stateless")]
impl Drop for ModernSubscriptionGuard {
fn drop(&mut self) {
self.registry
.subscriptions
.lock()
.unwrap()
.remove(&self.key);
}
}
#[cfg(feature = "oauth")]
#[derive(Clone)]
pub(crate) struct OAuthConfig {
pub(crate) metadata: crate::oauth::ProtectedResourceMetadata,
}
pub struct HttpTransport {
service_source: ServiceSource,
protocol_support: ProtocolSupport,
validate_origin: bool,
allowed_origins: Vec<String>,
validate_host: bool,
allowed_hosts: Vec<String>,
session_config: SessionConfig,
sampling_enabled: bool,
optional_sessions: bool,
session_store: Arc<dyn crate::session_store::SessionStore>,
event_store: Arc<dyn crate::event_store::EventStore>,
auto_reinit_sessions: bool,
external_notifications: Option<NotificationReceiver>,
#[cfg(feature = "stateless")]
stateless_config: Option<crate::stateless::StatelessConfig>,
#[cfg(feature = "stateless")]
stamp_server_info: bool,
#[cfg(feature = "oauth")]
oauth_config: Option<OAuthConfig>,
sse_responses: bool,
max_body_size: usize,
}
impl HttpTransport {
pub fn new(router: McpRouter) -> Self {
Self {
service_source: ServiceSource::Router {
router,
factory: identity_factory(),
},
protocol_support: ProtocolSupport::default(),
validate_origin: true,
allowed_origins: vec![],
validate_host: true,
allowed_hosts: vec![],
session_config: SessionConfig::default(),
sampling_enabled: false,
optional_sessions: true,
session_store: Arc::new(crate::session_store::MemorySessionStore::new()),
event_store: Arc::new(crate::event_store::MemoryEventStore::new()),
auto_reinit_sessions: false,
external_notifications: None,
#[cfg(feature = "stateless")]
stateless_config: None,
#[cfg(feature = "stateless")]
stamp_server_info: true,
#[cfg(feature = "oauth")]
oauth_config: None,
sse_responses: false,
max_body_size: DEFAULT_MAX_BODY_SIZE,
}
}
pub fn from_service<S>(service: S) -> Self
where
S: tower::Service<
RouterRequest,
Response = RouterResponse,
Error = std::convert::Infallible,
> + Clone
+ Send
+ 'static,
S::Future: Send,
{
Self {
service_source: ServiceSource::Service(Arc::new(std::sync::Mutex::new(
BoxCloneService::new(service),
))),
protocol_support: ProtocolSupport::default(),
validate_origin: true,
allowed_origins: vec![],
validate_host: true,
allowed_hosts: vec![],
session_config: SessionConfig::default(),
sampling_enabled: false,
optional_sessions: true,
session_store: Arc::new(crate::session_store::MemorySessionStore::new()),
event_store: Arc::new(crate::event_store::MemoryEventStore::new()),
auto_reinit_sessions: false,
external_notifications: None,
#[cfg(feature = "stateless")]
stateless_config: None,
#[cfg(feature = "stateless")]
stamp_server_info: true,
#[cfg(feature = "oauth")]
oauth_config: None,
sse_responses: false,
max_body_size: DEFAULT_MAX_BODY_SIZE,
}
}
pub fn with_notifications(router: McpRouter, notification_rx: NotificationReceiver) -> Self {
Self {
external_notifications: Some(notification_rx),
..Self::new(router)
}
}
pub fn external_notifications(mut self, notification_rx: NotificationReceiver) -> Self {
self.external_notifications = Some(notification_rx);
self
}
pub fn with_sampling(mut self) -> Self {
self.sampling_enabled = true;
self
}
pub fn require_sessions(mut self) -> Self {
self.optional_sessions = false;
self
}
pub fn protocol_support(mut self, support: ProtocolSupport) -> Self {
self.protocol_support = support;
self
}
pub fn protocol_versions<I, S>(
mut self,
versions: I,
) -> std::result::Result<Self, ProtocolSupportError>
where
I: IntoIterator<Item = S>,
S: Into<String>,
{
self.protocol_support = ProtocolSupport::try_new(versions)?;
Ok(self)
}
pub fn sse_responses(mut self, enabled: bool) -> Self {
self.sse_responses = enabled;
self
}
#[cfg(feature = "stateless")]
pub fn stamp_server_info(mut self, enabled: bool) -> Self {
self.stamp_server_info = enabled;
self
}
pub fn max_body_size(mut self, bytes: usize) -> Self {
self.max_body_size = bytes;
self
}
#[cfg(feature = "stateless")]
pub fn stateless(mut self, config: crate::stateless::StatelessConfig) -> Self {
self.stateless_config = Some(config);
self
}
pub fn disable_origin_validation(mut self) -> Self {
self.validate_origin = false;
self
}
pub fn allowed_origins(mut self, origins: Vec<String>) -> Self {
self.allowed_origins = origins;
self
}
pub fn disable_host_validation(mut self) -> Self {
self.validate_host = false;
self
}
pub fn allowed_hosts(mut self, hosts: Vec<String>) -> Self {
self.allowed_hosts = hosts;
self
}
pub fn session_config(mut self, config: SessionConfig) -> Self {
self.session_config = config;
self
}
pub fn session_ttl(mut self, ttl: Duration) -> Self {
self.session_config.ttl = ttl;
self
}
pub fn max_sessions(mut self, max: usize) -> Self {
self.session_config.max_sessions = Some(max);
self
}
pub fn session_store(mut self, store: Arc<dyn crate::session_store::SessionStore>) -> Self {
self.session_store = store;
self
}
pub fn event_store(mut self, store: Arc<dyn crate::event_store::EventStore>) -> Self {
self.event_store = store;
self
}
pub fn auto_reinitialize_sessions(mut self, enabled: bool) -> Self {
self.auto_reinit_sessions = enabled;
self
}
#[cfg(feature = "oauth")]
pub fn oauth(mut self, metadata: crate::oauth::ProtectedResourceMetadata) -> Self {
self.oauth_config = Some(OAuthConfig { metadata });
self
}
pub fn layer<L>(mut self, layer: L) -> Self
where
L: tower::Layer<McpRouter> + Send + Sync + 'static,
L::Service:
tower::Service<RouterRequest, Response = RouterResponse> + Clone + Send + 'static,
<L::Service as tower::Service<RouterRequest>>::Error: std::fmt::Display + Send,
<L::Service as tower::Service<RouterRequest>>::Future: Send,
{
match &mut self.service_source {
ServiceSource::Router { factory, .. } => {
*factory = Arc::new(move |router: McpRouter| {
let annotations = router.tool_annotations_map();
let wrapped = layer.layer(router);
tower::util::BoxCloneService::new(InjectAnnotations::new(
CatchError::new(wrapped),
annotations,
))
});
}
ServiceSource::Service(_) => {
panic!(
"layer() cannot be used with from_service() — \
wrap the service with middleware before passing it in"
);
}
}
self
}
fn build_state(&self) -> Arc<AppState> {
#[cfg(feature = "stateless")]
let modern_subscriptions = Arc::new(ModernSubscriptionRegistry::new(
match &self.service_source {
ServiceSource::Router { router, .. } if self.stamp_server_info => {
Some(router.implementation())
}
_ => None,
},
));
#[cfg(feature = "stateless")]
let service_source = match &self.service_source {
ServiceSource::Router { router, factory } => {
let (tx, mut rx) = notification_channel(256);
let direct_subscriptions = modern_subscriptions.clone();
router.attach_modern_notification_sink(Arc::new(move |notification| {
direct_subscriptions.publish(notification)
}));
let subscriptions = modern_subscriptions.clone();
tokio::spawn(async move {
while let Some(notification) = rx.recv().await {
subscriptions.publish(¬ification);
}
});
ServiceSource::Router {
router: router.clone().with_notification_sender(tx),
factory: factory.clone(),
}
}
ServiceSource::Service(service) => ServiceSource::Service(service.clone()),
};
#[cfg(not(feature = "stateless"))]
let service_source = self.service_source.clone();
let sessions = Arc::new(SessionRegistry::new(
self.session_config.clone(),
self.sampling_enabled,
self.session_store.clone(),
self.event_store.clone(),
service_source.clone(),
self.auto_reinit_sessions,
));
let cleanup_sessions = sessions.clone();
let cleanup_interval = self.session_config.cleanup_interval;
tokio::spawn(async move {
loop {
tokio::time::sleep(cleanup_interval).await;
cleanup_sessions.cleanup_expired().await;
}
});
Arc::new(AppState {
service_source,
protocol_support: self.protocol_support.clone(),
sessions,
validate_origin: self.validate_origin,
allowed_origins: self.allowed_origins.clone(),
validate_host: self.validate_host,
allowed_hosts: self.allowed_hosts.clone(),
optional_sessions: self.optional_sessions,
strict_initialization: self.session_config.strict_initialization,
#[cfg(feature = "stateless")]
stateless_config: self.stateless_config.clone(),
#[cfg(feature = "stateless")]
stamp_server_info: self.stamp_server_info,
#[cfg(feature = "stateless")]
modern_subscriptions,
sse_responses: self.sse_responses,
max_body_size: self.max_body_size,
})
}
pub fn into_router(self) -> Router {
let (router, _handle) = self.into_router_with_handle();
router
}
pub fn into_router_with_handle(mut self) -> (Router, SessionHandle) {
let external_rx = self.external_notifications.take();
let state = self.build_state();
let handle = SessionHandle {
store: state.sessions.clone(),
#[cfg(feature = "stateless")]
modern_subscriptions: state.modern_subscriptions.clone(),
};
spawn_external_notification_fanout(
external_rx,
state.sessions.clone(),
#[cfg(feature = "stateless")]
state.modern_subscriptions.clone(),
);
let router = Router::new()
.route("/", post(handle_post))
.route("/", get(handle_get))
.route("/", delete(handle_delete))
.route("/health", get(handle_health))
.with_state(state);
#[cfg(feature = "oauth")]
let router = self.add_oauth_route(router, "");
(router, handle)
}
pub fn into_router_at(self, path: &str) -> Router {
let (router, _handle) = self.into_router_at_with_handle(path);
router
}
pub fn into_router_at_with_handle(mut self, path: &str) -> (Router, SessionHandle) {
let external_rx = self.external_notifications.take();
let state = self.build_state();
let handle = SessionHandle {
store: state.sessions.clone(),
#[cfg(feature = "stateless")]
modern_subscriptions: state.modern_subscriptions.clone(),
};
spawn_external_notification_fanout(
external_rx,
state.sessions.clone(),
#[cfg(feature = "stateless")]
state.modern_subscriptions.clone(),
);
let mcp_router = Router::new()
.route("/", post(handle_post))
.route("/", get(handle_get))
.route("/", delete(handle_delete))
.route("/health", get(handle_health))
.with_state(state);
let router = Router::new().nest(path, mcp_router);
#[cfg(feature = "oauth")]
let router = self.add_oauth_route(router, path);
(router, handle)
}
pub async fn serve(self, addr: &str) -> Result<()> {
let listener = tokio::net::TcpListener::bind(addr)
.await
.map_err(|e| Error::Transport(format!("Failed to bind to {}: {}", addr, e)))?;
tracing::info!("MCP HTTP transport listening on {}", addr);
let router = self.into_router();
axum::serve(listener, router)
.await
.map_err(|e| Error::Transport(format!("Server error: {}", e)))?;
Ok(())
}
#[cfg(feature = "oauth")]
fn add_oauth_route(&self, router: Router, base_path: &str) -> Router {
if let Some(ref config) = self.oauth_config {
let metadata = config.metadata.clone();
let well_known_path = if base_path.is_empty() {
crate::oauth::ProtectedResourceMetadata::well_known_path().to_string()
} else {
format!(
"{}{}",
base_path.trim_end_matches('/'),
crate::oauth::ProtectedResourceMetadata::well_known_path()
)
};
router.route(
&well_known_path,
get(move || {
let m = metadata.clone();
async move { axum::Json(m) }
}),
)
} else {
router
}
}
}
fn spawn_external_notification_fanout(
rx: Option<NotificationReceiver>,
sessions: Arc<SessionRegistry>,
#[cfg(feature = "stateless")] modern_subscriptions: Arc<ModernSubscriptionRegistry>,
) {
let Some(mut rx) = rx else {
return;
};
tokio::spawn(async move {
while let Some(notification) = rx.recv().await {
#[cfg(feature = "stateless")]
modern_subscriptions.publish(¬ification);
if let Some(json) = crate::transport::stdio::serialize_notification(¬ification) {
sessions.broadcast_to_all(&json).await;
}
}
tracing::debug!("External notification channel closed; fan-out task exiting");
});
}
fn is_localhost_origin(origin: &str) -> bool {
if let Some(rest) = origin
.strip_prefix("http://")
.or_else(|| origin.strip_prefix("https://"))
{
is_localhost_host(rest)
} else {
false
}
}
fn is_localhost_host(host: &str) -> bool {
let host_only = if host.starts_with('[') {
host.split(']')
.next()
.unwrap_or(host)
.trim_start_matches('[')
} else {
host.split(':').next().unwrap_or(host)
};
matches!(host_only, "localhost" | "127.0.0.1" | "::1")
}
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())
}
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 {
parsed.get("method").is_none()
&& (parsed.get("result").is_some() || parsed.get("error").is_some())
}
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"
)
}
#[cfg(feature = "stateless")]
fn modern_response_status(response: &JsonRpcResponse) -> StatusCode {
let JsonRpcResponse::Error(error) = response else {
return StatusCode::OK;
};
if error.error.code == ErrorCode::MethodNotFound as i32 {
StatusCode::NOT_FOUND
} else if error.error.code == McpErrorCode::MissingRequiredClientCapability.code() {
StatusCode::BAD_REQUEST
} else {
StatusCode::OK
}
}
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()))
}
})
}
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)),
);
}
};
#[cfg(feature = "oauth")]
let http_extensions = parts.extensions;
#[cfg(not(feature = "oauth"))]
let _ = parts.extensions;
let parsed: serde_json::Value = match serde_json::from_str(&body) {
Ok(v) => v,
Err(e) => {
return json_rpc_error_response(
None,
JsonRpcError::parse_error(format!("Invalid JSON: {}", e)),
);
}
};
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 sep_2243_mode = super::http_headers::mode_for_version(&body_version);
if let Err(error) = super::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 = super::http_headers::mode_for_version(version);
if let Err(err) = super::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 request: JsonRpcRequest = match serde_json::from_value(parsed) {
Ok(r) => r,
Err(e) => {
return json_rpc_error_response(
None,
JsonRpcError::parse_error(format!("Invalid request: {}", e)),
);
}
};
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_initialized();
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);
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 request: JsonRpcRequest = match serde_json::from_value(parsed) {
Ok(r) => r,
Err(e) => {
return json_rpc_error_response(
None,
JsonRpcError::parse_error(format!("Invalid request: {}", e)),
);
}
};
let mut service = match &state.service_source {
ServiceSource::Router { router, factory } => {
let ephemeral = router.with_fresh_session();
ephemeral.session().mark_initialized();
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);
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).await;
}
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 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"),
);
}
}
}
if !is_init
&& let Some(version) = get_protocol_version(&headers)
&& !state.protocol_support.contains(&version)
{
let id = extract_request_id(&parsed);
return json_rpc_error_response(
id,
JsonRpcError::unsupported_protocol_version(
version,
state.protocol_support.versions().iter().map(String::as_str),
),
);
}
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 = super::http_headers::mode_for_version(&sep_2243_version);
if let Err(err) = super::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 request: JsonRpcRequest = match serde_json::from_value(parsed) {
Ok(r) => r,
Err(e) => {
return json_rpc_error_response(
None,
JsonRpcError::parse_error(format!("Invalid request: {}", e)),
);
}
};
let mut service = JsonRpcService::new(session.make_service());
#[allow(unused_mut)]
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);
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)
}
#[cfg(feature = "stateless")]
fn is_stateless_protocol_version(version: &str) -> bool {
version == PROTOCOL_VERSION_2026_07_28
}
#[cfg(feature = "stateless")]
fn stamp_server_info(response: &mut JsonRpcResponse, implementation: &Implementation) {
let JsonRpcResponse::Result(result) = response else {
return;
};
let Some(obj) = result.result.as_object_mut() else {
return;
};
let meta = obj
.entry("_meta")
.or_insert_with(|| serde_json::Value::Object(Default::default()));
let Some(meta_obj) = meta.as_object_mut() else {
return;
};
if let Ok(value) = serde_json::to_value(implementation) {
meta_obj.insert("io.modelcontextprotocol/serverInfo".to_string(), value);
}
}
#[cfg(feature = "stateless")]
struct CancelOnDisconnect(Option<crate::context::CancellationToken>);
#[cfg(feature = "stateless")]
impl CancelOnDisconnect {
fn arm(token: crate::context::CancellationToken) -> Self {
Self(Some(token))
}
fn disarm(&mut self) {
self.0 = None;
}
}
#[cfg(feature = "stateless")]
impl Drop for CancelOnDisconnect {
fn drop(&mut self) {
if let Some(token) = self.0.take() {
token.cancel();
}
}
}
#[cfg(feature = "stateless")]
struct StatelessSseContext {
version: String,
method: String,
cancel_guard: CancelOnDisconnect,
server_identity: Option<Implementation>,
subscriptions: Arc<ModernSubscriptionRegistry>,
}
#[cfg(feature = "stateless")]
fn stateless_sse_with_notifications(
first: crate::context::ServerNotification,
call: std::pin::Pin<
Box<dyn std::future::Future<Output = crate::error::Result<JsonRpcResponse>> + Send>,
>,
rx: crate::context::NotificationReceiver,
request: StatelessSseContext,
) -> Response {
struct Ctx {
call: Option<
std::pin::Pin<
Box<dyn std::future::Future<Output = crate::error::Result<JsonRpcResponse>> + Send>,
>,
>,
rx: crate::context::NotificationReceiver,
rx_open: bool,
queue: std::collections::VecDeque<String>,
terminal: Option<String>,
version: String,
method: String,
cancel_guard: CancelOnDisconnect,
server_identity: Option<Implementation>,
subscriptions: Arc<ModernSubscriptionRegistry>,
}
let mut queue = std::collections::VecDeque::new();
if !request.subscriptions.publish(&first)
&& let Some(json) = crate::transport::stdio::serialize_notification(&first)
{
queue.push_back(json);
}
let ctx = Ctx {
call: Some(call),
rx,
rx_open: true,
queue,
terminal: None,
version: request.version,
method: request.method,
cancel_guard: request.cancel_guard,
server_identity: request.server_identity,
subscriptions: request.subscriptions,
};
let stream = futures::stream::unfold(ctx, |mut ctx| async move {
loop {
if let Some(json) = ctx.queue.pop_front() {
return Some((
Ok::<_, Infallible>(Event::default().event(SSE_MESSAGE_EVENT).data(json)),
ctx,
));
}
if let Some(json) = ctx.terminal.take() {
return Some((
Ok(Event::default().event(SSE_MESSAGE_EVENT).data(json)),
ctx,
));
}
let mut call = ctx.call.take()?;
tokio::select! {
result = &mut call => {
ctx.cancel_guard.disarm();
while let Ok(n) = ctx.rx.try_recv() {
if !ctx.subscriptions.publish(&n)
&& let Some(json) =
crate::transport::stdio::serialize_notification(&n)
{
ctx.queue.push_back(json);
}
}
let terminal_json = match result {
Ok(mut response) => {
if ctx.method == "initialize"
&& let JsonRpcResponse::Result(ref mut r) = response
&& let Some(pv) = r.result.get_mut("protocolVersion")
{
*pv = serde_json::Value::String(ctx.version.clone());
}
apply_protocol_result_fields(
&mut response,
&ctx.method,
&ctx.version,
);
if let Some(ref identity) = ctx.server_identity {
stamp_server_info(&mut response, identity);
}
serde_json::to_string(&response).ok()
}
Err(e) => Some(
serde_json::json!({
"jsonrpc": "2.0",
"id": serde_json::Value::Null,
"error": JsonRpcError::internal_error(e.to_string()),
})
.to_string(),
),
};
ctx.terminal = terminal_json;
}
maybe = ctx.rx.recv(), if ctx.rx_open => {
match maybe {
Some(n) => {
if !ctx.subscriptions.publish(&n)
&& let Some(json) =
crate::transport::stdio::serialize_notification(&n)
{
ctx.queue.push_back(json);
}
}
None => ctx.rx_open = false,
}
ctx.call = Some(call);
}
}
}
});
Sse::new(stream)
.keep_alive(
axum::response::sse::KeepAlive::new()
.interval(Duration::from_secs(30))
.text("ping"),
)
.into_response()
}
#[cfg(feature = "stateless")]
fn listen_request_declares_tasks(parsed: &serde_json::Value) -> bool {
parsed
.get("params")
.and_then(crate::stateless::StatelessRequestMeta::from_params)
.and_then(|meta| meta.client_capabilities)
.and_then(|capabilities| capabilities.extensions)
.is_some_and(|declared| {
declared.contains_key(tower_mcp_types::protocol::TASKS_EXTENSION_ID)
})
}
#[cfg(feature = "stateless")]
async fn handle_modern_subscriptions_listen_sse(
state: Arc<AppState>,
parsed: &serde_json::Value,
) -> Response {
let id = extract_request_id(parsed);
let Some(subscription_id) = id.clone() else {
return json_rpc_error_response_with_status(
None,
JsonRpcError::invalid_request("subscriptions/listen requires a request id"),
StatusCode::BAD_REQUEST,
);
};
let params = match parsed
.get("params")
.cloned()
.ok_or_else(|| JsonRpcError::invalid_params("subscriptions/listen requires params"))
.and_then(|value| {
serde_json::from_value::<SubscriptionsListenParams>(value)
.map_err(|error| JsonRpcError::invalid_params(error.to_string()))
}) {
Ok(params) => params,
Err(error) => {
return json_rpc_error_response_with_status(id, error, StatusCode::BAD_REQUEST);
}
};
let Some(requested) = params.notifications else {
return json_rpc_error_response_with_status(
id,
JsonRpcError::invalid_params("subscriptions/listen requires a notifications filter"),
StatusCode::BAD_REQUEST,
);
};
if requested.task_ids.is_some() && !listen_request_declares_tasks(parsed) {
return json_rpc_error_response_with_status(
id,
JsonRpcError::missing_required_client_capability(
crate::router::tasks_client_capabilities(),
),
StatusCode::BAD_REQUEST,
);
}
let accepted = accepted_subscription_filter(requested, state.tasks_extension_enabled());
let (rx, guard) = state
.modern_subscriptions
.register(subscription_id.clone(), accepted.clone());
let acknowledgment = serde_json::json!({
"jsonrpc": "2.0",
"method": "notifications/subscriptions/acknowledged",
"params": {
"_meta": {
"io.modelcontextprotocol/subscriptionId": subscription_id
},
"notifications": accepted
}
})
.to_string();
struct ModernListenStream {
first: Option<String>,
rx: mpsc::UnboundedReceiver<String>,
_guard: ModernSubscriptionGuard,
}
let stream = futures::stream::unfold(
ModernListenStream {
first: Some(acknowledgment),
rx,
_guard: guard,
},
|mut state| async move {
let message = match state.first.take() {
Some(first) => Some(first),
None => state.rx.recv().await,
}?;
Some((
Ok::<_, Infallible>(Event::default().event(SSE_MESSAGE_EVENT).data(message)),
state,
))
},
);
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_static(PROTOCOL_VERSION_2026_07_28),
);
response
}
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()
}
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()
}
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()
}
}
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()
}
fn json_rpc_error_response(
id: Option<crate::protocol::RequestId>,
error: JsonRpcError,
) -> Response {
let response = JsonRpcResponse::error(id, error);
axum::Json(response).into_response()
}
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
}
#[cfg(test)]
mod tests {
use super::*;
use axum::body::Body;
use axum::http::Request;
use tower::ServiceExt;
fn create_test_router() -> McpRouter {
McpRouter::new().server_info("test-server", "1.0.0")
}
#[test]
fn final_result_fields_are_method_and_version_aware() {
for method in [
"server/discover",
"tools/list",
"prompts/list",
"resources/list",
"resources/read",
"resources/templates/list",
] {
let mut response =
JsonRpcResponse::result(RequestId::Number(1), serde_json::json!({"value": true}));
apply_protocol_result_fields(&mut response, method, PROTOCOL_VERSION_2026_07_28);
let json = serde_json::to_value(response).unwrap();
assert_eq!(json["result"]["resultType"], "complete", "{method}");
assert_eq!(json["result"]["ttlMs"], 0, "{method}");
assert_eq!(json["result"]["cacheScope"], "private", "{method}");
}
let mut ordinary =
JsonRpcResponse::result(RequestId::Number(1), serde_json::json!({"content": []}));
apply_protocol_result_fields(&mut ordinary, "tools/call", PROTOCOL_VERSION_2026_07_28);
let json = serde_json::to_value(ordinary).unwrap();
assert_eq!(json["result"]["resultType"], "complete");
assert!(json["result"].get("ttlMs").is_none());
assert!(json["result"].get("cacheScope").is_none());
}
#[test]
fn final_result_fields_preserve_explicit_values_and_legacy_wire_shape() {
let explicit = serde_json::json!({
"contents": [],
"ttlMs": 42,
"cacheScope": "public"
});
let mut response = JsonRpcResponse::result(RequestId::Number(1), explicit.clone());
apply_protocol_result_fields(&mut response, "resources/read", PROTOCOL_VERSION_2026_07_28);
let json = serde_json::to_value(response).unwrap();
assert_eq!(json["result"]["ttlMs"], 42);
assert_eq!(json["result"]["cacheScope"], "public");
for discriminator in ["input_required", "task"] {
let mut response = JsonRpcResponse::result(
RequestId::Number(1),
serde_json::json!({"resultType": discriminator}),
);
apply_protocol_result_fields(&mut response, "tools/call", PROTOCOL_VERSION_2026_07_28);
let json = serde_json::to_value(response).unwrap();
assert_eq!(json["result"]["resultType"], discriminator);
}
let mut legacy = JsonRpcResponse::result(RequestId::Number(1), explicit);
let before = serde_json::to_value(&legacy).unwrap();
apply_protocol_result_fields(&mut legacy, "resources/read", "2025-11-25");
assert_eq!(serde_json::to_value(legacy).unwrap(), before);
}
#[tokio::test]
#[cfg(feature = "stateless")]
async fn modern_subscription_registry_filters_and_tags_notifications() {
let registry = Arc::new(ModernSubscriptionRegistry::default());
let (mut rx, guard) = registry.register(
RequestId::String("listen-1".to_string()),
SubscriptionFilter {
tools_list_changed: Some(true),
..SubscriptionFilter::default()
},
);
assert!(registry.publish(&ServerNotification::PromptsListChanged));
assert!(matches!(
rx.try_recv(),
Err(mpsc::error::TryRecvError::Empty)
));
assert!(registry.publish(&ServerNotification::ToolsListChanged));
let message = rx.recv().await.expect("matching notification");
let json: serde_json::Value = serde_json::from_str(&message).unwrap();
assert_eq!(json["method"], "notifications/tools/list_changed");
assert_eq!(
json["params"]["_meta"]["io.modelcontextprotocol/subscriptionId"],
"listen-1"
);
drop(guard);
assert!(registry.subscriptions.lock().unwrap().is_empty());
}
#[tokio::test]
async fn runtime_protocol_allowlist_drives_discovery() {
let transport = HttpTransport::new(create_test_router())
.disable_origin_validation()
.protocol_versions(["2025-03-26"])
.unwrap();
let app = transport.into_router();
let request = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "server/discover"
})
.to_string(),
))
.unwrap();
let response = app.oneshot(request).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert_eq!(
json["result"]["supportedVersions"],
serde_json::json!(["2025-03-26"])
);
}
#[tokio::test]
async fn test_oversized_body_rejected_with_413() {
let transport = HttpTransport::new(create_test_router())
.disable_origin_validation()
.max_body_size(1024);
let app = transport.into_router();
let padding = "x".repeat(2048);
let body = format!(
r#"{{"jsonrpc":"2.0","id":1,"method":"ping","params":{{"pad":"{}"}}}}"#,
padding
);
let request = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.body(Body::from(body))
.unwrap();
let response = app.oneshot(request).await.unwrap();
assert_eq!(response.status(), StatusCode::PAYLOAD_TOO_LARGE);
}
#[tokio::test]
async fn test_oversized_content_length_rejected_without_reading() {
let transport = HttpTransport::new(create_test_router())
.disable_origin_validation()
.max_body_size(1024);
let app = transport.into_router();
let request = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.header("Content-Length", "10485760")
.body(Body::from(r#"{"jsonrpc":"2.0","id":1,"method":"ping"}"#))
.unwrap();
let response = app.oneshot(request).await.unwrap();
assert_eq!(response.status(), StatusCode::PAYLOAD_TOO_LARGE);
}
#[tokio::test]
async fn test_body_within_limit_accepted() {
let transport = HttpTransport::new(create_test_router())
.disable_origin_validation()
.max_body_size(1024);
let app = transport.into_router();
let request = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.body(Body::from(r#"{"jsonrpc":"2.0","id":1,"method":"ping"}"#))
.unwrap();
let response = app.oneshot(request).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
}
#[tokio::test]
async fn test_initialize_creates_session() {
let transport = HttpTransport::new(create_test_router()).disable_origin_validation();
let app = transport.into_router();
let request = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2025-11-25",
"capabilities": {},
"clientInfo": {
"name": "test-client",
"version": "1.0.0"
}
}
})
.to_string(),
))
.unwrap();
let response = app.oneshot(request).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
assert!(response.headers().contains_key(MCP_SESSION_ID_HEADER));
assert_eq!(
response
.headers()
.get(MCP_PROTOCOL_VERSION_HEADER)
.and_then(|v| v.to_str().ok()),
Some("2025-11-25")
);
}
#[tokio::test]
async fn test_protocol_version_header_on_subsequent_requests() {
let transport = HttpTransport::new(create_test_router()).disable_origin_validation();
let app = transport.into_router();
let init_request = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2025-03-26",
"capabilities": {},
"clientInfo": {
"name": "test-client",
"version": "1.0.0"
}
}
})
.to_string(),
))
.unwrap();
let init_response = app.clone().oneshot(init_request).await.unwrap();
let session_id = init_response
.headers()
.get(MCP_SESSION_ID_HEADER)
.unwrap()
.to_str()
.unwrap()
.to_string();
assert_eq!(
init_response
.headers()
.get(MCP_PROTOCOL_VERSION_HEADER)
.and_then(|v| v.to_str().ok()),
Some("2025-03-26")
);
let initialized_request = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.header(MCP_SESSION_ID_HEADER, &session_id)
.header(MCP_PROTOCOL_VERSION_HEADER, "2025-03-26")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"method": "notifications/initialized"
})
.to_string(),
))
.unwrap();
app.clone().oneshot(initialized_request).await.unwrap();
let list_request = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.header(MCP_SESSION_ID_HEADER, &session_id)
.header(MCP_PROTOCOL_VERSION_HEADER, "2025-03-26")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 2,
"method": "tools/list"
})
.to_string(),
))
.unwrap();
let response = app.oneshot(list_request).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(
response
.headers()
.get(MCP_PROTOCOL_VERSION_HEADER)
.and_then(|v| v.to_str().ok()),
Some("2025-03-26")
);
}
#[tokio::test]
async fn unsupported_protocol_version_returns_spec_shape_error() {
let transport = HttpTransport::new(create_test_router()).disable_origin_validation();
let app = transport.into_router();
let init_request = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2025-11-25",
"capabilities": {},
"clientInfo": { "name": "t", "version": "0" }
}
})
.to_string(),
))
.unwrap();
let init_response = app.clone().oneshot(init_request).await.unwrap();
let session_id = init_response
.headers()
.get(MCP_SESSION_ID_HEADER)
.unwrap()
.to_str()
.unwrap()
.to_string();
let bad = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json")
.header(MCP_SESSION_ID_HEADER, &session_id)
.header(MCP_PROTOCOL_VERSION_HEADER, "1999-01-01")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 99,
"method": "tools/list"
})
.to_string(),
))
.unwrap();
let response = app.oneshot(bad).await.unwrap();
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert_eq!(json["error"]["code"].as_i64().unwrap(), -32022);
assert_eq!(json["error"]["data"]["requested"], "1999-01-01");
let supported = json["error"]["data"]["supported"]
.as_array()
.expect("supported must be an array");
assert!(supported.contains(&serde_json::json!("2025-11-25")));
assert_eq!(json["id"], 99);
assert!(
json["error"]["data"].get("supportedVersions").is_none(),
"error data must use 'supported', not 'supportedVersions': {:?}",
json["error"]["data"]
);
let expected: Vec<serde_json::Value> = crate::COMPILED_PROTOCOL_VERSIONS
.iter()
.map(|v| serde_json::json!(v))
.collect();
assert_eq!(
supported, &expected,
"data.supported must exactly match COMPILED_PROTOCOL_VERSIONS"
);
}
#[cfg(feature = "stateless")]
#[tokio::test]
async fn stateless_unsupported_protocol_version_returns_spec_shape_error() {
let transport = HttpTransport::new(create_test_router()).disable_origin_validation();
let app = transport.into_router();
let req = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json")
.header(MCP_PROTOCOL_VERSION_HEADER, "2099-01-01")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 42,
"method": "tools/list"
})
.to_string(),
))
.unwrap();
let response = app.oneshot(req).await.unwrap();
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert_eq!(
json["error"]["code"].as_i64().unwrap(),
-32022,
"must return UnsupportedProtocolVersion (-32022): {json}"
);
assert_eq!(
json["error"]["data"]["requested"], "2099-01-01",
"data.requested must echo the version: {json}"
);
assert!(
json["error"]["data"].get("supportedVersions").is_none(),
"error data must use 'supported', not 'supportedVersions': {json}"
);
let supported = json["error"]["data"]["supported"]
.as_array()
.expect("data.supported must be an array");
let expected: Vec<serde_json::Value> = crate::COMPILED_PROTOCOL_VERSIONS
.iter()
.map(|v| serde_json::json!(v))
.collect();
assert_eq!(
supported, &expected,
"data.supported must exactly match COMPILED_PROTOCOL_VERSIONS"
);
}
#[tokio::test]
async fn sep_2243_lenient_mode_accepts_missing_headers() {
let transport = HttpTransport::new(create_test_router()).disable_origin_validation();
let app = transport.into_router();
let init = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2025-11-25",
"capabilities": {},
"clientInfo": { "name": "t", "version": "0" }
}
})
.to_string(),
))
.unwrap();
let init_response = app.clone().oneshot(init).await.unwrap();
assert_eq!(init_response.status(), StatusCode::OK);
let session_id = init_response
.headers()
.get(MCP_SESSION_ID_HEADER)
.unwrap()
.to_str()
.unwrap()
.to_string();
let req = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.header(MCP_SESSION_ID_HEADER, &session_id)
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 2,
"method": "tools/list"
})
.to_string(),
))
.unwrap();
let response = app.oneshot(req).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
}
#[tokio::test]
async fn sep_2243_lenient_mode_validates_present_headers() {
let transport = HttpTransport::new(create_test_router()).disable_origin_validation();
let app = transport.into_router();
let init = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.header(MCP_METHOD_HEADER, "initialize")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2025-11-25",
"capabilities": {},
"clientInfo": { "name": "t", "version": "0" }
}
})
.to_string(),
))
.unwrap();
let init_response = app.clone().oneshot(init).await.unwrap();
assert_eq!(init_response.status(), StatusCode::OK);
let session_id = init_response
.headers()
.get(MCP_SESSION_ID_HEADER)
.unwrap()
.to_str()
.unwrap()
.to_string();
let req = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json")
.header(MCP_SESSION_ID_HEADER, &session_id)
.header(MCP_METHOD_HEADER, "ping")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 2,
"method": "tools/list"
})
.to_string(),
))
.unwrap();
let response = app.oneshot(req).await.unwrap();
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert_eq!(json["error"]["code"].as_i64().unwrap(), -32020);
assert!(
json["error"]["message"]
.as_str()
.unwrap()
.contains("Mcp-Method")
);
assert_eq!(json["id"], 2);
}
#[tokio::test]
#[cfg(not(feature = "stateless"))]
async fn sep_2243_strict_mode_tools_call_with_matching_headers() {
use crate::{CallToolResult, ToolBuilder};
let router = McpRouter::new().server_info("t", "1.0.0").tool(
ToolBuilder::new("echo")
.description("echo")
.handler(|args: serde_json::Value| async move {
Ok(CallToolResult::text(args.to_string()))
})
.build(),
);
let transport = HttpTransport::new(router).disable_origin_validation();
let app = transport.into_router();
let init = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.header(MCP_METHOD_HEADER, "initialize")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2026-07-28",
"capabilities": {},
"clientInfo": { "name": "t", "version": "0" }
}
})
.to_string(),
))
.unwrap();
let init_response = app.clone().oneshot(init).await.unwrap();
assert_eq!(init_response.status(), StatusCode::OK);
let session_id = init_response
.headers()
.get(MCP_SESSION_ID_HEADER)
.unwrap()
.to_str()
.unwrap()
.to_string();
let negotiated_version = init_response
.headers()
.get(MCP_PROTOCOL_VERSION_HEADER)
.unwrap()
.to_str()
.unwrap()
.to_string();
let req = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json")
.header(MCP_SESSION_ID_HEADER, &session_id)
.header(MCP_PROTOCOL_VERSION_HEADER, &negotiated_version)
.header(MCP_METHOD_HEADER, "tools/call")
.header(MCP_NAME_HEADER, "echo")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 2,
"method": "tools/call",
"params": {
"name": "echo",
"arguments": {"message": "hi"}
}
})
.to_string(),
))
.unwrap();
let response = app.oneshot(req).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
}
#[tokio::test]
async fn sep_2243_tools_call_mcp_name_mismatch_rejected() {
use crate::{CallToolResult, ToolBuilder};
let router = McpRouter::new().server_info("t", "1.0.0").tool(
ToolBuilder::new("echo")
.description("echo")
.handler(|args: serde_json::Value| async move {
Ok(CallToolResult::text(args.to_string()))
})
.build(),
);
let transport = HttpTransport::new(router).disable_origin_validation();
let app = transport.into_router();
let init = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2025-11-25",
"capabilities": {},
"clientInfo": { "name": "t", "version": "0" }
}
})
.to_string(),
))
.unwrap();
let init_response = app.clone().oneshot(init).await.unwrap();
let session_id = init_response
.headers()
.get(MCP_SESSION_ID_HEADER)
.unwrap()
.to_str()
.unwrap()
.to_string();
let req = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json")
.header(MCP_SESSION_ID_HEADER, &session_id)
.header(MCP_METHOD_HEADER, "tools/call")
.header(MCP_NAME_HEADER, "not-echo")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 7,
"method": "tools/call",
"params": {
"name": "echo",
"arguments": {"message": "hi"}
}
})
.to_string(),
))
.unwrap();
let response = app.oneshot(req).await.unwrap();
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert_eq!(json["error"]["code"].as_i64().unwrap(), -32020);
let msg = json["error"]["message"].as_str().unwrap();
assert!(msg.contains("Mcp-Name"), "got: {msg}");
assert_eq!(json["id"], 7);
}
#[tokio::test]
async fn sep_2243_mcp_param_base64_decoded_and_matched() {
use crate::{CallToolResult, ToolBuilder};
let router = McpRouter::new().server_info("t", "1.0.0").tool(
ToolBuilder::new("echo")
.description("echo")
.handler(|args: serde_json::Value| async move {
Ok(CallToolResult::text(args.to_string()))
})
.build(),
);
let transport = HttpTransport::new(router).disable_origin_validation();
let app = transport.into_router();
let init = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2025-11-25",
"capabilities": {},
"clientInfo": { "name": "t", "version": "0" }
}
})
.to_string(),
))
.unwrap();
let init_response = app.clone().oneshot(init).await.unwrap();
let session_id = init_response
.headers()
.get(MCP_SESSION_ID_HEADER)
.unwrap()
.to_str()
.unwrap()
.to_string();
let req = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json")
.header(MCP_SESSION_ID_HEADER, &session_id)
.header(MCP_METHOD_HEADER, "tools/call")
.header(MCP_NAME_HEADER, "echo")
.header("mcp-param-message", "=?base64?SGVsbG8=?=")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 5,
"method": "tools/call",
"params": {
"name": "echo",
"arguments": {"message": "Hello"}
}
})
.to_string(),
))
.unwrap();
let response = app.oneshot(req).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
}
#[cfg(feature = "stateless")]
#[tokio::test]
async fn sep_2243_final_request_requires_schema_annotated_header() {
use crate::extract::RawArgs;
use crate::{CallToolResult, ToolBuilder};
let tool = ToolBuilder::new("route")
.input_schema(serde_json::json!({
"type": "object",
"properties": {
"tenant_id": {
"type": "string",
"x-mcp-header": "Tenant"
}
}
}))
.extractor_handler((), |RawArgs(args): RawArgs| async move {
Ok(CallToolResult::text(args["tenant_id"].to_string()))
})
.build();
let app = HttpTransport::new(
McpRouter::new()
.server_info("header-test", "1.0.0")
.tool(tool),
)
.disable_origin_validation()
.into_router();
let request = |custom_header: Option<&'static str>| {
let mut builder = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header(MCP_PROTOCOL_VERSION_HEADER, PROTOCOL_VERSION_2026_07_28)
.header(MCP_METHOD_HEADER, "tools/call")
.header(MCP_NAME_HEADER, "route");
if let Some(value) = custom_header {
builder = builder.header("Mcp-Param-Tenant", value);
}
builder
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 9,
"method": "tools/call",
"params": {
"name": "route",
"arguments": {"tenant_id": "acme"},
"_meta": {
"io.modelcontextprotocol/protocolVersion":
PROTOCOL_VERSION_2026_07_28,
"io.modelcontextprotocol/clientCapabilities": {}
}
}
})
.to_string(),
))
.unwrap()
};
let response = app.clone().oneshot(request(None)).await.unwrap();
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
let error: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert_eq!(error["id"], 9);
assert_eq!(error["error"]["code"], -32020);
let response = app.oneshot(request(Some("acme"))).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
}
#[tokio::test]
async fn sep_2243_notification_with_matching_method_header_accepted() {
let transport = HttpTransport::new(create_test_router()).disable_origin_validation();
let app = transport.into_router();
let init = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2025-11-25",
"capabilities": {},
"clientInfo": { "name": "t", "version": "0" }
}
})
.to_string(),
))
.unwrap();
let init_response = app.clone().oneshot(init).await.unwrap();
let session_id = init_response
.headers()
.get(MCP_SESSION_ID_HEADER)
.unwrap()
.to_str()
.unwrap()
.to_string();
let req = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.header(MCP_SESSION_ID_HEADER, &session_id)
.header(MCP_METHOD_HEADER, "notifications/initialized")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"method": "notifications/initialized"
})
.to_string(),
))
.unwrap();
let response = app.oneshot(req).await.unwrap();
assert_eq!(response.status(), StatusCode::ACCEPTED);
}
#[tokio::test]
async fn test_request_without_session_fails() {
let transport = HttpTransport::new(create_test_router())
.disable_origin_validation()
.require_sessions();
let app = transport.into_router();
let request = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "tools/list"
})
.to_string(),
))
.unwrap();
let response = app.oneshot(request).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert!(json.get("error").is_some());
assert_eq!(json["error"]["code"], -32006); }
#[tokio::test]
async fn test_delete_session() {
let transport = HttpTransport::new(create_test_router()).disable_origin_validation();
let app = transport.into_router();
let init_request = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2025-11-25",
"capabilities": {},
"clientInfo": {
"name": "test-client",
"version": "1.0.0"
}
}
})
.to_string(),
))
.unwrap();
let response = app.clone().oneshot(init_request).await.unwrap();
let session_id = response
.headers()
.get(MCP_SESSION_ID_HEADER)
.unwrap()
.to_str()
.unwrap()
.to_string();
let delete_request = Request::builder()
.method("DELETE")
.uri("/")
.header(MCP_SESSION_ID_HEADER, &session_id)
.body(Body::empty())
.unwrap();
let response = app.clone().oneshot(delete_request).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let list_request = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header(MCP_SESSION_ID_HEADER, &session_id)
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 2,
"method": "tools/list"
})
.to_string(),
))
.unwrap();
let response = app.oneshot(list_request).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert!(json.get("error").is_some());
assert_eq!(json["error"]["code"], -32005); }
#[tokio::test]
async fn test_custom_session_store_receives_create_and_delete() {
use crate::session_store::{MemorySessionStore, SessionStore as PublicSessionStore};
let store = Arc::new(MemorySessionStore::new());
let store_dyn: Arc<dyn PublicSessionStore> = store.clone();
let transport = HttpTransport::new(create_test_router())
.disable_origin_validation()
.session_store(store_dyn);
let (app, handle) = transport.into_router_with_handle();
let init_request = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2025-11-25",
"capabilities": {},
"clientInfo": { "name": "test-client", "version": "1.0.0" }
}
})
.to_string(),
))
.unwrap();
let response = app.clone().oneshot(init_request).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let session_id = response
.headers()
.get(MCP_SESSION_ID_HEADER)
.unwrap()
.to_str()
.unwrap()
.to_string();
assert_eq!(store.len().await, 1);
let record = store
.load(&session_id)
.await
.unwrap()
.expect("expected session to be persisted");
assert_eq!(record.id, session_id);
let client_info = record
.client_info
.expect("client_info should be populated after initialize");
assert_eq!(client_info.name, "test-client");
assert_eq!(client_info.version, "1.0.0");
assert!(
record.client_capabilities.is_some(),
"client_capabilities should be populated after initialize"
);
assert!(handle.terminate_session(&session_id).await);
assert_eq!(store.len().await, 0);
assert!(store.load(&session_id).await.unwrap().is_none());
}
#[tokio::test]
async fn test_session_store_record_carries_negotiated_protocol_version() {
use crate::session_store::{MemorySessionStore, SessionStore as PublicSessionStore};
let store = Arc::new(MemorySessionStore::new());
let store_dyn: Arc<dyn PublicSessionStore> = store.clone();
let transport = HttpTransport::new(create_test_router())
.disable_origin_validation()
.session_store(store_dyn);
let app = transport.into_router();
let init_request = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2025-03-26",
"capabilities": {},
"clientInfo": { "name": "v-client", "version": "2.0.0" }
}
})
.to_string(),
))
.unwrap();
let response = app.oneshot(init_request).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let session_id = response
.headers()
.get(MCP_SESSION_ID_HEADER)
.unwrap()
.to_str()
.unwrap()
.to_string();
let record = store
.load(&session_id)
.await
.unwrap()
.expect("session should be persisted");
assert_eq!(record.protocol_version, "2025-03-26");
let client_info = record.client_info.expect("client_info should be populated");
assert_eq!(client_info.name, "v-client");
}
#[tokio::test]
async fn test_restored_session_exposes_original_client_info() {
use crate::session_store::{MemorySessionStore, SessionStore as PublicSessionStore};
let store = Arc::new(MemorySessionStore::new());
let store_dyn: Arc<dyn PublicSessionStore> = store.clone();
let session_id = {
let transport = HttpTransport::new(create_test_router())
.disable_origin_validation()
.session_store(store_dyn.clone());
let app = transport.into_router();
let init_request = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2025-11-25",
"capabilities": { "roots": {} },
"clientInfo": {
"name": "original-client",
"version": "3.1.4"
}
}
})
.to_string(),
))
.unwrap();
let response = app.oneshot(init_request).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
response
.headers()
.get(MCP_SESSION_ID_HEADER)
.unwrap()
.to_str()
.unwrap()
.to_string()
};
let stored = store
.load(&session_id)
.await
.unwrap()
.expect("record should survive transport drop");
assert_eq!(
stored.client_info.as_ref().map(|c| c.name.as_str()),
Some("original-client")
);
let transport2 = HttpTransport::new(create_test_router())
.disable_origin_validation()
.session_store(store_dyn);
let app2 = transport2.into_router();
let list_request = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header(MCP_SESSION_ID_HEADER, &session_id)
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "tools/list"
})
.to_string(),
))
.unwrap();
let response = app2.oneshot(list_request).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert!(
json.get("result").is_some(),
"expected tools/list result, got {json}"
);
let after_restore = store
.load(&session_id)
.await
.unwrap()
.expect("record should still be present after restore");
let client_info = after_restore
.client_info
.expect("restored record should retain client_info");
assert_eq!(client_info.name, "original-client");
assert_eq!(client_info.version, "3.1.4");
assert!(
after_restore.client_capabilities.is_some(),
"restored record should retain client_capabilities"
);
}
#[tokio::test]
async fn test_auto_reinitialize_marks_synthetic_client_info() {
use crate::session_store::{MemorySessionStore, SessionStore as PublicSessionStore};
let store = Arc::new(MemorySessionStore::new());
let store_dyn: Arc<dyn PublicSessionStore> = store.clone();
let transport = HttpTransport::new(create_test_router())
.disable_origin_validation()
.session_store(store_dyn)
.auto_reinitialize_sessions(true);
let app = transport.into_router();
let list_request = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header(MCP_SESSION_ID_HEADER, "made-up-id")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "tools/list"
})
.to_string(),
))
.unwrap();
let response = app.oneshot(list_request).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let record = store
.load("made-up-id")
.await
.unwrap()
.expect("auto-reinitialize should persist a record");
assert_eq!(
record.client_info.as_ref().map(|c| c.name.as_str()),
Some("auto-recovered")
);
}
#[tokio::test]
async fn test_custom_event_store_buffers_and_purges() {
use crate::event_store::{EventStore as PublicEventStore, MemoryEventStore};
let events = Arc::new(MemoryEventStore::new());
let events_dyn: Arc<dyn PublicEventStore> = events.clone();
let session = Arc::new(Session::new(
create_test_router(),
false,
identity_factory(),
events_dyn,
));
session.buffer_event(0, "first".to_string()).await;
session.buffer_event(1, "second".to_string()).await;
assert_eq!(events.total_events().await, 2);
let replayed = events.replay_after(&session.id, 0).await.unwrap();
assert_eq!(replayed.len(), 1);
assert_eq!(replayed[0].id, 1);
assert_eq!(replayed[0].data, "second");
events.purge_session(&session.id).await.unwrap();
assert_eq!(events.total_events().await, 0);
}
#[tokio::test]
async fn test_restore_from_store_serves_unknown_session_id() {
use crate::session_store::{MemorySessionStore, SessionRecord, SessionStore};
let store = Arc::new(MemorySessionStore::new());
let store_dyn: Arc<dyn SessionStore> = store.clone();
let mut seeded = SessionRecord::new(
"shared-session".to_string(),
"2025-11-25".to_string(),
Duration::from_secs(60),
);
store.create(&mut seeded).await.unwrap();
let seeded_id = seeded.id;
let transport = HttpTransport::new(create_test_router())
.disable_origin_validation()
.session_store(store_dyn);
let app = transport.into_router();
let list_request = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header(MCP_SESSION_ID_HEADER, &seeded_id)
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "tools/list"
})
.to_string(),
))
.unwrap();
let response = app.oneshot(list_request).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert!(
json.get("result").is_some(),
"expected tools/list result, got {json}"
);
}
#[tokio::test]
async fn test_auto_reinitialize_serves_unknown_session_without_store_record() {
let transport = HttpTransport::new(create_test_router())
.disable_origin_validation()
.auto_reinitialize_sessions(true);
let app = transport.into_router();
let list_request = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header(MCP_SESSION_ID_HEADER, "client-made-up-id")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "tools/list"
})
.to_string(),
))
.unwrap();
let response = app.oneshot(list_request).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert!(
json.get("result").is_some(),
"expected tools/list result, got {json}"
);
}
#[tokio::test]
async fn test_unknown_session_without_restore_or_auto_reinit_returns_error() {
let transport = HttpTransport::new(create_test_router()).disable_origin_validation();
let app = transport.into_router();
let list_request = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header(MCP_SESSION_ID_HEADER, "never-seen-before")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "tools/list"
})
.to_string(),
))
.unwrap();
let response = app.oneshot(list_request).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert!(json.get("error").is_some(), "expected error, got {json}");
assert_eq!(json["error"]["code"], -32005); }
#[tokio::test]
async fn test_session_expiration() {
let config = SessionConfig::with_ttl(Duration::from_millis(50))
.cleanup_interval(Duration::from_millis(10));
let transport = HttpTransport::new(create_test_router())
.disable_origin_validation()
.session_config(config);
let app = transport.into_router();
let init_request = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2025-11-25",
"capabilities": {},
"clientInfo": {
"name": "test-client",
"version": "1.0.0"
}
}
})
.to_string(),
))
.unwrap();
let response = app.clone().oneshot(init_request).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let session_id = response
.headers()
.get(MCP_SESSION_ID_HEADER)
.unwrap()
.to_str()
.unwrap()
.to_string();
tokio::time::sleep(Duration::from_millis(100)).await;
let list_request = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header(MCP_SESSION_ID_HEADER, &session_id)
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 2,
"method": "tools/list"
})
.to_string(),
))
.unwrap();
let response = app.oneshot(list_request).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert!(json.get("error").is_some());
assert_eq!(json["error"]["code"], -32005); }
#[tokio::test]
async fn test_layer_with_identity() {
let transport = HttpTransport::new(create_test_router())
.disable_origin_validation()
.layer(tower::layer::util::Identity::new());
let app = transport.into_router();
let request = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2025-11-25",
"capabilities": {},
"clientInfo": {
"name": "test-client",
"version": "1.0.0"
}
}
})
.to_string(),
))
.unwrap();
let response = app.oneshot(request).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
assert!(response.headers().contains_key(MCP_SESSION_ID_HEADER));
}
#[tokio::test]
async fn test_layer_with_timeout() {
use std::time::Duration;
use tower::timeout::TimeoutLayer;
let transport = HttpTransport::new(create_test_router())
.disable_origin_validation()
.layer(TimeoutLayer::new(Duration::from_secs(30)));
let app = transport.into_router();
let request = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2025-11-25",
"capabilities": {},
"clientInfo": {
"name": "test-client",
"version": "1.0.0"
}
}
})
.to_string(),
))
.unwrap();
let response = app.oneshot(request).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
assert!(response.headers().contains_key(MCP_SESSION_ID_HEADER));
}
#[tokio::test]
async fn test_layer_middleware_error_produces_jsonrpc_error() {
use std::time::Duration;
use tower::timeout::TimeoutLayer;
let slow_tool = crate::tool::ToolBuilder::new("slow")
.description("A slow tool")
.handler(|_: serde_json::Value| async move {
tokio::time::sleep(Duration::from_secs(10)).await;
Ok(crate::CallToolResult::text("done"))
})
.build();
let router = McpRouter::new()
.server_info("test-server", "1.0.0")
.tool(slow_tool);
let transport = HttpTransport::new(router)
.disable_origin_validation()
.layer(TimeoutLayer::new(Duration::from_millis(1)));
let app = transport.into_router();
let init_request = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2025-11-25",
"capabilities": {},
"clientInfo": {
"name": "test-client",
"version": "1.0.0"
}
}
})
.to_string(),
))
.unwrap();
let response = app.clone().oneshot(init_request).await.unwrap();
let session_id = response
.headers()
.get(MCP_SESSION_ID_HEADER)
.unwrap()
.to_str()
.unwrap()
.to_string();
let tool_request = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header(MCP_SESSION_ID_HEADER, &session_id)
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 2,
"method": "tools/call",
"params": {
"name": "slow",
"arguments": {}
}
})
.to_string(),
))
.unwrap();
let response = app.oneshot(tool_request).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert!(
json.get("error").is_some(),
"Expected JSON-RPC error response, got: {}",
json
);
}
#[tokio::test]
async fn test_max_sessions_limit() {
let config = SessionConfig::default().max_sessions(1);
let transport = HttpTransport::new(create_test_router())
.disable_origin_validation()
.session_config(config);
let app = transport.into_router();
let init_request1 = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2025-11-25",
"capabilities": {},
"clientInfo": {
"name": "test-client",
"version": "1.0.0"
}
}
})
.to_string(),
))
.unwrap();
let response = app.clone().oneshot(init_request1).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let init_request2 = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 2,
"method": "initialize",
"params": {
"protocolVersion": "2025-11-25",
"capabilities": {},
"clientInfo": {
"name": "test-client-2",
"version": "1.0.0"
}
}
})
.to_string(),
))
.unwrap();
let response = app.oneshot(init_request2).await.unwrap();
assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE);
}
#[tokio::test]
async fn test_session_event_buffering() {
let session = Session::new(
create_test_router(),
false,
identity_factory(),
Arc::new(crate::event_store::MemoryEventStore::new()),
);
session.buffer_event(0, "event0".to_string()).await;
session.buffer_event(1, "event1".to_string()).await;
session.buffer_event(2, "event2".to_string()).await;
let events = session.get_events_after(0).await;
assert_eq!(events.len(), 2);
assert_eq!(events[0].id, 1);
assert_eq!(events[0].data, "event1");
assert_eq!(events[1].id, 2);
assert_eq!(events[1].data, "event2");
let events = session.get_events_after(1).await;
assert_eq!(events.len(), 1);
assert_eq!(events[0].id, 2);
let events = session.get_events_after(2).await;
assert!(events.is_empty());
}
#[tokio::test]
async fn test_session_event_counter_increments() {
let session = Session::new(
create_test_router(),
false,
identity_factory(),
Arc::new(crate::event_store::MemoryEventStore::new()),
);
assert_eq!(session.next_event_id(), 0);
assert_eq!(session.next_event_id(), 1);
assert_eq!(session.next_event_id(), 2);
}
#[tokio::test]
async fn test_session_event_buffer_limit() {
let session = Session::new(
create_test_router(),
false,
identity_factory(),
Arc::new(crate::event_store::MemoryEventStore::new()),
);
for i in 0..10 {
session.buffer_event(i, format!("event{}", i)).await;
}
let events = session.get_events_after(0).await;
assert_eq!(events.len(), 9);
}
#[tokio::test]
async fn test_session_handle_count() {
let transport = HttpTransport::new(create_test_router()).disable_origin_validation();
let (app, handle) = transport.into_router_with_handle();
assert_eq!(handle.session_count().await, 0);
let request = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2025-11-25",
"capabilities": {},
"clientInfo": {
"name": "test-client",
"version": "1.0.0"
}
}
})
.to_string(),
))
.unwrap();
let response = app.oneshot(request).await.unwrap();
assert_eq!(response.status(), 200);
assert_eq!(handle.session_count().await, 1);
}
#[tokio::test]
async fn test_session_handle_list_and_terminate() {
let transport = HttpTransport::new(create_test_router()).disable_origin_validation();
let (app, handle) = transport.into_router_with_handle();
assert!(handle.list_sessions().await.is_empty());
let request = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2025-11-25",
"capabilities": {},
"clientInfo": {
"name": "test-client",
"version": "1.0.0"
}
}
})
.to_string(),
))
.unwrap();
let response = app.oneshot(request).await.unwrap();
assert_eq!(response.status(), 200);
let sessions = handle.list_sessions().await;
assert_eq!(sessions.len(), 1);
assert!(!sessions[0].id.is_empty());
let session_id = sessions[0].id.clone();
assert!(handle.terminate_session(&session_id).await);
assert_eq!(handle.session_count().await, 0);
assert!(!handle.terminate_session(&session_id).await);
}
#[tokio::test]
async fn test_request_without_session_id_rejected() {
let transport = HttpTransport::new(create_test_router())
.disable_origin_validation()
.require_sessions();
let app = transport.into_router();
let request = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "tools/list",
"params": {}
})
.to_string(),
))
.unwrap();
let response = app.oneshot(request).await.unwrap();
assert_eq!(response.status(), StatusCode::OK); let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert!(json["error"].is_object());
}
#[tokio::test]
async fn test_invalid_session_id_returns_error() {
let transport = HttpTransport::new(create_test_router()).disable_origin_validation();
let app = transport.into_router();
let request = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json")
.header("mcp-session-id", "nonexistent-session-id")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "tools/list",
"params": {}
})
.to_string(),
))
.unwrap();
let response = app.oneshot(request).await.unwrap();
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert_eq!(json["error"]["code"].as_i64().unwrap(), -32005); }
#[tokio::test]
async fn test_notification_returns_accepted() {
let transport = HttpTransport::new(create_test_router()).disable_origin_validation();
let app = transport.into_router();
let init_req = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2025-11-25",
"capabilities": {},
"clientInfo": { "name": "test", "version": "1.0" }
}
})
.to_string(),
))
.unwrap();
let resp = app.clone().oneshot(init_req).await.unwrap();
let session_id = resp
.headers()
.get(MCP_SESSION_ID_HEADER)
.unwrap()
.to_str()
.unwrap()
.to_string();
let notif = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("mcp-session-id", &session_id)
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"method": "notifications/initialized"
})
.to_string(),
))
.unwrap();
let response = app.oneshot(notif).await.unwrap();
assert_eq!(response.status(), StatusCode::ACCEPTED);
}
#[tokio::test]
async fn test_invalid_json_returns_parse_error() {
let transport = HttpTransport::new(create_test_router()).disable_origin_validation();
let app = transport.into_router();
let request = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json")
.body(Body::from("not valid json{{{"))
.unwrap();
let response = app.oneshot(request).await.unwrap();
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&body).unwrap();
tower_mcp_types::testing::assert_jsonrpc_error_response(&json);
assert!(
json["id"].is_null(),
"id must be null on parse error: {json}"
);
assert_eq!(json["error"]["code"].as_i64().unwrap(), -32700);
}
#[tokio::test]
async fn test_session_config_max_sessions() {
let transport = HttpTransport::new(create_test_router())
.disable_origin_validation()
.session_config(SessionConfig::default().max_sessions(1));
let app = transport.into_router();
let init1 = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2025-11-25",
"capabilities": {},
"clientInfo": { "name": "test1", "version": "1.0" }
}
})
.to_string(),
))
.unwrap();
let resp1 = app.clone().oneshot(init1).await.unwrap();
assert_eq!(resp1.status(), StatusCode::OK);
let init2 = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 2,
"method": "initialize",
"params": {
"protocolVersion": "2025-11-25",
"capabilities": {},
"clientInfo": { "name": "test2", "version": "1.0" }
}
})
.to_string(),
))
.unwrap();
let resp2 = app.oneshot(init2).await.unwrap();
assert_eq!(resp2.status(), StatusCode::SERVICE_UNAVAILABLE);
}
#[tokio::test]
async fn test_delete_terminates_session() {
let transport = HttpTransport::new(create_test_router()).disable_origin_validation();
let app = transport.into_router();
let init_req = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2025-11-25",
"capabilities": {},
"clientInfo": { "name": "test", "version": "1.0" }
}
})
.to_string(),
))
.unwrap();
let resp = app.clone().oneshot(init_req).await.unwrap();
let session_id = resp
.headers()
.get(MCP_SESSION_ID_HEADER)
.unwrap()
.to_str()
.unwrap()
.to_string();
let delete_req = Request::builder()
.method("DELETE")
.uri("/")
.header("mcp-session-id", &session_id)
.body(Body::empty())
.unwrap();
let resp = app.clone().oneshot(delete_req).await.unwrap();
assert!(resp.status().is_success());
let list_req = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json")
.header("mcp-session-id", &session_id)
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 2,
"method": "tools/list",
"params": {}
})
.to_string(),
))
.unwrap();
let resp = app.oneshot(list_req).await.unwrap();
let body = axum::body::to_bytes(resp.into_body(), usize::MAX)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert_eq!(json["error"]["code"].as_i64().unwrap(), -32005);
}
#[test]
fn test_is_localhost_origin_http() {
assert!(is_localhost_origin("http://localhost"));
assert!(is_localhost_origin("http://localhost:3000"));
assert!(is_localhost_origin("http://127.0.0.1"));
assert!(is_localhost_origin("http://127.0.0.1:8080"));
assert!(is_localhost_origin("http://[::1]"));
assert!(is_localhost_origin("http://[::1]:3000"));
}
#[test]
fn test_is_localhost_origin_https() {
assert!(is_localhost_origin("https://localhost"));
assert!(is_localhost_origin("https://127.0.0.1:443"));
}
#[test]
fn test_is_not_localhost_origin() {
assert!(!is_localhost_origin("http://example.com"));
assert!(!is_localhost_origin("http://evil-localhost.com"));
assert!(!is_localhost_origin("http://localhost.evil.com"));
assert!(!is_localhost_origin("ftp://localhost"));
assert!(!is_localhost_origin("localhost"));
assert!(!is_localhost_origin(""));
}
#[tokio::test]
async fn test_origin_validation_rejects_cross_origin() {
let transport = HttpTransport::new(create_test_router());
let app = transport.into_router();
let req = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.header("Origin", "http://evil.com")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2025-11-25",
"capabilities": {},
"clientInfo": { "name": "test", "version": "1.0" }
}
})
.to_string(),
))
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::FORBIDDEN);
}
#[tokio::test]
async fn test_origin_validation_allows_localhost() {
let transport = HttpTransport::new(create_test_router());
let app = transport.into_router();
let req = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.header("Origin", "http://localhost:3000")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2025-11-25",
"capabilities": {},
"clientInfo": { "name": "test", "version": "1.0" }
}
})
.to_string(),
))
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
}
#[tokio::test]
async fn test_origin_validation_allows_configured_origin() {
let transport = HttpTransport::new(create_test_router())
.allowed_origins(vec!["https://my-app.example.com".to_string()]);
let app = transport.into_router();
let req = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.header("Origin", "https://my-app.example.com")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2025-11-25",
"capabilities": {},
"clientInfo": { "name": "test", "version": "1.0" }
}
})
.to_string(),
))
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
}
#[tokio::test]
async fn test_origin_validation_rejects_unconfigured_origin() {
let transport = HttpTransport::new(create_test_router())
.allowed_origins(vec!["https://my-app.example.com".to_string()]);
let app = transport.into_router();
let req = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.header("Origin", "https://other-app.example.com")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2025-11-25",
"capabilities": {},
"clientInfo": { "name": "test", "version": "1.0" }
}
})
.to_string(),
))
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::FORBIDDEN);
}
#[tokio::test]
async fn test_origin_validation_no_header_allowed() {
let transport = HttpTransport::new(create_test_router());
let app = transport.into_router();
let req = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2025-11-25",
"capabilities": {},
"clientInfo": { "name": "test", "version": "1.0" }
}
})
.to_string(),
))
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
}
#[tokio::test]
async fn test_disabled_origin_validation_allows_any() {
let transport = HttpTransport::new(create_test_router()).disable_origin_validation();
let app = transport.into_router();
let req = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.header("Origin", "http://evil.com")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2025-11-25",
"capabilities": {},
"clientInfo": { "name": "test", "version": "1.0" }
}
})
.to_string(),
))
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
}
fn initialize_body() -> Body {
Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2025-11-25",
"capabilities": {},
"clientInfo": { "name": "test", "version": "1.0" }
}
})
.to_string(),
)
}
#[test]
fn test_is_localhost_host_variants() {
assert!(is_localhost_host("localhost"));
assert!(is_localhost_host("localhost:3000"));
assert!(is_localhost_host("127.0.0.1"));
assert!(is_localhost_host("127.0.0.1:8080"));
assert!(is_localhost_host("[::1]"));
assert!(is_localhost_host("[::1]:3000"));
assert!(!is_localhost_host("evil.com"));
assert!(!is_localhost_host("api.example.com:8443"));
assert!(!is_localhost_host("10.0.0.1"));
}
#[tokio::test]
async fn test_host_validation_allows_localhost() {
let transport = HttpTransport::new(create_test_router())
.allowed_hosts(vec!["api.example.com".to_string()]);
let app = transport.into_router();
let req = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.header("Host", "127.0.0.1:3000")
.body(initialize_body())
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
}
#[tokio::test]
async fn test_host_validation_allows_configured_host() {
let transport = HttpTransport::new(create_test_router())
.allowed_hosts(vec!["api.example.com".to_string()]);
let app = transport.into_router();
let req = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.header("Host", "api.example.com")
.body(initialize_body())
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
}
#[tokio::test]
async fn test_host_validation_rejects_unconfigured_host() {
let transport = HttpTransport::new(create_test_router())
.allowed_hosts(vec!["api.example.com".to_string()]);
let app = transport.into_router();
let req = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.header("Host", "evil.com")
.body(initialize_body())
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::BAD_REQUEST);
}
#[tokio::test]
async fn test_host_validation_no_allowlist_accepts_any_host() {
let transport = HttpTransport::new(create_test_router());
let app = transport.into_router();
let req = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.header("Host", "any.example.com")
.body(initialize_body())
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
}
#[tokio::test]
async fn test_disabled_host_validation_allows_any_with_allowlist() {
let transport = HttpTransport::new(create_test_router())
.disable_host_validation()
.allowed_hosts(vec!["api.example.com".to_string()]);
let app = transport.into_router();
let req = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.header("Host", "evil.com")
.body(initialize_body())
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
}
#[test]
fn test_effective_host_prefers_header() {
let mut headers = HeaderMap::new();
headers.insert(header::HOST, HeaderValue::from_static("api.example.com"));
let uri: axum::http::Uri = "http://other.example.com/path".parse().unwrap();
assert_eq!(effective_host(&headers, &uri), Some("api.example.com"));
}
#[test]
fn test_effective_host_falls_back_to_authority() {
let headers = HeaderMap::new();
let uri: axum::http::Uri = "http://api.example.com/path".parse().unwrap();
assert_eq!(effective_host(&headers, &uri), Some("api.example.com"));
}
#[test]
fn test_effective_host_returns_none_when_both_missing() {
let headers = HeaderMap::new();
let uri: axum::http::Uri = "/path".parse().unwrap();
assert_eq!(effective_host(&headers, &uri), None);
}
async fn init_session(app: &Router) -> String {
let req = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2025-11-25",
"capabilities": {},
"clientInfo": { "name": "test", "version": "1.0" }
}
})
.to_string(),
))
.unwrap();
let resp = app.clone().oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
resp.headers()
.get(MCP_SESSION_ID_HEADER)
.and_then(|v| v.to_str().ok())
.map(|s| s.to_string())
.expect("initialize must return a session id")
}
#[tokio::test]
async fn test_external_notification_reaches_single_session() {
let (notif_tx, notif_rx) = notification_channel(8);
let transport = HttpTransport::with_notifications(create_test_router(), notif_rx);
let (app, session_handle) = transport.into_router_with_handle();
let session_id = init_session(&app).await;
let mut rx = {
let sessions = session_handle.store.sessions.read().await;
let session = sessions
.get(&session_id)
.expect("session should be registered");
session.notifications_tx.subscribe()
};
notif_tx
.send(crate::context::ServerNotification::ResourceUpdated {
uri: "claude://chats/abc".to_string(),
})
.await
.unwrap();
let json = tokio::time::timeout(Duration::from_secs(1), rx.recv())
.await
.expect("notification should arrive within timeout")
.expect("broadcast channel closed");
assert!(json.contains("notifications/resources/updated"));
assert!(json.contains("claude://chats/abc"));
}
#[tokio::test]
async fn test_external_notification_fans_out_to_all_sessions() {
let (notif_tx, notif_rx) = notification_channel(8);
let transport = HttpTransport::with_notifications(create_test_router(), notif_rx);
let (app, session_handle) = transport.into_router_with_handle();
let session_a = init_session(&app).await;
let session_b = init_session(&app).await;
assert_ne!(session_a, session_b);
let (mut rx_a, mut rx_b) = {
let sessions = session_handle.store.sessions.read().await;
let a = sessions.get(&session_a).unwrap();
let b = sessions.get(&session_b).unwrap();
(
a.notifications_tx.subscribe(),
b.notifications_tx.subscribe(),
)
};
notif_tx
.send(crate::context::ServerNotification::ResourcesListChanged)
.await
.unwrap();
let json_a = tokio::time::timeout(Duration::from_secs(1), rx_a.recv())
.await
.unwrap()
.unwrap();
let json_b = tokio::time::timeout(Duration::from_secs(1), rx_b.recv())
.await
.unwrap()
.unwrap();
assert!(json_a.contains("notifications/resources/list_changed"));
assert!(json_b.contains("notifications/resources/list_changed"));
}
#[tokio::test]
async fn test_external_notifications_builder_method() {
let (notif_tx, notif_rx) = notification_channel(8);
let transport = HttpTransport::new(create_test_router()).external_notifications(notif_rx);
let (app, session_handle) = transport.into_router_with_handle();
let session_id = init_session(&app).await;
let mut rx = {
let sessions = session_handle.store.sessions.read().await;
sessions
.get(&session_id)
.unwrap()
.notifications_tx
.subscribe()
};
notif_tx
.send(crate::context::ServerNotification::ToolsListChanged)
.await
.unwrap();
let json = tokio::time::timeout(Duration::from_secs(1), rx.recv())
.await
.unwrap()
.unwrap();
assert!(json.contains("notifications/tools/list_changed"));
}
#[tokio::test]
async fn test_default_transport_has_no_external_fanout_task() {
let transport = HttpTransport::new(create_test_router());
let (app, _handle) = transport.into_router_with_handle();
let _session_id = init_session(&app).await;
}
#[tokio::test]
#[cfg(feature = "stateless")]
async fn stateless_v2026_initialize_is_method_not_found() {
let transport = HttpTransport::new(create_test_router())
.disable_origin_validation()
.disable_host_validation();
let app = transport.into_router();
let req = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.header(MCP_METHOD_HEADER, "initialize")
.header(MCP_PROTOCOL_VERSION_HEADER, PROTOCOL_VERSION_2026_07_28)
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2026-07-28",
"capabilities": {},
"clientInfo": { "name": "sc", "version": "1.0" },
"_meta": {
"io.modelcontextprotocol/protocolVersion": "2026-07-28",
"io.modelcontextprotocol/clientCapabilities": {}
}
}
})
.to_string(),
))
.unwrap();
let response = app.oneshot(req).await.unwrap();
assert_eq!(response.status(), StatusCode::NOT_FOUND);
assert!(
!response.headers().contains_key(MCP_SESSION_ID_HEADER),
"removed final method must not create a session"
);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert_eq!(json["id"], 1);
assert_eq!(json["error"]["code"], ErrorCode::MethodNotFound.code());
}
#[tokio::test]
#[cfg(feature = "stateless")]
async fn stateless_v2026_rejects_missing_required_meta_with_http_400() {
let app = HttpTransport::new(create_test_router())
.disable_origin_validation()
.disable_host_validation()
.into_router();
let request = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json")
.header(MCP_METHOD_HEADER, "server/discover")
.header(MCP_PROTOCOL_VERSION_HEADER, PROTOCOL_VERSION_2026_07_28)
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 101,
"method": "server/discover",
"params": {}
})
.to_string(),
))
.unwrap();
let response = app.oneshot(request).await.unwrap();
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert_eq!(json["id"], 101);
assert_eq!(json["error"]["code"], ErrorCode::InvalidParams.code());
}
#[tokio::test]
#[cfg(feature = "stateless")]
async fn stateless_v2026_rejects_invalid_meta_and_extension_keys_with_http_400() {
let app = HttpTransport::new(create_test_router())
.disable_origin_validation()
.disable_host_validation()
.into_router();
let build_request =
|id: i64, extra_meta: serde_json::Value, extensions: serde_json::Value| {
let mut meta = serde_json::json!({
"io.modelcontextprotocol/protocolVersion": PROTOCOL_VERSION_2026_07_28,
"io.modelcontextprotocol/clientCapabilities": {
"extensions": extensions
}
});
meta.as_object_mut()
.unwrap()
.extend(extra_meta.as_object().unwrap().clone());
Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json")
.header(MCP_METHOD_HEADER, "server/discover")
.header(MCP_PROTOCOL_VERSION_HEADER, PROTOCOL_VERSION_2026_07_28)
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": id,
"method": "server/discover",
"params": { "_meta": meta }
})
.to_string(),
))
.unwrap()
};
for request in [
build_request(
111,
serde_json::json!({"com.example/-invalid": true}),
serde_json::json!({}),
),
build_request(
112,
serde_json::json!({}),
serde_json::json!({"unprefixed": {}}),
),
build_request(
113,
serde_json::json!({}),
serde_json::json!({"com.example/feature": true}),
),
] {
let response = app.clone().oneshot(request).await.unwrap();
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert_eq!(json["error"]["code"], ErrorCode::InvalidParams.code());
}
}
#[tokio::test]
#[cfg(feature = "stateless")]
async fn stateless_v2026_rejects_missing_protocol_header_with_http_400() {
let app = HttpTransport::new(create_test_router())
.disable_origin_validation()
.disable_host_validation()
.into_router();
let request = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json")
.header(MCP_METHOD_HEADER, "server/discover")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 102,
"method": "server/discover",
"params": {
"_meta": {
"io.modelcontextprotocol/protocolVersion": "2026-07-28",
"io.modelcontextprotocol/clientCapabilities": {}
}
}
})
.to_string(),
))
.unwrap();
let response = app.oneshot(request).await.unwrap();
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert_eq!(json["id"], 102);
assert_eq!(json["error"]["code"], McpErrorCode::HeaderMismatch.code());
}
#[tokio::test]
#[cfg(feature = "stateless")]
async fn stateless_v2026_unknown_method_is_http_404() {
let app = HttpTransport::new(create_test_router())
.disable_origin_validation()
.disable_host_validation()
.into_router();
let request = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json")
.header(MCP_METHOD_HEADER, "unknown/method")
.header(MCP_PROTOCOL_VERSION_HEADER, PROTOCOL_VERSION_2026_07_28)
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 103,
"method": "unknown/method",
"params": {
"_meta": {
"io.modelcontextprotocol/protocolVersion": "2026-07-28",
"io.modelcontextprotocol/clientCapabilities": {}
}
}
})
.to_string(),
))
.unwrap();
let response = app.oneshot(request).await.unwrap();
assert_eq!(response.status(), StatusCode::NOT_FOUND);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert_eq!(json["id"], 103);
assert_eq!(json["error"]["code"], ErrorCode::MethodNotFound.code());
}
#[tokio::test]
#[cfg(feature = "stateless")]
async fn stateless_v2026_ignores_legacy_session_and_resumption_headers() {
let app = HttpTransport::new(create_test_router())
.disable_origin_validation()
.disable_host_validation()
.into_router();
let request = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json")
.header(MCP_METHOD_HEADER, "tools/list")
.header(MCP_PROTOCOL_VERSION_HEADER, PROTOCOL_VERSION_2026_07_28)
.header(MCP_SESSION_ID_HEADER, "legacy-session-that-does-not-exist")
.header(LAST_EVENT_ID_HEADER, "legacy-event")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 104,
"method": "tools/list",
"params": {
"_meta": {
"io.modelcontextprotocol/protocolVersion": "2026-07-28",
"io.modelcontextprotocol/clientCapabilities": {}
}
}
})
.to_string(),
))
.unwrap();
let response = app.oneshot(request).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
assert!(!response.headers().contains_key(MCP_SESSION_ID_HEADER));
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert!(json["result"]["tools"].is_array());
}
#[tokio::test]
#[cfg(feature = "stateless")]
async fn stateless_v2026_enforces_tool_client_capability_requirements() {
use crate::{CallToolResult, SamplingCapability, ToolBuilder};
let tool = ToolBuilder::new("sample")
.no_params_handler(|| async { Ok(CallToolResult::text("ok")) })
.build()
.require_client_capabilities(ClientCapabilities {
sampling: Some(SamplingCapability::default()),
..ClientCapabilities::default()
});
let router = McpRouter::new()
.server_info("test-server", "1.0.0")
.tool(tool);
let app = HttpTransport::new(router)
.disable_origin_validation()
.disable_host_validation()
.into_router();
let build_request = |id: i64, capabilities: serde_json::Value| {
Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json")
.header(MCP_METHOD_HEADER, "tools/call")
.header(MCP_NAME_HEADER, "sample")
.header(MCP_PROTOCOL_VERSION_HEADER, PROTOCOL_VERSION_2026_07_28)
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": id,
"method": "tools/call",
"params": {
"name": "sample",
"arguments": {},
"_meta": {
"io.modelcontextprotocol/protocolVersion": "2026-07-28",
"io.modelcontextprotocol/clientCapabilities": capabilities
}
}
})
.to_string(),
))
.unwrap()
};
let response = app
.clone()
.oneshot(build_request(105, serde_json::json!({})))
.await
.unwrap();
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert_eq!(json["id"], 105);
assert_eq!(
json["error"]["code"],
McpErrorCode::MissingRequiredClientCapability.code()
);
assert_eq!(
json["error"]["data"]["requiredCapabilities"],
serde_json::json!({ "sampling": {} })
);
let response = app
.oneshot(build_request(106, serde_json::json!({ "sampling": {} })))
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert_eq!(json["result"]["content"][0]["text"], "ok");
}
#[tokio::test]
#[cfg(feature = "stateless")]
async fn stateless_v2026_tools_call_without_session_succeeds() {
use crate::{CallToolResult, ToolBuilder};
let router = McpRouter::new().server_info("t", "1.0.0").tool(
ToolBuilder::new("echo")
.description("echo")
.handler(|args: serde_json::Value| async move {
Ok(CallToolResult::text(args.to_string()))
})
.build(),
);
let transport = HttpTransport::new(router).disable_origin_validation();
let app = transport.into_router();
let req = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json")
.header(MCP_PROTOCOL_VERSION_HEADER, "2026-07-28")
.header(MCP_METHOD_HEADER, "tools/call")
.header(MCP_NAME_HEADER, "echo")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "tools/call",
"params": {
"name": "echo",
"arguments": {"message": "hello"},
"_meta": {
"io.modelcontextprotocol/protocolVersion": "2026-07-28",
"io.modelcontextprotocol/clientInfo": {
"name": "sc", "version": "1.0"
},
"io.modelcontextprotocol/clientCapabilities": {}
}
}
})
.to_string(),
))
.unwrap();
let response = app.oneshot(req).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
assert!(
!response.headers().contains_key(MCP_SESSION_ID_HEADER),
"stateless tools/call must not set mcp-session-id"
);
assert_eq!(
response
.headers()
.get(MCP_PROTOCOL_VERSION_HEADER)
.and_then(|v| v.to_str().ok()),
Some("2026-07-28")
);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert!(
json.get("result").is_some(),
"expected tools/call result, got: {json}"
);
assert_eq!(json["result"]["resultType"], "complete");
}
#[cfg(feature = "stateless")]
fn stateless_tools_call_request() -> Request<Body> {
Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json")
.header(MCP_PROTOCOL_VERSION_HEADER, "2026-07-28")
.header(MCP_METHOD_HEADER, "tools/call")
.header(MCP_NAME_HEADER, "echo")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "tools/call",
"params": {
"name": "echo",
"arguments": {"message": "hello"},
"_meta": {
"io.modelcontextprotocol/protocolVersion": "2026-07-28",
"io.modelcontextprotocol/clientInfo": {
"name": "sc", "version": "1.0"
},
"io.modelcontextprotocol/clientCapabilities": {}
}
}
})
.to_string(),
))
.unwrap()
}
#[cfg(feature = "stateless")]
fn echo_router() -> McpRouter {
use crate::{CallToolResult, ToolBuilder};
McpRouter::new().server_info("t", "1.0.0").tool(
ToolBuilder::new("echo")
.description("echo")
.handler(|args: serde_json::Value| async move {
Ok(CallToolResult::text(args.to_string()))
})
.build(),
)
}
#[tokio::test]
#[cfg(feature = "stateless")]
async fn stateless_v2026_response_stamps_server_info_by_default() {
let transport = HttpTransport::new(echo_router()).disable_origin_validation();
let app = transport.into_router();
let response = app.oneshot(stateless_tools_call_request()).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert_eq!(
json["result"]["_meta"]["io.modelcontextprotocol/serverInfo"]["name"], "t",
"expected serverInfo stamped into result._meta, got: {json}"
);
assert_eq!(
json["result"]["_meta"]["io.modelcontextprotocol/serverInfo"]["version"],
"1.0.0"
);
}
#[tokio::test]
#[cfg(feature = "stateless")]
async fn stateless_v2026_response_omits_server_info_when_disabled() {
let transport = HttpTransport::new(echo_router())
.disable_origin_validation()
.stamp_server_info(false);
let app = transport.into_router();
let response = app.oneshot(stateless_tools_call_request()).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert!(
json["result"].get("_meta").is_none(),
"expected no _meta when stamping is disabled, got: {json}"
);
}
#[tokio::test]
async fn stateless_v2025_initialize_still_gets_session_id() {
let transport = HttpTransport::new(create_test_router()).disable_origin_validation();
let app = transport.into_router();
let req = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0", "id": 1, "method": "initialize",
"params": {
"protocolVersion": "2025-11-25",
"capabilities": {},
"clientInfo": { "name": "old-client", "version": "1.0" }
}
})
.to_string(),
))
.unwrap();
let response = app.oneshot(req).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
assert!(
response.headers().contains_key(MCP_SESSION_ID_HEADER),
"2025-11-25 initialize must return mcp-session-id"
);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert!(
json["result"].get("resultType").is_none(),
"legacy result must remain unchanged: {json}"
);
}
#[tokio::test]
async fn stateless_v2025_tools_list_without_session_rejected() {
let transport = HttpTransport::new(create_test_router())
.disable_origin_validation()
.require_sessions();
let app = transport.into_router();
let req = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json")
.header(MCP_PROTOCOL_VERSION_HEADER, "2025-11-25")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0", "id": 1, "method": "tools/list"
})
.to_string(),
))
.unwrap();
let response = app.oneshot(req).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert!(json.get("error").is_some(), "expected error, got: {json}");
assert_eq!(
json["error"]["code"].as_i64().unwrap(),
-32006,
"expected SessionRequired (-32006)"
);
}
#[tokio::test]
#[cfg(feature = "stateless")]
async fn stateless_v2026_tools_list_without_session_succeeds() {
let transport = HttpTransport::new(create_test_router()).disable_origin_validation();
let app = transport.into_router();
let req = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json")
.header(MCP_PROTOCOL_VERSION_HEADER, "2026-07-28")
.header(MCP_METHOD_HEADER, "tools/list")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "tools/list",
"params": {
"_meta": {
"io.modelcontextprotocol/protocolVersion": "2026-07-28",
"io.modelcontextprotocol/clientCapabilities": {}
}
}
})
.to_string(),
))
.unwrap();
let response = app.oneshot(req).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
assert!(
!response.headers().contains_key(MCP_SESSION_ID_HEADER),
"stateless tools/list must not set mcp-session-id"
);
assert_eq!(
response
.headers()
.get(MCP_PROTOCOL_VERSION_HEADER)
.and_then(|v| v.to_str().ok()),
Some("2026-07-28")
);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert!(
json["result"]["tools"].is_array(),
"expected tools array in result, got: {json}"
);
assert_eq!(json["result"]["resultType"], "complete");
assert_eq!(json["result"]["ttlMs"], 0);
assert_eq!(json["result"]["cacheScope"], "private");
}
#[tokio::test]
#[cfg(feature = "stateless")]
async fn stateless_v2026_notification_returns_202_no_session() {
let transport = HttpTransport::new(create_test_router()).disable_origin_validation();
let (app, handle) = transport.into_router_with_handle();
let req = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json")
.header(MCP_PROTOCOL_VERSION_HEADER, "2026-07-28")
.header(MCP_METHOD_HEADER, "notifications/cancelled")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"method": "notifications/cancelled",
"params": {
"requestId": 99,
"reason": "test",
"_meta": {
"io.modelcontextprotocol/protocolVersion": "2026-07-28",
"io.modelcontextprotocol/clientCapabilities": {}
}
}
})
.to_string(),
))
.unwrap();
let response = app.oneshot(req).await.unwrap();
assert_eq!(
response.status(),
StatusCode::ACCEPTED,
"stateless notification must return 202 ACCEPTED"
);
assert!(
!response.headers().contains_key(MCP_SESSION_ID_HEADER),
"stateless notification must not set mcp-session-id"
);
assert_eq!(
handle.session_count().await,
0,
"stateless notification must not create a session"
);
}
#[tokio::test]
#[cfg(feature = "stateless")]
async fn stateless_v2026_missing_mcp_method_returns_400() {
let transport = HttpTransport::new(create_test_router()).disable_origin_validation();
let app = transport.into_router();
let req = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json")
.header(MCP_PROTOCOL_VERSION_HEADER, "2026-07-28")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "tools/list",
"params": {
"_meta": {
"io.modelcontextprotocol/protocolVersion": "2026-07-28",
"io.modelcontextprotocol/clientCapabilities": {}
}
}
})
.to_string(),
))
.unwrap();
let response = app.oneshot(req).await.unwrap();
assert_eq!(
response.status(),
StatusCode::BAD_REQUEST,
"missing Mcp-Method must return HTTP 400"
);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert!(json.get("error").is_some(), "expected error, got: {json}");
assert_eq!(
json["error"]["code"].as_i64().unwrap(),
-32020,
"expected HeaderMismatch (-32020)"
);
assert!(
json["error"]["message"]
.as_str()
.unwrap_or("")
.contains("Mcp-Method"),
"error message must mention Mcp-Method, got: {json}"
);
}
#[tokio::test]
async fn sse_responses_false_returns_application_json() {
let transport = HttpTransport::new(create_test_router())
.disable_origin_validation()
.sse_responses(false);
let app = transport.into_router();
let request = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2025-11-25",
"capabilities": {},
"clientInfo": {"name": "test", "version": "0.1"}
}
})
.to_string(),
))
.unwrap();
let response = app.oneshot(request).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let ct = response
.headers()
.get(header::CONTENT_TYPE)
.and_then(|v| v.to_str().ok())
.unwrap_or("");
assert!(
ct.contains("application/json"),
"sse_responses(false) should return application/json, got: {ct}"
);
}
#[tokio::test]
async fn sse_responses_true_returns_text_event_stream_with_valid_json() {
let transport = HttpTransport::new(create_test_router())
.disable_origin_validation()
.sse_responses(true);
let app = transport.into_router();
let init_body = serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2025-11-25",
"capabilities": {},
"clientInfo": {"name": "test", "version": "0.1"}
}
})
.to_string();
let request = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.body(Body::from(init_body))
.unwrap();
let response = app.oneshot(request).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let ct = response
.headers()
.get(header::CONTENT_TYPE)
.and_then(|v| v.to_str().ok())
.unwrap_or("");
assert!(
ct.contains("text/event-stream"),
"sse_responses(true) should return text/event-stream, got: {ct}"
);
let bytes = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
let body_text = String::from_utf8_lossy(&bytes);
assert!(
body_text.contains("event: message"),
"SSE body missing 'event: message': {body_text}"
);
assert!(
body_text.contains("data: "),
"SSE body missing 'data: ' line: {body_text}"
);
let data_line = body_text
.lines()
.find(|l| l.starts_with("data: "))
.expect("no data: line in SSE body");
let json_str = data_line.trim_start_matches("data: ");
let val: serde_json::Value =
serde_json::from_str(json_str).expect("data: line is not valid JSON");
assert_eq!(val["jsonrpc"], "2.0", "jsonrpc version mismatch: {val}");
assert_eq!(val["id"], 1, "id mismatch: {val}");
assert!(
val["result"].is_object(),
"result should be an object: {val}"
);
assert_eq!(
val["result"]["protocolVersion"].as_str(),
Some("2025-11-25"),
"protocolVersion missing or wrong: {val}"
);
}
#[tokio::test]
async fn sse_responses_true_tools_list_returns_valid_sse() {
let transport = HttpTransport::new(create_test_router())
.disable_origin_validation()
.sse_responses(true);
let app = transport.into_router();
let init_request = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2025-11-25",
"capabilities": {},
"clientInfo": {"name": "test", "version": "0.1"}
}
})
.to_string(),
))
.unwrap();
let init_response = app.clone().oneshot(init_request).await.unwrap();
assert_eq!(init_response.status(), StatusCode::OK);
let session_id = init_response
.headers()
.get(MCP_SESSION_ID_HEADER)
.and_then(|v| v.to_str().ok())
.map(|s| s.to_string())
.expect("missing session ID from initialize");
let notif_request = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.header(MCP_SESSION_ID_HEADER, &session_id)
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"method": "notifications/initialized"
})
.to_string(),
))
.unwrap();
app.clone().oneshot(notif_request).await.unwrap();
let list_request = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.header(MCP_SESSION_ID_HEADER, &session_id)
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 2,
"method": "tools/list",
"params": {}
})
.to_string(),
))
.unwrap();
let list_response = app.oneshot(list_request).await.unwrap();
assert_eq!(list_response.status(), StatusCode::OK);
let ct = list_response
.headers()
.get(header::CONTENT_TYPE)
.and_then(|v| v.to_str().ok())
.unwrap_or("");
assert!(
ct.contains("text/event-stream"),
"tools/list with sse_responses(true) should return text/event-stream, got: {ct}"
);
let bytes = axum::body::to_bytes(list_response.into_body(), usize::MAX)
.await
.unwrap();
let body_text = String::from_utf8_lossy(&bytes);
let data_line = body_text
.lines()
.find(|l| l.starts_with("data: "))
.expect("no data: line in SSE body for tools/list");
let json_str = data_line.trim_start_matches("data: ");
let val: serde_json::Value =
serde_json::from_str(json_str).expect("tools/list data: line is not valid JSON");
assert_eq!(val["jsonrpc"], "2.0");
assert_eq!(val["id"], 2);
assert!(
val["result"]["tools"].is_array(),
"tools/list result.tools should be an array: {val}"
);
}
async fn do_initialize(app: &axum::Router) -> String {
let init_request = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 1,
"method": "initialize",
"params": {
"protocolVersion": "2025-11-25",
"capabilities": {},
"clientInfo": { "name": "test-client", "version": "1.0.0" }
}
})
.to_string(),
))
.unwrap();
let response = app.clone().oneshot(init_request).await.unwrap();
response
.headers()
.get(MCP_SESSION_ID_HEADER)
.unwrap()
.to_str()
.unwrap()
.to_string()
}
#[tokio::test]
async fn tools_list_before_initialized_notification_returns_error() {
let transport = HttpTransport::new(create_test_router()).disable_origin_validation();
let app = transport.into_router();
let session_id = do_initialize(&app).await;
let list_request = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.header(MCP_SESSION_ID_HEADER, &session_id)
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 2,
"method": "tools/list"
})
.to_string(),
))
.unwrap();
let response = app.oneshot(list_request).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert!(
json.get("error").is_some(),
"expected error when notifications/initialized not sent, got: {json}"
);
assert_eq!(
json["error"]["code"].as_i64().unwrap(),
-32600,
"expected InvalidRequest (-32600), got: {json}"
);
assert!(
json["error"]["message"]
.as_str()
.unwrap_or("")
.contains("notifications/initialized"),
"error message should mention notifications/initialized, got: {json}"
);
}
#[tokio::test]
async fn tools_list_after_initialized_notification_succeeds() {
let transport = HttpTransport::new(create_test_router()).disable_origin_validation();
let app = transport.into_router();
let session_id = do_initialize(&app).await;
let notif_request = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.header(MCP_SESSION_ID_HEADER, &session_id)
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"method": "notifications/initialized"
})
.to_string(),
))
.unwrap();
app.clone().oneshot(notif_request).await.unwrap();
let list_request = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.header(MCP_SESSION_ID_HEADER, &session_id)
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 2,
"method": "tools/list"
})
.to_string(),
))
.unwrap();
let response = app.oneshot(list_request).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert!(
json.get("result").is_some(),
"expected success after notifications/initialized, got: {json}"
);
}
#[tokio::test]
async fn notifications_initialized_itself_always_accepted() {
let transport = HttpTransport::new(create_test_router()).disable_origin_validation();
let app = transport.into_router();
let session_id = do_initialize(&app).await;
let notif_request = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.header(MCP_SESSION_ID_HEADER, &session_id)
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"method": "notifications/initialized"
})
.to_string(),
))
.unwrap();
let response = app.oneshot(notif_request).await.unwrap();
assert_eq!(
response.status(),
StatusCode::ACCEPTED,
"notifications/initialized must return 202 ACCEPTED"
);
}
#[tokio::test]
async fn strict_initialization_false_allows_tools_list_without_notification() {
let config = SessionConfig {
strict_initialization: false,
..Default::default()
};
let transport = HttpTransport::new(create_test_router())
.disable_origin_validation()
.session_config(config);
let app = transport.into_router();
let session_id = do_initialize(&app).await;
let list_request = Request::builder()
.method("POST")
.uri("/")
.header("Content-Type", "application/json")
.header("Accept", "application/json, text/event-stream")
.header(MCP_SESSION_ID_HEADER, &session_id)
.body(Body::from(
serde_json::json!({
"jsonrpc": "2.0",
"id": 2,
"method": "tools/list"
})
.to_string(),
))
.unwrap();
let response = app.oneshot(list_request).await.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.unwrap();
let json: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert!(
json.get("result").is_some(),
"expected success with strict_initialization=false, got: {json}"
);
}
}