fix(mcp): answer the connect challenge before reading a session-less POST body

The connect gate only needs the body to tell initialize from other methods, so the peek now runs on demand through a callback and the consumed ASGI messages replay to the handler. A gated route answers 401 before the body arrives again, as the merge base did

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-10-03 09:23:03 +00:00
parent 393853834f
commit d7a2d28672
3 changed files with 227 additions and 50 deletions

View file

@ -15,7 +15,7 @@ never imports a concrete provider.
from __future__ import annotations from __future__ import annotations
import itertools import itertools
from collections.abc import Mapping from collections.abc import Awaitable, Callable, Mapping
from dataclasses import dataclass from dataclasses import dataclass
from typing import TYPE_CHECKING, Final, Protocol, runtime_checkable from typing import TYPE_CHECKING, Final, Protocol, runtime_checkable
@ -162,7 +162,7 @@ async def preflight_caller_sign_in(
*, *,
root_path: str, root_path: str,
resource_metadata: str | None, resource_metadata: str | None,
connecting: bool, connecting: Callable[[], Awaitable[bool]],
) -> None: ) -> None:
"""Run every provider's connect-time check against the subject token, so a bearer the IdP will """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 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): case Unavailable(fail_open=True):
continue continue
case Unavailable(detail=detail, fail_open=False) if connecting: case Unavailable(detail=detail, fail_open=False):
raise HTTPException(status_code=503, detail=detail) if await connecting():
case Unavailable(): raise HTTPException(status_code=503, detail=detail)
continue
case _ as verdict: case _ as verdict:
assert_never(verdict) assert_never(verdict)

View file

@ -8,12 +8,13 @@ import asyncio
import contextlib import contextlib
import contextvars import contextvars
import hashlib import hashlib
import itertools
import json import json
import os import os
import time import time
import types import types
from collections import Counter 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 from typing import TYPE_CHECKING, Final, NoReturn, Protocol
import httpx import httpx
@ -1426,6 +1427,41 @@ if MCP_AVAILABLE:
return consumed_messages, b"".join(body_chunks) 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( async def _handle_stale_mcp_session(
scope: Scope, scope: Scope,
receive: Receive, receive: Receive,
@ -1608,7 +1644,7 @@ if MCP_AVAILABLE:
mcp_server_auth_headers: dict[str, dict[str, str]] | None, mcp_server_auth_headers: dict[str, dict[str, str]] | None,
user_api_key_auth: UserAPIKeyAuth | None, user_api_key_auth: UserAPIKeyAuth | None,
client_ip: str | None, client_ip: str | None,
connecting: bool, connecting: Callable[[], Awaitable[bool]],
allowed_server_ids: set[str] | None = None, allowed_server_ids: set[str] | None = None,
raw_headers: Mapping[str, str] | None = None, raw_headers: Mapping[str, str] | None = None,
) -> None: ) -> None:
@ -2073,25 +2109,13 @@ if MCP_AVAILABLE:
toolset_allowed_server_ids = await _toolset_server_ids(active_toolset_id) toolset_allowed_server_ids = await _toolset_server_ids(active_toolset_id)
named_session_id: Final = _get_session_id_from_scope(scope) named_session_id: Final = _get_session_id_from_scope(scope)
names_live_session: Final = named_session_id is not None and ( names_live_session: Final = (
named_session_id in _stateful_session_owners or named_session_id in _stateful_server_instances() named_session_id is not None and named_session_id in _stateful_server_instances()
) )
consumed_messages, connect_body = ( connect_peek: Final = _ConnectBodyPeek(
await _read_request_body_for_routing(receive) receive, peekable=scope.get("method") == "POST" and not names_live_session
if scope.get("method") == "POST" and not names_live_session
else ([], b"")
) )
connecting: Final = _is_initialize_request(connect_body) receive = connect_peek.receive
# 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
# https://datatracker.ietf.org/doc/html/rfc9728#name-www-authenticate-response # https://datatracker.ietf.org/doc/html/rfc9728#name-www-authenticate-response
# Must run after toolset scoping so the challenge set is derived # 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, mcp_server_auth_headers=mcp_server_auth_headers,
user_api_key_auth=user_api_key_auth, user_api_key_auth=user_api_key_auth,
client_ip=_client_ip, client_ip=_client_ip,
connecting=connecting, connecting=connect_peek.connecting,
allowed_server_ids=toolset_allowed_server_ids, allowed_server_ids=toolset_allowed_server_ids,
raw_headers=raw_headers, raw_headers=raw_headers,
) )
@ -2177,13 +2201,10 @@ if MCP_AVAILABLE:
return return
session_id = _get_session_id_from_scope(scope) session_id = _get_session_id_from_scope(scope)
session_messages, session_body = ( session_body: Final = (
await _read_request_body_for_routing(receive) await connect_peek.read() if scope.get("method") == "POST" and names_live_session else b""
if scope.get("method") == "POST" and names_live_session
else ([], b"")
) )
consumed_messages.extend(session_messages) body: Final = await connect_peek.body() or session_body
body: Final = connect_body or session_body
is_initialize: Final = _is_initialize_request(body) is_initialize: Final = _is_initialize_request(body)
use_stateful: Final = bool(session_id or is_initialize) use_stateful: Final = bool(session_id or is_initialize)
@ -2443,7 +2464,7 @@ if MCP_AVAILABLE:
mcp_server_auth_headers=mcp_server_auth_headers, mcp_server_auth_headers=mcp_server_auth_headers,
user_api_key_auth=user_api_key_auth, user_api_key_auth=user_api_key_auth,
client_ip=_sse_client_ip, client_ip=_sse_client_ip,
connecting=scope["method"] == "GET", connecting=_known_connecting(scope["method"] == "GET"),
allowed_server_ids=toolset_allowed_server_ids, allowed_server_ids=toolset_allowed_server_ids,
raw_headers=raw_headers, raw_headers=raw_headers,
) )

View file

@ -53,6 +53,10 @@ def test_mcp_available_on_sdk2():
assert MCP_AVAILABLE is True assert MCP_AVAILABLE is True
async def _connecting() -> bool:
return True
def _rendered_log_message(call): def _rendered_log_message(call):
message = str(call.args[0]) message = str(call.args[0])
values = call.args[1:] 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 == [] 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 @pytest.mark.asyncio
async def test_stateful_mcp_session_post_hands_the_whole_body_to_the_session_manager(): 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 """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"), user_api_key_auth=UserAPIKeyAuth(api_key="sk-granted-docs", user_id="u-1"),
client_ip=None, client_ip=None,
raw_headers={"x-litellm-api-key": "sk-granted-docs"}, raw_headers={"x-litellm-api-key": "sk-granted-docs"},
connecting=True, connecting=_connecting,
) )
@pytest.mark.asyncio @pytest.mark.asyncio
@ -9297,7 +9454,7 @@ class TestConnectPreflightRoutesLikeTheScopedRouter:
user_api_key_auth=UserAPIKeyAuth(api_key="sk-collision", user_id="u-1"), user_api_key_auth=UserAPIKeyAuth(api_key="sk-collision", user_id="u-1"),
client_ip=None, client_ip=None,
raw_headers={"x-litellm-api-key": "sk-collision"}, raw_headers={"x-litellm-api-key": "sk-collision"},
connecting=True, connecting=_connecting,
) )
@pytest.mark.asyncio @pytest.mark.asyncio
@ -9405,7 +9562,7 @@ class TestConnectPreflightRoutesLikeTheScopedRouter:
user_api_key_auth=UserAPIKeyAuth(api_key="sk-group", user_id="u-1"), user_api_key_auth=UserAPIKeyAuth(api_key="sk-group", user_id="u-1"),
client_ip=None, client_ip=None,
raw_headers={"x-litellm-api-key": "sk-group"}, raw_headers={"x-litellm-api-key": "sk-group"},
connecting=True, connecting=_connecting,
) )
assert outcome is None, ( assert outcome is None, (
@ -10728,7 +10885,7 @@ class TestPreemptive401ModeAware:
mcp_server_auth_headers=None, mcp_server_auth_headers=None,
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key"), user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key"),
client_ip=None, client_ip=None,
connecting=True, connecting=_connecting,
) )
async def _connect_with_a_grant(self, server, requested: str, path: str) -> HTTPException: async def _connect_with_a_grant(self, server, requested: str, path: str) -> HTTPException:
@ -10751,7 +10908,7 @@ class TestPreemptive401ModeAware:
mcp_server_auth_headers=None, mcp_server_auth_headers=None,
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key"), user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key"),
client_ip=None, client_ip=None,
connecting=True, connecting=_connecting,
) )
return exc.value return exc.value
@ -10833,7 +10990,7 @@ class TestPreemptive401ModeAware:
mcp_server_auth_headers=None, mcp_server_auth_headers=None,
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key"), user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key"),
client_ip=None, client_ip=None,
connecting=True, connecting=_connecting,
) )
assert outcome is None, "an unselected aggregate connect must pass without a sign-in challenge" 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, mcp_server_auth_headers=None,
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key"), user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key"),
client_ip=None, client_ip=None,
connecting=True, connecting=_connecting,
) )
assert exc.value.status_code == 401 assert exc.value.status_code == 401
@ -11045,7 +11202,7 @@ class TestSingleServerPreflightReachesIdJag:
mcp_server_auth_headers=None, mcp_server_auth_headers=None,
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"), user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"),
client_ip=None, client_ip=None,
connecting=True, connecting=_connecting,
) )
@pytest.mark.asyncio @pytest.mark.asyncio
@ -11104,7 +11261,7 @@ class TestSingleServerPreflightReachesIdJag:
mcp_server_auth_headers=None, mcp_server_auth_headers=None,
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"), user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"),
client_ip=None, client_ip=None,
connecting=True, connecting=_connecting,
) )
assert exc.value.status_code == 401 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 "", "x-litellm-api-key": user_api_key_auth.api_key if user_api_key_auth else "",
"authorization": self.SUBJECT_HEADERS["Authorization"], "authorization": self.SUBJECT_HEADERS["Authorization"],
}, },
connecting=True, connecting=_connecting,
) )
return allowed_lookup, preflight return allowed_lookup, preflight
@ -11236,7 +11393,7 @@ class TestOboChallengeGateKeepsBaseConnectRules:
user_api_key_auth=UserAPIKeyAuth(api_key="sk-1234", user_id="u-1"), user_api_key_auth=UserAPIKeyAuth(api_key="sk-1234", user_id="u-1"),
client_ip=None, client_ip=None,
raw_headers={"x-litellm-api-key": "sk-1234", "authorization": "Bearer sk-1234"}, raw_headers={"x-litellm-api-key": "sk-1234", "authorization": "Bearer sk-1234"},
connecting=True, connecting=_connecting,
) )
@pytest.mark.asyncio @pytest.mark.asyncio
@ -12081,7 +12238,7 @@ class TestConnectChallengeResolver:
mcp_server_auth_headers=None, mcp_server_auth_headers=None,
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"), user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"),
client_ip=None, client_ip=None,
connecting=True, connecting=_connecting,
) )
finally: finally:
litellm.logging_callback_manager.remove_callback_from_list_by_object( litellm.logging_callback_manager.remove_callback_from_list_by_object(
@ -12123,7 +12280,7 @@ class TestConnectChallengeResolver:
mcp_server_auth_headers=None, mcp_server_auth_headers=None,
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"), user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"),
client_ip=None, client_ip=None,
connecting=True, connecting=_connecting,
) )
finally: finally:
litellm.logging_callback_manager.remove_callback_from_list_by_object( litellm.logging_callback_manager.remove_callback_from_list_by_object(
@ -12176,7 +12333,7 @@ class TestConnectChallengeResolver:
mcp_server_auth_headers=None, mcp_server_auth_headers=None,
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"), user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"),
client_ip=None, client_ip=None,
connecting=True, connecting=_connecting,
) )
assert exc.value.status_code == 401 assert exc.value.status_code == 401
@ -12244,7 +12401,7 @@ class TestConnectChallengeResolver:
mcp_server_auth_headers=None, mcp_server_auth_headers=None,
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"), user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"),
client_ip=client_ip, client_ip=client_ip,
connecting=True, connecting=_connecting,
) )
finally: finally:
litellm.logging_callback_manager.remove_callback_from_list_by_object( litellm.logging_callback_manager.remove_callback_from_list_by_object(
@ -12293,7 +12450,7 @@ class TestConnectChallengeResolver:
mcp_server_auth_headers=None, mcp_server_auth_headers=None,
user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"), user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-virtual-key", user_id="u-1"),
client_ip=None, client_ip=None,
connecting=True, connecting=_connecting,
) )
finally: finally:
litellm.logging_callback_manager.remove_callback_from_list_by_object( 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 """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.""" 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 from litellm.proxy._experimental.mcp_server import server as server_module
server = _catalog_server() server = _catalog_server()