fix(server): cap /mcp body size (SRV-03)

The /mcp nest (to_mcp streamable HTTP service) collects the raw body
itself, so axum's extractor-based DefaultBodyLimit never applies and a
multi-GB chunked POST buffers entirely in memory (review-001 SRV-03).

Wrap the nest with an explicit counting-body middleware: 8 MiB cap
(headroom over the gateway's 2 MiB for JSON-RPC batch payloads),
Content-Length-declared oversizes rejected before reading, streaming
oversizes cut off mid-read and answered 413 (rmcp maps body-read errors
to 500, so the middleware sources the status itself).

Verification: cargo check (default + mcp + all-features); mcp test
suite pending tree unstuck from parallel agent's WIP.
This commit is contained in:
2026-08-29 08:47:50 +00:00
parent cc34c08e4e
commit a943d142c4
2 changed files with 199 additions and 0 deletions
+1
View File
@@ -43,6 +43,7 @@ uuid = { version = "1", features = ["v4"] }
futures = "0.3"
openapiv3 = "2"
http = "1"
http-body-util = "0.1"
url = "2"
bytes = "1"
jsonschema = { version = "0.46", default-features = false }
+198
View File
@@ -160,6 +160,7 @@ fn build_router(state: RouterState, extra_routes: Option<Router>) -> Router {
"/mcp",
crate::adapters::to_mcp_service(std::sync::Arc::new(dispatch)),
)
.layer(axum::middleware::from_fn(mcp_body_limit))
.layer(from_fn_with_state(
auth_state.clone(),
bearer_auth_middleware,
@@ -240,6 +241,115 @@ async fn rejected_reserved_path() -> axum::response::Response {
use axum::middleware::from_fn_with_state;
/// Cap the `/mcp` body at 8 MiB (feature `mcp`).
///
/// The nested rmcp `StreamableHttpService` collects the raw request body
/// itself (`expect_json` → `body.collect()`), bypassing axum's
/// extractor-based default limit: `DefaultBodyLimit` works by inserting
/// an extension that `FromRequest` extractors consult, so it has no
/// effect on a service that reads the body directly (rmcp 1.8
/// `server_side_http::expect_json` never checks it). This middleware is
/// both the cap and the status source: it wraps the body in a counting
/// stream that stops at [`MCP_BODY_LIMIT`] with an explicit error and
/// post-checks a exceedance flag to answer `413 Payload Too Large`,
/// replacing whatever the inner service answered (rmcp maps body-read
/// errors to `500`).
///
/// The limit is 8 MiB — headroom over the gateway's 2 MiB whole-body
/// default for JSON-RPC batch payloads on the MCP surface.
#[cfg(feature = "mcp")]
const MCP_BODY_LIMIT: usize = 8 * 1024 * 1024;
#[cfg(feature = "mcp")]
const MCP_BODY_LIMIT_EXCEEDED: &str = "mcp body limit exceeded";
#[cfg(feature = "mcp")]
async fn mcp_body_limit(
req: axum::extract::Request,
next: axum::middleware::Next,
) -> axum::response::Response {
use axum::response::IntoResponse;
let (parts, body) = req.into_parts();
if let Some(len) = parts
.headers
.get(http::header::CONTENT_LENGTH)
.and_then(|v| v.to_str().ok())
.and_then(|v| v.parse::<usize>().ok())
{
if len > MCP_BODY_LIMIT {
return (
axum::http::StatusCode::PAYLOAD_TOO_LARGE,
"Payload Too Large: /mcp body exceeds the 8 MiB limit",
)
.into_response();
}
}
let exceeded = Arc::new(std::sync::atomic::AtomicBool::new(false));
let counting = CountingBody {
inner: body.into_data_stream(),
remaining: MCP_BODY_LIMIT,
exceeded: Arc::clone(&exceeded),
};
let mut limited_req =
axum::extract::Request::from_parts(parts, axum::body::Body::from_stream(counting));
limited_req.extensions_mut().insert(exceeded.clone());
let response = next.run(limited_req).await;
if exceeded.load(std::sync::atomic::Ordering::Relaxed) {
return (
axum::http::StatusCode::PAYLOAD_TOO_LARGE,
"Payload Too Large: /mcp body exceeds the 8 MiB limit",
)
.into_response();
}
response
}
#[cfg(feature = "mcp")]
struct CountingBody {
inner: axum::body::BodyDataStream,
remaining: usize,
exceeded: Arc<std::sync::atomic::AtomicBool>,
}
#[cfg(feature = "mcp")]
impl futures::Stream for CountingBody {
type Item = Result<bytes::Bytes, std::io::Error>;
fn poll_next(
mut self: std::pin::Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Option<Self::Item>> {
let this = &mut *self;
match std::pin::Pin::new(&mut this.inner).poll_next(cx) {
std::task::Poll::Ready(Some(Ok(data))) => {
let len = data.len();
if len > this.remaining {
this.remaining = 0;
this.exceeded
.store(true, std::sync::atomic::Ordering::Relaxed);
return std::task::Poll::Ready(Some(Err(std::io::Error::other(
MCP_BODY_LIMIT_EXCEEDED,
))));
}
this.remaining -= len;
std::task::Poll::Ready(Some(Ok(data)))
}
std::task::Poll::Ready(Some(Err(e))) => {
let _ = e;
std::task::Poll::Ready(Some(Err(std::io::Error::other(MCP_BODY_LIMIT_EXCEEDED))))
}
std::task::Poll::Pending => std::task::Poll::Pending,
std::task::Poll::Ready(None) => std::task::Poll::Ready(None),
}
}
}
#[async_trait]
impl alkcall::core::types::ProtocolHandler for HttpAdapter {
fn alpn(&self) -> &'static [u8] {
@@ -576,6 +686,94 @@ mod tests {
let _ = server_task.await;
}
#[cfg(feature = "mcp")]
#[tokio::test]
async fn mcp_rejects_oversized_body_declared_content_length_with_413() {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let oversized = MCP_BODY_LIMIT + 1;
let head = format!(
"POST /mcp HTTP/1.1\r\nHost: localhost\r\nContent-Type: application/json\r\nAccept: application/json, text/event-stream\r\nContent-Length: {oversized}\r\nConnection: close\r\n\r\n"
);
let adapter = HttpAdapter::new(provider(), empty_registry());
let (client, server) = tokio::io::duplex(64 * 1024);
let conn = Connection::from_bidi(server, b"http/1.1".to_vec(), None);
let auth = AuthContext::anonymous(b"http/1.1");
let server_task = tokio::spawn(async move {
let _ = ProtocolHandler::handle(&adapter, conn, &auth).await;
});
let mut client = client;
client.write_all(head.as_bytes()).await.unwrap();
let mut response = Vec::new();
let _ = tokio::time::timeout(
std::time::Duration::from_secs(5),
client.read_to_end(&mut response),
)
.await
.expect("read timed out")
.unwrap();
let text = String::from_utf8_lossy(&response);
assert!(
text.starts_with("HTTP/1.1 413 Payload Too Large"),
"declared oversized body got: {text}"
);
let _ = server_task.await;
}
#[cfg(feature = "mcp")]
#[tokio::test]
async fn mcp_rejects_oversized_chunked_body_with_413() {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let adapter = HttpAdapter::new(provider(), empty_registry());
let (client, server) = tokio::io::duplex(64 * 1024);
let conn = Connection::from_bidi(server, b"http/1.1".to_vec(), None);
let auth = AuthContext::anonymous(b"http/1.1");
let server_task = tokio::spawn(async move {
let _ = ProtocolHandler::handle(&adapter, conn, &auth).await;
});
let mut client = client;
client
.write_all(
b"POST /mcp HTTP/1.1\r\nHost: localhost\r\nContent-Type: application/json\r\nAccept: application/json, text/event-stream\r\nTransfer-Encoding: chunked\r\nConnection: close\r\n\r\n",
)
.await
.unwrap();
let chunk_size = 64 * 1024;
for _ in 0..(MCP_BODY_LIMIT / chunk_size) + 1 {
let header = format!("{chunk_size:x}\r\n");
client.write_all(header.as_bytes()).await.unwrap();
let chunk = vec![b'a'; chunk_size];
client.write_all(&chunk).await.unwrap();
client.write_all(b"\r\n").await.unwrap();
}
client.write_all(b"0\r\n\r\n").await.unwrap();
let mut response = Vec::new();
let _ = tokio::time::timeout(
std::time::Duration::from_secs(10),
client.read_to_end(&mut response),
)
.await
.expect("read timed out")
.unwrap();
let text = String::from_utf8_lossy(&response);
assert!(
text.starts_with("HTTP/1.1 413 Payload Too Large"),
"chunked oversized body got: {text}"
);
let _ = server_task.await;
}
struct StaticProvider;
impl IdentityProvider for StaticProvider {
fn resolve_from_fingerprint(&self, _: &str) -> Option<alkcall::core::auth::Identity> {