fix(mcp): cap routing-peek body size to bound pre-dispatch memory

Authenticated clients that POST without an mcp-session-id forced the proxy
to buffer the entire request body before routing, since the peek loop
drained every body chunk to decide whether the JSON-RPC method was
'initialize'. Cap the peek at 4 KB (more than enough for an initialize
envelope) and let the remainder stream through wrapped_receive into the
downstream handler.
This commit is contained in:
Sameerlite 2026-05-13 16:57:07 +00:00
parent 9d60219a32
commit 7dbeb657a0
2 changed files with 126 additions and 1 deletions

View file

@ -72,6 +72,11 @@ _byok_cred_cache: Dict[Tuple[str, str], Tuple[Optional[str], float]] = {}
_BYOK_CRED_CACHE_TTL = 60 # seconds
_BYOK_CRED_CACHE_MAX_SIZE = 4096 # cap to prevent unbounded growth
_STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS = 30 * 60
# Maximum bytes to peek when sniffing the JSON-RPC method on a no-session-id
# POST. An `initialize` envelope is a few hundred bytes; capping the peek
# prevents an authenticated client from forcing the proxy to buffer an
# arbitrarily large body just to make a routing decision.
_MCP_ROUTING_PEEK_MAX_BYTES = 4096
def _invalidate_byok_cred_cache(user_id: str, server_id: str) -> None:
@ -2782,10 +2787,21 @@ if MCP_AVAILABLE:
receive: Receive,
) -> Tuple[List[Message], bytes]:
"""
Consume request body messages for routing, returning them for replay.
Read just enough of the request body to decide whether this is a
JSON-RPC ``initialize`` call. Returns the consumed ASGI messages so
the caller can replay them faithfully to the downstream handler, and
the peeked body bytes (capped at ``_MCP_ROUTING_PEEK_MAX_BYTES``).
Stops reading from the wire as soon as either (a) we have peeked
``_MCP_ROUTING_PEEK_MAX_BYTES`` of body, or (b) the body is complete.
The remainder of an oversized body is streamed lazily through
``wrapped_receive`` in the caller — so an authenticated client cannot
force the proxy to buffer an arbitrarily large payload just to make a
routing decision.
"""
consumed_messages: List[Message] = []
body_chunks: List[bytes] = []
peeked_bytes = 0
while True:
message = await receive()
@ -2797,10 +2813,16 @@ if MCP_AVAILABLE:
body = message.get("body", b"") or b""
if body:
body_chunks.append(body)
peeked_bytes += len(body)
if not message.get("more_body", False):
break
if peeked_bytes >= _MCP_ROUTING_PEEK_MAX_BYTES:
# Stop draining; downstream replay will pull remaining chunks
# directly from the original `receive` via wrapped_receive.
break
return consumed_messages, b"".join(body_chunks)
async def _handle_stale_mcp_session(

View file

@ -1307,6 +1307,109 @@ async def test_mcp_routing_chunked_initialize_to_stateful():
), "chunked initialize (no session) should route to stateful, not stateless"
@pytest.mark.asyncio
async def test_mcp_routing_caps_body_peek_for_oversized_chunked_body():
"""
A no-session-id POST with a very large chunked body should not force
the proxy to buffer the entire body just to decide routing — the peek
should stop once ``_MCP_ROUTING_PEEK_MAX_BYTES`` worth of body has been
consumed, and the remaining chunks should stream through the original
receive into the downstream handler.
"""
try:
from litellm.proxy._experimental.mcp_server import server as mcp_server
from litellm.proxy._experimental.mcp_server.server import (
handle_streamable_http_mcp,
session_manager_stateful,
session_manager_stateless,
)
except ImportError:
pytest.skip("MCP server not available")
peek_cap = mcp_server._MCP_ROUTING_PEEK_MAX_BYTES
# First chunk fills the peek budget; subsequent chunks are oversized payload.
first_chunk = b"x" * peek_cap
oversized_tail = [b"y" * 65536 for _ in range(4)]
messages = [
{"type": "http.request", "body": first_chunk, "more_body": True},
*[
{"type": "http.request", "body": chunk, "more_body": True}
for chunk in oversized_tail
],
{"type": "http.request", "body": b"", "more_body": False},
]
receive_calls = {"count": 0}
async def receive():
idx = receive_calls["count"]
receive_calls["count"] += 1
return messages[idx]
scope = {
"type": "http",
"method": "POST",
"path": "/mcp/progress_test",
"headers": [
(b"content-type", b"application/json"),
(b"authorization", b"Bearer test-key"),
],
}
send = AsyncMock()
stateless_received_chunks = []
receive_count_at_dispatch = {"value": -1}
async def stateless_handle(s, r, se):
# Snapshot how many wire reads happened BEFORE dispatch — the cap
# check is meaningful only against pre-dispatch consumption.
receive_count_at_dispatch["value"] = receive_calls["count"]
# Drain the wrapped receive the same way the SDK would.
while True:
msg = await r()
if msg.get("type") != "http.request":
break
stateless_received_chunks.append(msg.get("body", b"") or b"")
if not msg.get("more_body", False):
break
async def stateful_handle(s, r, se):
raise AssertionError("non-initialize POST should not reach stateful manager")
with (
patch(
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
new_callable=AsyncMock,
return_value=(MagicMock(), None, ["progress_test"], None, None, None),
),
patch("litellm.proxy._experimental.mcp_server.server.set_auth_context"),
patch(
"litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED",
True,
),
patch.object(
session_manager_stateless, "handle_request", side_effect=stateless_handle
),
patch.object(
session_manager_stateful, "handle_request", side_effect=stateful_handle
),
patch.object(session_manager_stateless, "_server_instances", {}),
patch.object(session_manager_stateful, "_server_instances", {}),
):
await handle_streamable_http_mcp(scope, receive, send)
# The routing peek must stop pulling from the wire once the cap is reached.
# Without the cap fix, every chunk would have been pulled before dispatch,
# so this assertion guards against unbounded pre-dispatch buffering.
assert receive_count_at_dispatch["value"] == 1, (
"routing should stop reading after the peek cap is filled, "
f"but consumed {receive_count_at_dispatch['value']} chunks before dispatching"
)
# All chunks must still reach the downstream handler via replay+stream.
total_streamed = sum(len(b) for b in stateless_received_chunks)
assert total_streamed == len(first_chunk) + sum(len(b) for b in oversized_tail)
@pytest.mark.asyncio
async def test_stateful_mcp_requests_refresh_session_auth_context():
"""