From 9a4d21b1a68a8b0510876dd55ba47fe14fc1ce50 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sun, 20 Sep 2026 00:11:31 +0000 Subject: [PATCH] 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> --- .../mcp_server/auth/user_api_key_auth_mcp.py | 13 ++-- .../proxy/_experimental/mcp_server/server.py | 48 +++++++++----- .../auth/test_user_api_key_auth_mcp.py | 25 ++++++++ .../mcp_server/test_mcp_server.py | 62 ++++++++++++++++++- 4 files changed, 125 insertions(+), 23 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index b6a82b4df78..22ce7906615 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -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 diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 0428b068e4a..0af26acb920 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -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 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index 4380df194ed..fb5a33b040b 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py @@ -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"{}" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 5a7c0472fee..dda6e281841 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -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)) == {}