fix(mcp): defer POST body peek until after client allowlist check

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
Devin AI 2026-09-20 00:11:31 +00:00
parent 4a25b2f499
commit 9a4d21b1a6
4 changed files with 125 additions and 23 deletions

View file

@ -1,5 +1,5 @@
import re
from collections.abc import Mapping, Sequence
from collections.abc import Awaitable, Callable, Mapping, Sequence
from dataclasses import dataclass
from datetime import datetime, timezone
from types import MappingProxyType
@ -72,13 +72,16 @@ _EMPTY_TOOLSET_GRANTS: Final[Mapping[str, Sequence[str]]] = MappingProxyType({})
def _admission_request(scope: Scope) -> Request:
"""Request whose ``body()`` serves the routing layer's peeked JSON-RPC bytes
(``b"{}"`` when nothing was peeked) instead of the ASGI receive channel."""
"""Request whose ``body()`` serves the routing layer's lazily peeked JSON-RPC
bytes (``b"{}"`` when no peek callable was stashed) instead of the ASGI
receive channel."""
request: Final = Request(scope=scope)
peeked_body: Final[bytes] = scope.get(MCP_PEEKED_BODY_SCOPE_KEY, b"{}")
peeked_body: Final[Callable[[], Awaitable[bytes]] | None] = scope.get(MCP_PEEKED_BODY_SCOPE_KEY)
async def mock_body() -> bytes:
return peeked_body
if peeked_body is None:
return b"{}"
return await peeked_body()
request.body = mock_body
return request

View file

@ -3957,7 +3957,7 @@ if MCP_AVAILABLE:
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
``_LazyPeekedBody.receive`` in the caller — so an authenticated client cannot
force the proxy to buffer an arbitrarily large payload just to make a
routing decision.
"""
@ -3995,6 +3995,29 @@ if MCP_AVAILABLE:
return consumed_messages, b"".join(body_chunks)
class _LazyPeekedBody:
"""Reads the POST body for routing and auth only when first asked, so an
allowlist rejection never touches the receive channel. Whatever it
consumed is replayed to the downstream handler through ``receive``,
which falls through to the wire once the replay buffer is drained."""
__slots__ = ("_consumed_messages", "_peeked_body", "_receive")
def __init__(self, receive: Receive) -> None:
self._receive = receive
self._consumed_messages: list[Message] | None = None
self._peeked_body: bytes = b""
async def body(self) -> bytes:
if self._consumed_messages is None:
self._consumed_messages, self._peeked_body = await _read_request_body_for_routing(self._receive)
return self._peeked_body
async def receive(self) -> Message:
if self._consumed_messages:
return self._consumed_messages.pop(0)
return await self._receive()
async def _handle_stale_mcp_session(
scope: Scope,
receive: Receive,
@ -4550,23 +4573,14 @@ if MCP_AVAILABLE:
)(scope, receive, send)
return
path: Final[str] = scope.get("path", "")
consumed_messages: list[Message] = [] # mutable-ok: replay buffer for peeked ASGI messages
body = b""
if scope.get("method") == "POST":
consumed_messages, body = await _read_request_body_for_routing(receive)
if consumed_messages:
original_receive: Final = receive
peek: Final = _LazyPeekedBody(receive) if scope.get("method") == "POST" else None
if peek is not None:
receive = peek.receive # rebind-ok: replay peeked ASGI messages to the downstream handler
async def wrapped_receive() -> Message:
if consumed_messages:
return consumed_messages.pop(0)
return await original_receive()
async def peeked_json_object_body() -> bytes:
return _peeked_json_object_body(await peek.body()) or b"{}"
receive = wrapped_receive # rebind-ok: replay peeked ASGI messages to the downstream handler
peeked_object_body: Final = _peeked_json_object_body(body)
if peeked_object_body is not None:
scope[MCP_PEEKED_BODY_SCOPE_KEY] = peeked_object_body
is_initialize: Final = _is_initialize_request(body)
scope[MCP_PEEKED_BODY_SCOPE_KEY] = peeked_json_object_body
(
user_api_key_auth,
mcp_auth_header,
@ -4576,6 +4590,8 @@ if MCP_AVAILABLE:
raw_headers,
) = await extract_mcp_auth_context(scope, path)
reject_disallowed_mcp_client(StarletteRequest(scope).headers, user_api_key_auth)
body: Final = await peek.body() if peek is not None else b""
is_initialize: Final = _is_initialize_request(body)
scoped_server_endpoint: Final = len(_get_mcp_servers_in_path(path) or []) == 1
# Extract client IP for MCP access control

View file

@ -9437,3 +9437,28 @@ class TestScopedSessionAdmission:
def test_scope_field_cannot_be_forged_through_construction(self):
forged = UserAPIKeyAuth(user_id="u1", mcp_session_resource_server_id="any-server")
assert forged.mcp_session_resource_server_id is None
@pytest.mark.asyncio
async def test_admission_request_body_serves_stashed_peek_callable():
from litellm.constants import MCP_PEEKED_BODY_SCOPE_KEY
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import _admission_request
jsonrpc_body = b'{"jsonrpc":"2.0","id":1,"method":"tools/list","params":{}}'
async def peek() -> bytes:
return jsonrpc_body
with_peek = _admission_request(
{
"type": "http",
"method": "POST",
"path": "/mcp",
"headers": [],
MCP_PEEKED_BODY_SCOPE_KEY: peek,
}
)
assert await with_peek.body() == jsonrpc_body
without_peek = _admission_request({"type": "http", "method": "POST", "path": "/mcp", "headers": []})
assert await without_peek.body() == b"{}"

View file

@ -2223,6 +2223,64 @@ async def test_streamable_http_admits_listed_or_unrestricted_clients_and_hands_t
send.assert_not_awaited()
@pytest.mark.asyncio
async def test_streamable_http_admission_auth_sees_peeked_jsonrpc_body() -> None:
"""The lazy peek must expose the JSON-RPC body to admission auth (via the
scope callable) while the replayed receive still hands the full body to the
downstream session manager."""
from litellm.constants import MCP_PEEKED_BODY_SCOPE_KEY
from litellm.proxy._experimental.mcp_server import server as mcp_module
scope: Final[Scope] = {"type": "http", "method": "POST", "path": "/mcp", "headers": []}
receive: Final = AsyncMock(
side_effect=[
{"type": "http.request", "body": _INITIALIZE[:20], "more_body": True},
{"type": "http.request", "body": _INITIALIZE[20:], "more_body": False},
]
)
send: Final = AsyncMock()
downstream_bodies: Final[list[bytes]] = []
auth_bodies: Final[list[bytes]] = []
async def extract(auth_scope: Scope, _: str):
auth_bodies.append(await auth_scope[MCP_PEEKED_BODY_SCOPE_KEY]())
return (UserAPIKeyAuth(user_id="allowlist-user", jwt_claims=_LISTED_JWT), None, None, None, None, {})
async def handle_request(_: Scope, downstream_receive: Receive, __: Send) -> None:
downstream_bodies.append(await _drain_body(downstream_receive))
stateful_handle: Final = AsyncMock(side_effect=handle_request)
stateless_handle: Final = AsyncMock()
with (
patch( # test-quality-ok: the ASGI handler resolves auth through a module-level function; no injection seam
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
new_callable=AsyncMock,
side_effect=extract,
),
patch( # test-quality-ok: module flag guarding lazy session-manager startup; no injection seam
"litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", True
),
patch( # test-quality-ok: the allowlist is read off this module global; no injection seam
"litellm.proxy.proxy_server.general_settings", _ALLOWLIST_SETTINGS
),
patch( # test-quality-ok: session managers are module singletons; the downstream call is the observable
"litellm.proxy._experimental.mcp_server.server.session_manager_stateful",
SimpleNamespace(handle_request=stateful_handle),
),
patch( # test-quality-ok: session managers are module singletons; the downstream call is the observable
"litellm.proxy._experimental.mcp_server.server.session_manager_stateless",
SimpleNamespace(handle_request=stateless_handle),
),
):
await mcp_module.handle_streamable_http_mcp(scope, receive, send)
assert auth_bodies == [_INITIALIZE]
assert downstream_bodies == [_INITIALIZE]
stateless_handle.assert_not_awaited()
send.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize(
("jwt_claims", "headers", "admitted"),
@ -2434,7 +2492,7 @@ async def test_mcp_routing_stashes_peeked_body_for_auth():
await handle_streamable_http_mcp(scope, receive, send)
assert stateless_called, "tools/list without a session should route to the stateless manager"
assert scope[MCP_PEEKED_BODY_SCOPE_KEY] == jsonrpc_body
assert await scope[MCP_PEEKED_BODY_SCOPE_KEY]() == jsonrpc_body
request_data: Final = await _read_request_body(_admission_request(scope))
assert request_data.get("method") == "tools/list"
@ -2494,7 +2552,7 @@ async def test_mcp_routing_batch_body_is_not_stashed_for_auth():
):
await handle_streamable_http_mcp(scope, receive, send)
assert MCP_PEEKED_BODY_SCOPE_KEY not in scope
assert await scope[MCP_PEEKED_BODY_SCOPE_KEY]() == b"{}"
assert await _read_request_body(_admission_request(scope)) == {}