mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
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:
parent
9d60219a32
commit
7dbeb657a0
2 changed files with 126 additions and 1 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue