From 608840b82368fdff599041276c23947a42ef45e0 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 15 Sep 2026 18:41:10 +0000 Subject: [PATCH] fix(mcp): expose peeked JSON-RPC body to auth so discovery methods skip budget checks Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/auth/user_api_key_auth_mcp.py | 22 +++- .../proxy/_experimental/mcp_server/server.py | 41 +++--- .../mcp_server/test_mcp_server.py | 119 ++++++++++++++++++ 3 files changed, 156 insertions(+), 26 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 b0d57cb6228..93e41d63c5d 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 @@ -67,9 +67,24 @@ if TYPE_CHECKING: from litellm.proxy.utils import PrismaClient +MCP_PEEKED_BODY_SCOPE_KEY: Final = "litellm_mcp_peeked_body" + _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: Final = Request(scope=scope) + peeked_body: Final[bytes] = scope.get(MCP_PEEKED_BODY_SCOPE_KEY, b"{}") + + async def mock_body(): + return peeked_body + + request.body = mock_body + return request + + def _as_list(values: Sequence[str] | None) -> list[str] | None: # mutable-ok: resolver returns a list """Widen a read-only allowlist back to the mutable list the resolver's own contract returns, preserving the ``None`` that means "no restriction".""" @@ -440,12 +455,7 @@ class MCPRequestHandler: if mcp_servers_header == "" or (mcp_servers is not None and len(mcp_servers) == 0): mcp_servers = [] # Create a proper Request object with mock body method to avoid ASGI receive channel issues - request: Final = Request(scope=scope) - - async def mock_body(): - return b"{}" - - request.body = mock_body + request: Final = _admission_request(scope) # Inline import — auth_utils participates in a proxy import cycle. from litellm.proxy.auth.auth_utils import ( # noqa: PLC0415 get_request_route, diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 7feb1fd468d..b71c2acf91e 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -33,6 +33,7 @@ from litellm.llms.custom_httpx.http_handler import ( httpxSpecialProvider, ) from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( + MCP_PEEKED_BODY_SCOPE_KEY, MCPRequestHandler, _is_mcp_admitted_user_subject, ) @@ -4389,6 +4390,25 @@ if MCP_AVAILABLE: """Handle MCP requests through StreamableHTTP.""" try: path: Final[str] = scope.get("path", "") + consumed_messages: list[Message] = [] + body = b"" + if scope.get("method") == "POST": + consumed_messages, body = await _read_request_body_for_routing(receive) + if consumed_messages: + original_receive: Final = receive + + async def wrapped_receive(): + if consumed_messages: + return consumed_messages.pop(0) + return await original_receive() + + receive = wrapped_receive + try: + if isinstance(json.loads(body), dict): + scope[MCP_PEEKED_BODY_SCOPE_KEY] = body + except (json.JSONDecodeError, TypeError): + pass + is_initialize: Final = _is_initialize_request(body) ( user_api_key_auth, mcp_auth_header, @@ -4467,8 +4487,6 @@ if MCP_AVAILABLE: # - No session ID + initialize → stateful (so client gets mcp-session-id) # - No session ID + other → stateless (curl, Inspector, Notion) session_id = _get_session_id_from_scope(scope) - is_initialize = False - consumed_messages: list[Message] = [] # Owner-binding: a live stateful session may only be driven by the # caller that created it. Reject mismatches with 403 so a leaked @@ -4476,8 +4494,7 @@ if MCP_AVAILABLE: # # Run before ``_handle_stale_mcp_session`` so a non-owner cannot # force-clean another caller's residual tracking entries via a - # stale DELETE, and before peeking the request body so the 403 - # response sees a pristine ``receive`` channel. + # stale DELETE. if session_id: expected_owner: Final = _stateful_session_owners.get(session_id) request_owner = _owner_fingerprint_for(user_api_key_auth, oauth2_headers, _client_ip) @@ -4506,11 +4523,6 @@ if MCP_AVAILABLE: return session_id = _get_session_id_from_scope(scope) - body = b"" - if scope.get("method") == "POST": - consumed_messages, body = await _read_request_body_for_routing(receive) - is_initialize = _is_initialize_request(body) - use_stateful: Final = bool(session_id or is_initialize) target_manager: Final = session_manager_stateful if use_stateful else session_manager_stateless @@ -4539,17 +4551,6 @@ if MCP_AVAILABLE: await too_many_response(scope, receive, send) return - # Replay body messages if we consumed them for peeking - original_receive: Final = receive - if consumed_messages: - - async def wrapped_receive(): - if consumed_messages: - return consumed_messages.pop(0) - return await original_receive() - - receive = wrapped_receive - # Serialize requests on the same stateful session so concurrent # callers don't clobber each other's auth context mid-flight. # 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 f5e4a420496..4d60b338c21 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 @@ -2127,6 +2127,125 @@ async def test_mcp_routing_chunked_initialize_to_stateful(): ) +@pytest.mark.asyncio +async def test_mcp_routing_stashes_peeked_body_for_auth(): + """The auth-time Request built by ``process_mcp_request`` cannot read the ASGI + receive channel; the routing peek must stash the real JSON-RPC body in scope so + ``is_mcp_discovery_request`` can see ``method`` and an over-budget key is not + 429'd on zero-spend ``initialize`` / ``tools/list``.""" + try: + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( + MCP_PEEKED_BODY_SCOPE_KEY, + _admission_request, + ) + from litellm.proxy._experimental.mcp_server.server import ( + handle_streamable_http_mcp, + session_manager_stateful, + session_manager_stateless, + ) + from litellm.proxy.common_utils.http_parsing_utils import _read_request_body + except ImportError: + pytest.skip("MCP server not available") + + jsonrpc_body: Final = b'{"jsonrpc":"2.0","id":2,"method":"tools/list","params":{}}' + scope = { + "type": "http", + "method": "POST", + "path": "/mcp", + "headers": [ + (b"content-type", b"application/json"), + (b"authorization", b"Bearer test-key"), + ], + } + receive = AsyncMock(side_effect=[{"type": "http.request", "body": jsonrpc_body, "more_body": False}]) + send = AsyncMock() + stateless_called = [] + + async def stateless_handle(s, r, se): + stateless_called.append(1) + + with ( + patch( + "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", + new_callable=AsyncMock, + return_value=(MagicMock(), None, None, None, None, None), + ), + patch( + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + new_callable=AsyncMock, + return_value=[MagicMock()], + ), + 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_stateless, "_server_instances", {}), + patch.object(session_manager_stateful, "_server_instances", {}), + ): + 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 + request_data: Final = await _read_request_body(_admission_request(scope)) + assert request_data.get("method") == "tools/list" + + +@pytest.mark.asyncio +async def test_mcp_routing_batch_body_is_not_stashed_for_auth(): + """A JSON-RPC batch (array) body must not be exposed to auth: it can carry a + spend-bearing ``tools/call`` alongside discovery methods, so it fails closed + and stays budget-enforced.""" + try: + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( + MCP_PEEKED_BODY_SCOPE_KEY, + _admission_request, + ) + from litellm.proxy._experimental.mcp_server.server import ( + handle_streamable_http_mcp, + session_manager_stateless, + ) + from litellm.proxy.common_utils.http_parsing_utils import _read_request_body + except ImportError: + pytest.skip("MCP server not available") + + batch_body: Final = b'[{"jsonrpc":"2.0","id":1,"method":"tools/list"},{"jsonrpc":"2.0","id":2,"method":"tools/call"}]' + scope = { + "type": "http", + "method": "POST", + "path": "/mcp", + "headers": [ + (b"content-type", b"application/json"), + (b"authorization", b"Bearer test-key"), + ], + } + receive = AsyncMock(side_effect=[{"type": "http.request", "body": batch_body, "more_body": False}]) + send = AsyncMock() + + with ( + patch( + "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", + new_callable=AsyncMock, + return_value=(MagicMock(), None, None, None, None, None), + ), + patch( + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + new_callable=AsyncMock, + return_value=[MagicMock()], + ), + patch( + "litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", + True, + ), + patch.object(session_manager_stateless, "handle_request", new=AsyncMock()), + patch.object(session_manager_stateless, "_server_instances", {}), + ): + await handle_streamable_http_mcp(scope, receive, send) + + assert MCP_PEEKED_BODY_SCOPE_KEY not in scope + assert await _read_request_body(_admission_request(scope)) == {} + + @pytest.mark.asyncio async def test_mcp_routing_caps_body_peek_for_oversized_chunked_body(): """