#[non_exhaustive]
pub enum Transport {
Stdio,
#[cfg(feature = "transport-http")]
Http(HttpConfig),
}
#[cfg(feature = "transport-http")]
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct HttpConfig {
pub bind: std::net::SocketAddr,
pub path: String,
pub max_request_body_bytes: usize,
pub max_concurrent_sessions: usize,
}
#[cfg(feature = "transport-http")]
impl HttpConfig {
pub const DEFAULT_MAX_REQUEST_BODY_BYTES: usize = 4 * 1024 * 1024;
pub const DEFAULT_MAX_CONCURRENT_SESSIONS: usize = 100;
pub fn new(bind: std::net::SocketAddr, path: impl Into<String>) -> Self {
Self {
bind,
path: path.into(),
max_request_body_bytes: Self::DEFAULT_MAX_REQUEST_BODY_BYTES,
max_concurrent_sessions: Self::DEFAULT_MAX_CONCURRENT_SESSIONS,
}
}
#[must_use]
pub const fn with_max_request_body_bytes(mut self, bytes: usize) -> Self {
self.max_request_body_bytes = bytes;
self
}
#[must_use]
pub const fn with_max_concurrent_sessions(mut self, max: usize) -> Self {
self.max_concurrent_sessions = max;
self
}
}
use rmcp::ServiceExt as _;
#[cfg(feature = "transport-http")]
use rmcp::model::{ClientJsonRpcMessage, ServerJsonRpcMessage};
#[cfg(feature = "transport-http")]
use rmcp::transport::streamable_http_server::session::local::{
LocalSessionManager, LocalSessionManagerError,
};
#[cfg(feature = "transport-http")]
use rmcp::transport::streamable_http_server::session::{
ServerSseMessage, SessionId, SessionManager,
};
pub(crate) struct ShutdownSignal {
#[cfg(unix)]
sigterm: Option<tokio::signal::unix::Signal>,
#[cfg(unix)]
sigint: Option<tokio::signal::unix::Signal>,
#[cfg(windows)]
ctrl_c: Option<tokio::signal::windows::CtrlC>,
}
impl ShutdownSignal {
pub(crate) fn new() -> Self {
#[cfg(unix)]
{
use tokio::signal::unix::{SignalKind, signal};
let sigterm = match signal(SignalKind::terminate()) {
Ok(sigterm) => Some(sigterm),
Err(e) => {
tracing::warn!(
"SIGTERM handler registration failed ({e}), SIGTERM will not be caught"
);
None
}
};
let sigint = match signal(SignalKind::interrupt()) {
Ok(sigint) => Some(sigint),
Err(e) => {
tracing::warn!(
"SIGINT handler registration failed ({e}), SIGINT will not be caught"
);
None
}
};
Self { sigterm, sigint }
}
#[cfg(windows)]
{
let ctrl_c = match tokio::signal::windows::ctrl_c() {
Ok(ctrl_c) => Some(ctrl_c),
Err(e) => {
tracing::warn!("Ctrl-C handler registration failed ({e})");
None
}
};
Self { ctrl_c }
}
#[cfg(not(any(unix, windows)))]
{
Self {}
}
}
pub(crate) async fn recv(&mut self) {
#[cfg(unix)]
{
match (self.sigterm.as_mut(), self.sigint.as_mut()) {
(Some(sigterm), Some(sigint)) => {
tokio::select! {
_ = sigterm.recv() => {},
_ = sigint.recv() => {},
}
}
(Some(sigterm), None) => {
sigterm.recv().await;
}
(None, Some(sigint)) => {
sigint.recv().await;
}
(None, None) => {
let _ = tokio::signal::ctrl_c().await;
}
}
}
#[cfg(windows)]
{
match self.ctrl_c.as_mut() {
Some(ctrl_c) => {
ctrl_c.recv().await;
}
None => {
let _ = tokio::signal::ctrl_c().await;
}
}
}
#[cfg(not(any(unix, windows)))]
{
let _ = tokio::signal::ctrl_c().await;
}
}
}
pub(crate) async fn run_stdio(
mcp_server: crate::mcp::McplsServer,
peer_cell: &tokio::sync::OnceCell<rmcp::Peer<rmcp::RoleServer>>,
mut shutdown_signal: ShutdownSignal,
) -> Result<(), crate::Error> {
let service = tokio::select! {
result = mcp_server.serve(rmcp::transport::stdio()) => {
result.map_err(|e| crate::Error::McpServer(format!("Failed to start MCP server: {e}")))?
}
() = shutdown_signal.recv() => {
tracing::info!("shutdown signal received during handshake, stopping stdio transport");
return Ok(());
}
};
if let Err(e) = peer_cell.set(service.peer().clone()) {
tracing::debug!("Peer cell already set ({}), ignoring", e);
}
tokio::select! {
result = service.waiting() => result
.map(|_| ())
.map_err(|e| crate::Error::McpServer(format!("MCP server error: {e}"))),
() = shutdown_signal.recv() => {
tracing::info!("shutdown signal received, stopping stdio transport");
Ok(())
}
}
}
#[cfg(feature = "transport-http")]
#[allow(clippy::significant_drop_tightening)]
pub(crate) async fn run_http(
mcp_server: crate::mcp::McplsServer,
cfg: HttpConfig,
mut shutdown_signal: ShutdownSignal,
) -> Result<(), crate::Error> {
use std::sync::Arc;
use rmcp::transport::streamable_http_server::{
StreamableHttpServerConfig, StreamableHttpService,
};
use tokio_util::sync::CancellationToken;
let session_manager = Arc::new(CappedSessionManager::new(cfg.max_concurrent_sessions));
let cancel = CancellationToken::new();
let mcp_for_factory = mcp_server.clone();
let mut http_cfg = StreamableHttpServerConfig::default();
http_cfg.cancellation_token = cancel.clone();
http_cfg.max_request_body_bytes = cfg.max_request_body_bytes;
let service = StreamableHttpService::new(
move || Ok::<_, std::io::Error>(mcp_for_factory.clone()),
session_manager,
http_cfg,
);
let app = axum::Router::new()
.nest_service(&cfg.path, service.clone())
.route_service("/", service)
.layer(axum::middleware::from_fn(enforce_session_cap));
let listener = tokio::net::TcpListener::bind(cfg.bind)
.await
.map_err(|e| crate::Error::McpServer(format!("bind {}: {e}", cfg.bind)))?;
tracing::info!(addr = %cfg.bind, path = %cfg.path, "MCP HTTP transport listening");
if !cfg.bind.ip().is_loopback() {
tracing::warn!(
addr = %cfg.bind,
"binding to a non-loopback address: mcpls performs no authentication of its own on \
any transport — place this endpoint behind a reverse proxy that enforces \
authentication. The proxy must also rewrite the Host header, since rmcp's Host \
validation allows only localhost/127.0.0.1/::1 by default"
);
}
let cancel_for_force_timeout = cancel.clone();
let serve = axum::serve(listener, app).with_graceful_shutdown(async move {
shutdown_signal.recv().await;
cancel.cancel();
});
tokio::select! {
result = serve => result.map_err(|e| crate::Error::McpServer(format!("http serve: {e}"))),
() = async move {
cancel_for_force_timeout.cancelled().await;
tokio::time::sleep(HTTP_GRACEFUL_SHUTDOWN_TIMEOUT).await;
} => {
tracing::warn!(
timeout = ?HTTP_GRACEFUL_SHUTDOWN_TIMEOUT,
"HTTP graceful shutdown did not complete in time, proceeding with shutdown anyway"
);
Ok(())
}
}
}
#[cfg(feature = "transport-http")]
const HTTP_GRACEFUL_SHUTDOWN_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(30);
#[cfg(feature = "transport-http")]
struct CappedSessionManager {
inner: LocalSessionManager,
semaphore: std::sync::Arc<tokio::sync::Semaphore>,
permits:
tokio::sync::Mutex<std::collections::HashMap<SessionId, tokio::sync::OwnedSemaphorePermit>>,
}
#[cfg(feature = "transport-http")]
impl CappedSessionManager {
fn new(max_sessions: usize) -> Self {
Self {
inner: LocalSessionManager::default(),
semaphore: std::sync::Arc::new(tokio::sync::Semaphore::new(max_sessions)),
permits: tokio::sync::Mutex::new(std::collections::HashMap::new()),
}
}
}
#[cfg(feature = "transport-http")]
const SESSION_CAP_MARKER: &str = "mcpls-http-session-cap-reached";
#[cfg(feature = "transport-http")]
#[derive(Debug, thiserror::Error)]
enum CappedSessionManagerError {
#[error("{SESSION_CAP_MARKER}: maximum concurrent HTTP sessions already active")]
CapReached,
#[error(transparent)]
Inner(#[from] LocalSessionManagerError),
}
#[cfg(feature = "transport-http")]
impl SessionManager for CappedSessionManager {
type Error = CappedSessionManagerError;
type Transport = <LocalSessionManager as SessionManager>::Transport;
async fn create_session(&self) -> Result<(SessionId, Self::Transport), Self::Error> {
let permit = self
.semaphore
.clone()
.try_acquire_owned()
.map_err(|_| CappedSessionManagerError::CapReached)?;
let (id, transport) = self.inner.create_session().await?;
self.permits.lock().await.insert(id.clone(), permit);
Ok((id, transport))
}
async fn initialize_session(
&self,
id: &SessionId,
message: ClientJsonRpcMessage,
) -> Result<ServerJsonRpcMessage, Self::Error> {
Ok(self.inner.initialize_session(id, message).await?)
}
async fn has_session(&self, id: &SessionId) -> Result<bool, Self::Error> {
Ok(self.inner.has_session(id).await?)
}
async fn close_session(&self, id: &SessionId) -> Result<(), Self::Error> {
self.permits.lock().await.remove(id);
self.inner.close_session(id).await?;
Ok(())
}
async fn create_stream(
&self,
id: &SessionId,
message: ClientJsonRpcMessage,
) -> Result<impl futures::Stream<Item = ServerSseMessage> + Send + Sync + 'static, Self::Error>
{
Ok(self.inner.create_stream(id, message).await?)
}
async fn accept_message(
&self,
id: &SessionId,
message: ClientJsonRpcMessage,
) -> Result<(), Self::Error> {
Ok(self.inner.accept_message(id, message).await?)
}
async fn create_standalone_stream(
&self,
id: &SessionId,
) -> Result<impl futures::Stream<Item = ServerSseMessage> + Send + Sync + 'static, Self::Error>
{
Ok(self.inner.create_standalone_stream(id).await?)
}
async fn resume(
&self,
id: &SessionId,
last_event_id: String,
) -> Result<impl futures::Stream<Item = ServerSseMessage> + Send + Sync + 'static, Self::Error>
{
Ok(self.inner.resume(id, last_event_id).await?)
}
}
#[cfg(feature = "transport-http")]
async fn enforce_session_cap(
request: axum::extract::Request,
next: axum::middleware::Next,
) -> axum::response::Response {
let response = next.run(request).await;
if response.status() != axum::http::StatusCode::INTERNAL_SERVER_ERROR {
return response;
}
let (mut parts, body) = response.into_parts();
let Ok(bytes) = axum::body::to_bytes(body, 64 * 1024).await else {
return axum::response::Response::from_parts(
parts,
axum::body::Body::from("Internal Server Error"),
);
};
if bytes
.windows(SESSION_CAP_MARKER.len())
.any(|window| window == SESSION_CAP_MARKER.as_bytes())
{
parts.status = axum::http::StatusCode::TOO_MANY_REQUESTS;
parts.headers.insert(
axum::http::header::RETRY_AFTER,
axum::http::HeaderValue::from_static("1"),
);
return axum::response::Response::from_parts(
parts,
axum::body::Body::from("Too Many Requests: maximum concurrent HTTP sessions reached"),
);
}
axum::response::Response::from_parts(parts, axum::body::Body::from(bytes))
}
#[cfg(test)]
#[allow(clippy::unwrap_used)]
mod tests {
#[test]
fn test_transport_stdio_variant() {
let t = super::Transport::Stdio;
assert!(matches!(t, super::Transport::Stdio));
}
#[tokio::test]
async fn test_run_stdio_returns_promptly_when_stdin_is_already_closed() {
use std::path::PathBuf;
use std::sync::Arc;
use tokio::sync::Mutex;
use crate::bridge::{NotificationCache, ResourceSubscriptions, Translator};
use crate::mcp::McplsServer;
let translator = Arc::new(Translator::new());
let notification_cache = Arc::new(Mutex::new(NotificationCache::new()));
let workspace_roots: Arc<[PathBuf]> = Arc::from(Vec::new());
let subs = Arc::new(ResourceSubscriptions::new());
let server = McplsServer::new(translator, notification_cache, workspace_roots, subs, false);
let peer_cell = tokio::sync::OnceCell::new();
let outcome = tokio::time::timeout(
std::time::Duration::from_secs(2),
super::run_stdio(server, &peer_cell, super::ShutdownSignal::new()),
)
.await;
assert!(
outcome.is_ok(),
"run_stdio must not hang when stdin is already closed"
);
let result = outcome.unwrap();
assert!(
matches!(result, Err(crate::Error::McpServer(_))),
"expected a McpServer error from the failed handshake, got: {result:?}"
);
}
#[cfg(feature = "transport-http")]
mod http_tests {
use std::net::SocketAddr;
use super::super::{HttpConfig, Transport};
#[test]
fn test_http_config_fields() {
let addr: SocketAddr = "127.0.0.1:3000".parse().unwrap();
let cfg = HttpConfig::new(addr, "/mcp");
assert_eq!(cfg.bind, addr);
assert_eq!(cfg.path, "/mcp");
}
#[test]
fn test_http_config_clone() {
let cfg = HttpConfig::new("127.0.0.1:3001".parse().unwrap(), "/test");
let cloned = cfg.clone();
assert_eq!(cloned.bind, cfg.bind);
assert_eq!(cloned.path, cfg.path);
}
#[test]
fn test_transport_http_variant() {
let cfg = HttpConfig::new("127.0.0.1:3002".parse().unwrap(), "/mcp");
let t = Transport::Http(cfg);
assert!(matches!(t, Transport::Http(_)));
}
#[test]
fn test_http_config_new_uses_default_limits() {
let cfg = HttpConfig::new("127.0.0.1:3003".parse().unwrap(), "/mcp");
assert_eq!(
cfg.max_request_body_bytes,
HttpConfig::DEFAULT_MAX_REQUEST_BODY_BYTES
);
assert_eq!(
cfg.max_concurrent_sessions,
HttpConfig::DEFAULT_MAX_CONCURRENT_SESSIONS
);
}
#[test]
fn test_http_config_with_max_request_body_bytes_overrides_default() {
let cfg = HttpConfig::new("127.0.0.1:3004".parse().unwrap(), "/mcp")
.with_max_request_body_bytes(1024);
assert_eq!(cfg.max_request_body_bytes, 1024);
assert_eq!(
cfg.max_concurrent_sessions,
HttpConfig::DEFAULT_MAX_CONCURRENT_SESSIONS
);
}
#[test]
fn test_http_config_with_max_concurrent_sessions_overrides_default() {
let cfg = HttpConfig::new("127.0.0.1:3005".parse().unwrap(), "/mcp")
.with_max_concurrent_sessions(5);
assert_eq!(cfg.max_concurrent_sessions, 5);
assert_eq!(
cfg.max_request_body_bytes,
HttpConfig::DEFAULT_MAX_REQUEST_BODY_BYTES
);
}
#[tokio::test]
async fn test_run_http_binds() {
use std::path::PathBuf;
use std::sync::Arc;
use tokio::sync::Mutex;
use crate::bridge::{NotificationCache, ResourceSubscriptions, Translator};
use crate::mcp::McplsServer;
let translator = Arc::new(Translator::new());
let notification_cache = Arc::new(Mutex::new(NotificationCache::new()));
let workspace_roots: Arc<[PathBuf]> = Arc::from(Vec::new());
let subs = Arc::new(ResourceSubscriptions::new());
let server =
McplsServer::new(translator, notification_cache, workspace_roots, subs, false);
let probe = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = probe.local_addr().unwrap();
drop(probe);
let cfg = HttpConfig::new(addr, "/mcp");
let server_task = tokio::spawn(super::super::run_http(
server,
cfg,
super::super::ShutdownSignal::new(),
));
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
let connected = tokio::net::TcpStream::connect(addr).await;
assert!(
connected.is_ok(),
"HTTP listener should accept TCP connections"
);
server_task.abort();
}
#[tokio::test(start_paused = true)]
async fn test_run_http_does_not_self_terminate_without_signal() {
let probe = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = probe.local_addr().unwrap();
drop(probe);
let cfg = HttpConfig::new(addr, "/mcp");
let server_task = tokio::spawn(super::super::run_http(
test_server(),
cfg,
super::super::ShutdownSignal::new(),
));
for _ in 0..10 {
tokio::task::yield_now().await;
}
tokio::time::advance(
super::super::HTTP_GRACEFUL_SHUTDOWN_TIMEOUT + std::time::Duration::from_secs(5),
)
.await;
for _ in 0..10 {
tokio::task::yield_now().await;
}
assert!(
!server_task.is_finished(),
"run_http must still be serving after HTTP_GRACEFUL_SHUTDOWN_TIMEOUT of uptime \
with no shutdown signal sent"
);
server_task.abort();
}
#[tokio::test]
async fn test_run_http_bind_error() {
use std::path::PathBuf;
use std::sync::Arc;
use tokio::sync::Mutex;
use crate::bridge::{NotificationCache, ResourceSubscriptions, Translator};
use crate::mcp::McplsServer;
let occupied = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = occupied.local_addr().unwrap();
let translator = Arc::new(Translator::new());
let notification_cache = Arc::new(Mutex::new(NotificationCache::new()));
let workspace_roots: Arc<[PathBuf]> = Arc::from(Vec::new());
let subs = Arc::new(ResourceSubscriptions::new());
let server =
McplsServer::new(translator, notification_cache, workspace_roots, subs, false);
let cfg = HttpConfig::new(addr, "/mcp");
let result =
super::super::run_http(server, cfg, super::super::ShutdownSignal::new()).await;
assert!(
result.is_err(),
"run_http should fail when port is occupied"
);
drop(occupied);
}
fn test_server() -> crate::mcp::McplsServer {
use std::path::PathBuf;
use std::sync::Arc;
use tokio::sync::Mutex;
use crate::bridge::{NotificationCache, ResourceSubscriptions, Translator};
use crate::mcp::McplsServer;
let translator = Arc::new(Translator::new());
let notification_cache = Arc::new(Mutex::new(NotificationCache::new()));
let workspace_roots: Arc<[PathBuf]> = Arc::from(Vec::new());
let subs = Arc::new(ResourceSubscriptions::new());
McplsServer::new(translator, notification_cache, workspace_roots, subs, false)
}
async fn raw_http_post(
addr: SocketAddr,
path: &str,
extra_headers: &str,
body: &[u8],
) -> String {
use tokio::io::{AsyncReadExt as _, AsyncWriteExt as _};
let mut stream = tokio::net::TcpStream::connect(addr).await.unwrap();
let request = format!(
"POST {path} HTTP/1.1\r\nHost: {addr}\r\nConnection: close\r\n{extra_headers}Content-Length: {}\r\n\r\n",
body.len()
);
stream.write_all(request.as_bytes()).await.unwrap();
stream.write_all(body).await.unwrap();
let mut response = Vec::new();
let mut buf = [0u8; 8192];
loop {
match tokio::time::timeout(std::time::Duration::from_secs(2), stream.read(&mut buf))
.await
{
Ok(Ok(0)) | Err(_) => break,
Ok(Ok(n)) => response.extend_from_slice(&buf[..n]),
Ok(Err(e)) => panic!("read error: {e}"),
}
}
String::from_utf8_lossy(&response).into_owned()
}
#[tokio::test]
async fn test_run_http_rejects_oversized_body_with_413() {
let probe = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = probe.local_addr().unwrap();
drop(probe);
let cfg = HttpConfig::new(addr, "/mcp").with_max_request_body_bytes(64);
let server_task = tokio::spawn(super::super::run_http(
test_server(),
cfg,
super::super::ShutdownSignal::new(),
));
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
let oversized_body = vec![b'a'; 65];
let response = raw_http_post(
addr,
"/mcp",
"Accept: application/json, text/event-stream\r\nContent-Type: application/json\r\n",
&oversized_body,
)
.await;
assert!(
response.starts_with("HTTP/1.1 413"),
"expected 413 Payload Too Large, got: {response}"
);
server_task.abort();
}
#[tokio::test]
async fn test_run_http_accepts_body_within_limit() {
let probe = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = probe.local_addr().unwrap();
drop(probe);
let cfg = HttpConfig::new(addr, "/mcp").with_max_request_body_bytes(64);
let server_task = tokio::spawn(super::super::run_http(
test_server(),
cfg,
super::super::ShutdownSignal::new(),
));
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
let small_body = vec![b'a'; 32];
let response = raw_http_post(
addr,
"/mcp",
"Accept: application/json, text/event-stream\r\nContent-Type: application/json\r\n",
&small_body,
)
.await;
assert!(
!response.starts_with("HTTP/1.1 413"),
"body within limit must not be rejected as too large, got: {response}"
);
server_task.abort();
}
#[allow(clippy::significant_drop_tightening)]
#[tokio::test]
async fn test_capped_session_manager_enforces_hard_bound() {
use rmcp::transport::streamable_http_server::session::SessionManager as _;
let manager = super::super::CappedSessionManager::new(1);
let (first_id, _transport) = manager.create_session().await.unwrap();
let second_err = manager.create_session().await.map(|_| ()).unwrap_err();
assert!(
matches!(
second_err,
super::super::CappedSessionManagerError::CapReached
),
"expected CapReached once at capacity, got: {second_err:?}"
);
manager.close_session(&first_id).await.unwrap();
let (third_id, _transport) = manager.create_session().await.unwrap();
assert_ne!(first_id, third_id);
}
#[allow(clippy::significant_drop_tightening)]
#[tokio::test]
async fn test_capped_session_manager_bounds_concurrent_create_session() {
use rmcp::transport::streamable_http_server::session::SessionManager as _;
const MAX_SESSIONS: usize = 5;
const CONCURRENT_ATTEMPTS: usize = 25;
let manager =
std::sync::Arc::new(super::super::CappedSessionManager::new(MAX_SESSIONS));
let mut tasks = tokio::task::JoinSet::new();
for _ in 0..CONCURRENT_ATTEMPTS {
let manager = manager.clone();
tasks.spawn(async move { manager.create_session().await.is_ok() });
}
let mut succeeded = 0usize;
while let Some(result) = tasks.join_next().await {
if result.unwrap() {
succeeded += 1;
}
}
assert_eq!(
succeeded, MAX_SESSIONS,
"exactly max_sessions concurrent create_session calls must succeed"
);
}
#[tokio::test]
async fn test_enforce_session_cap_rewrites_capacity_marker_to_429() {
let app = axum::Router::new()
.route(
"/",
axum::routing::post(|| async {
(
axum::http::StatusCode::INTERNAL_SERVER_ERROR,
format!(
"Encounter an error when create session: {}: maximum concurrent \
HTTP sessions already active",
super::super::SESSION_CAP_MARKER
),
)
}),
)
.layer(axum::middleware::from_fn(super::super::enforce_session_cap));
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let server_task = tokio::spawn(async move {
axum::serve(listener, app).await.unwrap();
});
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
let response = raw_http_post(addr, "/", "", b"{}").await;
assert!(
response.starts_with("HTTP/1.1 429"),
"expected 429 for a marker-carrying 500, got: {response}"
);
assert!(
response.to_lowercase().contains("retry-after"),
"expected a Retry-After header, got: {response}"
);
server_task.abort();
}
#[tokio::test]
async fn test_enforce_session_cap_leaves_unrelated_500_untouched() {
let app = axum::Router::new()
.route(
"/",
axum::routing::post(|| async {
(
axum::http::StatusCode::INTERNAL_SERVER_ERROR,
"Encounter an error when create session: some unrelated failure",
)
}),
)
.layer(axum::middleware::from_fn(super::super::enforce_session_cap));
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let server_task = tokio::spawn(async move {
axum::serve(listener, app).await.unwrap();
});
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
let response = raw_http_post(addr, "/", "", b"{}").await;
assert!(
response.starts_with("HTTP/1.1 500"),
"unrelated 500s must not be rewritten to 429, got: {response}"
);
server_task.abort();
}
#[tokio::test]
async fn test_run_http_rejects_new_session_at_capacity_with_429() {
let probe = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = probe.local_addr().unwrap();
drop(probe);
let cfg = HttpConfig::new(addr, "/mcp").with_max_concurrent_sessions(1);
let server_task = tokio::spawn(super::super::run_http(
test_server(),
cfg,
super::super::ShutdownSignal::new(),
));
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
let initialize_body = br#"{"jsonrpc":"2.0","id":1,"method":"initialize","params":{"protocolVersion":"2025-06-18","capabilities":{},"clientInfo":{"name":"test","version":"0"}}}"#;
let accept_headers =
"Accept: application/json, text/event-stream\r\nContent-Type: application/json\r\n";
let first = raw_http_post(addr, "/mcp", accept_headers, initialize_body).await;
assert!(
first.starts_with("HTTP/1.1 200"),
"first initialize handshake should succeed, got: {first}"
);
let second = raw_http_post(addr, "/mcp", accept_headers, initialize_body).await;
assert!(
second.starts_with("HTTP/1.1 429"),
"second initialize handshake should be rejected once at capacity, got: {second}"
);
server_task.abort();
}
#[tokio::test]
async fn test_run_http_stateless_request_bypasses_session_cap() {
let probe = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = probe.local_addr().unwrap();
drop(probe);
let cfg = HttpConfig::new(addr, "/mcp").with_max_concurrent_sessions(1);
let server_task = tokio::spawn(super::super::run_http(
test_server(),
cfg,
super::super::ShutdownSignal::new(),
));
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
let accept_headers =
"Accept: application/json, text/event-stream\r\nContent-Type: application/json\r\n";
let legacy_initialize = br#"{"jsonrpc":"2.0","id":1,"method":"initialize","params":{"protocolVersion":"2025-06-18","capabilities":{},"clientInfo":{"name":"test","version":"0"}}}"#;
let legacy = raw_http_post(addr, "/mcp", accept_headers, legacy_initialize).await;
assert!(
legacy.starts_with("HTTP/1.1 200"),
"legacy initialize should succeed and consume the sole session slot, got: {legacy}"
);
let stateless_initialize = br#"{"jsonrpc":"2.0","id":2,"method":"initialize","params":{"protocolVersion":"2026-07-28","capabilities":{},"clientInfo":{"name":"test","version":"0"}}}"#;
let stateless = raw_http_post(addr, "/mcp", accept_headers, stateless_initialize).await;
assert!(
!stateless.starts_with("HTTP/1.1 429"),
"stateless requests must bypass the session cap entirely, got: {stateless}"
);
server_task.abort();
}
#[derive(Clone, Default)]
struct CapturedMessages(std::sync::Arc<std::sync::Mutex<Vec<String>>>);
impl<S: tracing::Subscriber> tracing_subscriber::Layer<S> for CapturedMessages {
fn on_event(
&self,
event: &tracing::Event<'_>,
_ctx: tracing_subscriber::layer::Context<'_, S>,
) {
struct MessageVisitor(String);
impl tracing::field::Visit for MessageVisitor {
fn record_debug(
&mut self,
field: &tracing::field::Field,
value: &dyn std::fmt::Debug,
) {
if field.name() == "message" {
self.0 = format!("{value:?}");
}
}
}
let mut visitor = MessageVisitor(String::new());
event.record(&mut visitor);
self.0.lock().unwrap().push(visitor.0);
}
}
#[tokio::test]
async fn test_run_http_non_loopback_bind_warns_to_use_reverse_proxy() {
use tracing_subscriber::layer::SubscriberExt as _;
let addr: SocketAddr = "0.0.0.0:0".parse().unwrap();
let cfg = HttpConfig::new(addr, "/mcp");
let captured = CapturedMessages::default();
let subscriber = tracing_subscriber::registry().with(captured.clone());
let guard = tracing::subscriber::set_default(subscriber);
let _ = tokio::time::timeout(
std::time::Duration::from_millis(200),
super::super::run_http(test_server(), cfg, super::super::ShutdownSignal::new()),
)
.await;
drop(guard);
let messages = captured.0.lock().unwrap().clone();
assert!(
messages.iter().any(|m| m.contains(
"place this endpoint behind a reverse proxy that enforces authentication"
)),
"expected reverse-proxy warning in captured tracing events, got: {messages:?}"
);
}
}
}