diff --git a/litellm/proxy/_experimental/mcp_server/caller_sign_in.py b/litellm/proxy/_experimental/mcp_server/caller_sign_in.py index 74637972f75..1d6d75b438f 100644 --- a/litellm/proxy/_experimental/mcp_server/caller_sign_in.py +++ b/litellm/proxy/_experimental/mcp_server/caller_sign_in.py @@ -15,7 +15,7 @@ never imports a concrete provider. from __future__ import annotations import itertools -from collections.abc import Mapping +from collections.abc import Awaitable, Callable, Mapping from dataclasses import dataclass from typing import TYPE_CHECKING, Final, Protocol, runtime_checkable @@ -162,7 +162,7 @@ async def preflight_caller_sign_in( *, root_path: str, resource_metadata: str | None, - connecting: bool, + connecting: Callable[[], Awaitable[bool]], ) -> None: """Run every provider's connect-time check against the subject token, so a bearer the IdP will reject surfaces as a challenge here rather than a JSON-RPC error at the first tool call. A fail-closed @@ -186,9 +186,8 @@ async def preflight_caller_sign_in( ) case Unavailable(fail_open=True): continue - case Unavailable(detail=detail, fail_open=False) if connecting: - raise HTTPException(status_code=503, detail=detail) - case Unavailable(): - continue + case Unavailable(detail=detail, fail_open=False): + if await connecting(): + raise HTTPException(status_code=503, detail=detail) case _ as verdict: assert_never(verdict) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 47d41325e74..4d9bdf41696 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -8,12 +8,13 @@ import asyncio import contextlib import contextvars import hashlib +import itertools import json import os import time import types from collections import Counter -from collections.abc import AsyncGenerator, AsyncIterator, Callable, Iterable, Mapping, Sequence +from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Callable, Iterable, Iterator, Mapping, Sequence from typing import TYPE_CHECKING, Final, NoReturn, Protocol import httpx @@ -1426,6 +1427,41 @@ if MCP_AVAILABLE: return consumed_messages, b"".join(body_chunks) + class _ConnectBodyPeek: + """Reads a session-less ``POST`` body only once a gate asks whether it is ``initialize``, so a challenge + that needs no body still answers before the body arrives; consumed messages replay through ``receive``.""" + + def __init__(self, receive: Receive, peekable: bool) -> None: + self._receive: Final = receive + self._peekable: Final = peekable + self._body: bytes | None = None + self._replay: Iterator[Message] = iter(()) + + async def read(self) -> bytes: + messages, body = await _read_request_body_for_routing(self._receive) + self._replay = itertools.chain(self._replay, messages) + return body + + async def body(self) -> bytes: + if not self._peekable: + return b"" + if self._body is None: + self._body = await self.read() + return self._body + + async def connecting(self) -> bool: + return _is_initialize_request(await self.body()) + + async def receive(self) -> Message: + replayed: Final = next(self._replay, None) + return replayed if replayed is not None else await self._receive() + + def _known_connecting(value: bool) -> Callable[[], Awaitable[bool]]: + async def answer() -> bool: + return value + + return answer + async def _handle_stale_mcp_session( scope: Scope, receive: Receive, @@ -1608,7 +1644,7 @@ if MCP_AVAILABLE: mcp_server_auth_headers: dict[str, dict[str, str]] | None, user_api_key_auth: UserAPIKeyAuth | None, client_ip: str | None, - connecting: bool, + connecting: Callable[[], Awaitable[bool]], allowed_server_ids: set[str] | None = None, raw_headers: Mapping[str, str] | None = None, ) -> None: @@ -2073,25 +2109,13 @@ if MCP_AVAILABLE: toolset_allowed_server_ids = await _toolset_server_ids(active_toolset_id) named_session_id: Final = _get_session_id_from_scope(scope) - names_live_session: Final = named_session_id is not None and ( - named_session_id in _stateful_session_owners or named_session_id in _stateful_server_instances() + names_live_session: Final = ( + named_session_id is not None and named_session_id in _stateful_server_instances() ) - consumed_messages, connect_body = ( - await _read_request_body_for_routing(receive) - if scope.get("method") == "POST" and not names_live_session - else ([], b"") + connect_peek: Final = _ConnectBodyPeek( + receive, peekable=scope.get("method") == "POST" and not names_live_session ) - connecting: Final = _is_initialize_request(connect_body) - - # Replay body messages if we consumed them for peeking - original_receive: Final = receive - - async def wrapped_receive(): - if consumed_messages: - return consumed_messages.pop(0) - return await original_receive() - - receive = wrapped_receive + receive = connect_peek.receive # https://datatracker.ietf.org/doc/html/rfc9728#name-www-authenticate-response # Must run after toolset scoping so the challenge set is derived @@ -2105,7 +2129,7 @@ if MCP_AVAILABLE: mcp_server_auth_headers=mcp_server_auth_headers, user_api_key_auth=user_api_key_auth, client_ip=_client_ip, - connecting=connecting, + connecting=connect_peek.connecting, allowed_server_ids=toolset_allowed_server_ids, raw_headers=raw_headers, ) @@ -2177,13 +2201,10 @@ if MCP_AVAILABLE: return session_id = _get_session_id_from_scope(scope) - session_messages, session_body = ( - await _read_request_body_for_routing(receive) - if scope.get("method") == "POST" and names_live_session - else ([], b"") + session_body: Final = ( + await connect_peek.read() if scope.get("method") == "POST" and names_live_session else b"" ) - consumed_messages.extend(session_messages) - body: Final = connect_body or session_body + body: Final = await connect_peek.body() or session_body is_initialize: Final = _is_initialize_request(body) use_stateful: Final = bool(session_id or is_initialize) @@ -2443,7 +2464,7 @@ if MCP_AVAILABLE: mcp_server_auth_headers=mcp_server_auth_headers, user_api_key_auth=user_api_key_auth, client_ip=_sse_client_ip, - connecting=scope["method"] == "GET", + connecting=_known_connecting(scope["method"] == "GET"), allowed_server_ids=toolset_allowed_server_ids, raw_headers=raw_headers, ) diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py index bfd759bc2a1..85b631251ef 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py @@ -53,6 +53,10 @@ def test_mcp_available_on_sdk2(): assert MCP_AVAILABLE is True +async def _connecting() -> bool: + return True + + def _rendered_log_message(call): message = str(call.args[0]) values = call.args[1:] @@ -3968,6 +3972,159 @@ async def test_initialize_naming_a_stale_session_still_meets_the_fail_closed_con assert sent_messages == [] +@pytest.mark.asyncio +async def test_connect_challenge_answers_before_the_body_is_read(monkeypatch): + """A session-less ``POST`` to a gated server without a subject token gets its RFC 9728 challenge straight + away: the gate must not wait for the body it never needs, so a client that withholds it still sees 401.""" + 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, + ) + except ImportError: + pytest.skip("MCP server not available") + + monkeypatch.delenv("SERVER_ROOT_PATH", raising=False) + server = _make_obo_server("obo_server") + mcp_operations.global_mcp_server_manager.registry.update({server.server_id: server}) + scope = { + "type": "http", + "method": "POST", + "scheme": "http", + "path": "/mcp/obo_server", + "root_path": "", + "query_string": b"", + "server": ("gw.example", 4000), + "client": ("10.0.0.7", 51000), + "headers": [ + (b"host", b"gw.example:4000"), + (b"content-type", b"application/json"), + (b"x-litellm-api-key", b"sk-litellm-virtual-key"), + ], + } + withheld_body: asyncio.Queue[Message] = asyncio.Queue() + sent_messages: list[Message] = [] + + async def capture_send(message: Message) -> None: + sent_messages.append(message) + + handle_request_mock = AsyncMock() + with ( + patch( + "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", + new_callable=AsyncMock, + return_value=( + UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"), + None, + ["obo_server"], + None, + None, + {"x-litellm-api-key": "sk-litellm-virtual-key"}, + ), + ), + patch.object( # test-quality-ok: the key's grant list lives in the DB; the real router runs on it + mcp_operations.global_mcp_server_manager, + "get_allowed_mcp_servers", + AsyncMock(return_value=[server.server_id]), + ), + patch.object(mcp_server, "_SESSION_MANAGERS_INITIALIZED", True), + patch.object(session_manager_stateful, "handle_request", side_effect=handle_request_mock), + patch.object(session_manager_stateful, "_server_instances", {}), + pytest.raises(HTTPException) as exc, + ): + await asyncio.wait_for(handle_streamable_http_mcp(scope, withheld_body.get, capture_send), timeout=2) + + assert exc.value.status_code == 401 + assert (exc.value.headers or {})["WWW-Authenticate"].startswith( + 'Bearer resource_metadata="http://gw.example:4000/.well-known/oauth-protected-resource/mcp/obo_server"' + ), exc.value.headers + handle_request_mock.assert_not_awaited() + assert sent_messages == [] + + +@pytest.mark.asyncio +async def test_initialize_naming_a_session_whose_transport_is_already_gone_still_meets_the_fail_closed_connect_gate(): + """While idle purge or cap eviction is still terminating a transport, its owner entry outlives the transport; + an ``initialize`` retried with that id is still a new connection and meets the fail-closed gate.""" + try: + from litellm.proxy._experimental.mcp_server import server as mcp_server + from litellm.proxy._experimental.mcp_server.caller_sign_in import Unavailable + from litellm.proxy._experimental.mcp_server.server import ( + handle_streamable_http_mcp, + session_manager_stateful, + ) + except ImportError: + pytest.skip("MCP server not available") + + server = _catalog_server() + guardrail = _CallerSignInGuardrail( + guardrail_name="sign-in-stub", + preflight_result=Unavailable("the Entra token endpoint could not be reached", fail_open=False), + ) + initialize = json.dumps({"jsonrpc": "2.0", "id": 1, "method": "initialize", "params": {}}).encode() + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/catalog", + "headers": [ + (b"content-type", b"application/json"), + (b"authorization", b"Bearer entra.jwt.token"), + (b"mcp-session-id", b"session-mid-termination"), + ], + } + incoming: asyncio.Queue[Message] = asyncio.Queue() + await incoming.put({"type": "http.request", "body": initialize, "more_body": False}) + sent_messages: list[Message] = [] + + async def capture_send(message: Message) -> None: + sent_messages.append(message) + + handle_request_mock = AsyncMock() + litellm.logging_callback_manager.add_litellm_callback(guardrail) + try: + with ( + patch( + "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", + new_callable=AsyncMock, + return_value=( + UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"), + None, + ["catalog"], + None, + None, + {"x-litellm-api-key": "sk-litellm-virtual-key", "authorization": "Bearer entra.jwt.token"}, + ), + ), + patch.object( + mcp_operations.global_mcp_server_manager, + "get_filtered_registry", + return_value={server.server_id: server}, + ), + patch.object(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server])), + patch.object(mcp_server, "_check_passthrough_upstream_auth", AsyncMock()), + patch.object(mcp_operations, "_raise_if_initialize_grants_no_mcp_servers", AsyncMock()), + patch.object(mcp_server, "_SESSION_MANAGERS_INITIALIZED", True), + patch.object(session_manager_stateful, "handle_request", side_effect=handle_request_mock), + patch.object(session_manager_stateful, "_server_instances", {}), + patch.object(mcp_server, "_owner_fingerprint_for", return_value="owner-fingerprint"), + patch.dict( + mcp_server._stateful_session_owners, {"session-mid-termination": "owner-fingerprint"}, clear=True + ), + pytest.raises(HTTPException) as exc, + ): + await asyncio.wait_for(handle_streamable_http_mcp(scope, incoming.get, capture_send), timeout=2) + finally: + litellm.logging_callback_manager.remove_callback_from_list_by_object( + litellm.callbacks, guardrail, require_self=False + ) + + assert exc.value.status_code == 503 + assert guardrail.preflight_calls == ["entra.jwt.token"] + handle_request_mock.assert_not_awaited() + assert sent_messages == [] + + @pytest.mark.asyncio async def test_stateful_mcp_session_post_hands_the_whole_body_to_the_session_manager(): """The owner's follow-up POST on a live session reaches the stateful manager with every body byte intact @@ -9208,7 +9365,7 @@ class TestConnectPreflightRoutesLikeTheScopedRouter: user_api_key_auth=UserAPIKeyAuth(api_key="sk-granted-docs", user_id="u-1"), client_ip=None, raw_headers={"x-litellm-api-key": "sk-granted-docs"}, - connecting=True, + connecting=_connecting, ) @pytest.mark.asyncio @@ -9297,7 +9454,7 @@ class TestConnectPreflightRoutesLikeTheScopedRouter: user_api_key_auth=UserAPIKeyAuth(api_key="sk-collision", user_id="u-1"), client_ip=None, raw_headers={"x-litellm-api-key": "sk-collision"}, - connecting=True, + connecting=_connecting, ) @pytest.mark.asyncio @@ -9405,7 +9562,7 @@ class TestConnectPreflightRoutesLikeTheScopedRouter: user_api_key_auth=UserAPIKeyAuth(api_key="sk-group", user_id="u-1"), client_ip=None, raw_headers={"x-litellm-api-key": "sk-group"}, - connecting=True, + connecting=_connecting, ) assert outcome is None, ( @@ -10728,7 +10885,7 @@ class TestPreemptive401ModeAware: mcp_server_auth_headers=None, user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key"), client_ip=None, - connecting=True, + connecting=_connecting, ) async def _connect_with_a_grant(self, server, requested: str, path: str) -> HTTPException: @@ -10751,7 +10908,7 @@ class TestPreemptive401ModeAware: mcp_server_auth_headers=None, user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key"), client_ip=None, - connecting=True, + connecting=_connecting, ) return exc.value @@ -10833,7 +10990,7 @@ class TestPreemptive401ModeAware: mcp_server_auth_headers=None, user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key"), client_ip=None, - connecting=True, + connecting=_connecting, ) assert outcome is None, "an unselected aggregate connect must pass without a sign-in challenge" @@ -10936,7 +11093,7 @@ class TestPreemptive401ModeAware: mcp_server_auth_headers=None, user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key"), client_ip=None, - connecting=True, + connecting=_connecting, ) assert exc.value.status_code == 401 @@ -11045,7 +11202,7 @@ class TestSingleServerPreflightReachesIdJag: mcp_server_auth_headers=None, user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"), client_ip=None, - connecting=True, + connecting=_connecting, ) @pytest.mark.asyncio @@ -11104,7 +11261,7 @@ class TestSingleServerPreflightReachesIdJag: mcp_server_auth_headers=None, user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"), client_ip=None, - connecting=True, + connecting=_connecting, ) assert exc.value.status_code == 401 @@ -11173,7 +11330,7 @@ class TestOboPreflightScopedToAllowedServers: "x-litellm-api-key": user_api_key_auth.api_key if user_api_key_auth else "", "authorization": self.SUBJECT_HEADERS["Authorization"], }, - connecting=True, + connecting=_connecting, ) return allowed_lookup, preflight @@ -11236,7 +11393,7 @@ class TestOboChallengeGateKeepsBaseConnectRules: user_api_key_auth=UserAPIKeyAuth(api_key="sk-1234", user_id="u-1"), client_ip=None, raw_headers={"x-litellm-api-key": "sk-1234", "authorization": "Bearer sk-1234"}, - connecting=True, + connecting=_connecting, ) @pytest.mark.asyncio @@ -12081,7 +12238,7 @@ class TestConnectChallengeResolver: mcp_server_auth_headers=None, user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"), client_ip=None, - connecting=True, + connecting=_connecting, ) finally: litellm.logging_callback_manager.remove_callback_from_list_by_object( @@ -12123,7 +12280,7 @@ class TestConnectChallengeResolver: mcp_server_auth_headers=None, user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"), client_ip=None, - connecting=True, + connecting=_connecting, ) finally: litellm.logging_callback_manager.remove_callback_from_list_by_object( @@ -12176,7 +12333,7 @@ class TestConnectChallengeResolver: mcp_server_auth_headers=None, user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"), client_ip=None, - connecting=True, + connecting=_connecting, ) assert exc.value.status_code == 401 @@ -12244,7 +12401,7 @@ class TestConnectChallengeResolver: mcp_server_auth_headers=None, user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"), client_ip=client_ip, - connecting=True, + connecting=_connecting, ) finally: litellm.logging_callback_manager.remove_callback_from_list_by_object( @@ -12293,7 +12450,7 @@ class TestConnectChallengeResolver: mcp_server_auth_headers=None, user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"), client_ip=None, - connecting=True, + connecting=_connecting, ) finally: litellm.logging_callback_manager.remove_callback_from_list_by_object( @@ -12310,7 +12467,7 @@ class TestConnectSignInPreflight: """A subject token the provider rejects must fail at connect with the RFC 9728 challenge, not as a JSON-RPC error on every tools/call where ``WWW-Authenticate`` is lost.""" - async def _connect(self, route_names, guardrail, allowed, raw_headers=None, connecting=True): + async def _connect(self, route_names, guardrail, allowed, raw_headers=None, connecting=_connecting): from litellm.proxy._experimental.mcp_server import server as server_module server = _catalog_server()