mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
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:
parent
4a25b2f499
commit
9a4d21b1a6
4 changed files with 125 additions and 23 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"{}"
|
||||
|
|
|
|||
|
|
@ -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)) == {}
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue