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
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)

View file

@ -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,
)

View file

@ -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()