fix(mcp): keep the fail-closed connect gate for an initialize naming a stale session

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-10-03 07:04:46 +00:00
parent bcfb081279
commit 393853834f
2 changed files with 85 additions and 3 deletions

View file

@ -2072,10 +2072,13 @@ if MCP_AVAILABLE:
user_api_key_auth = await _apply_toolset_scope(user_api_key_auth, active_toolset_id)
toolset_allowed_server_ids = await _toolset_server_ids(active_toolset_id)
session_header_present: Final = _get_session_id_from_scope(scope) is not None
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()
)
consumed_messages, connect_body = (
await _read_request_body_for_routing(receive)
if scope.get("method") == "POST" and not session_header_present
if scope.get("method") == "POST" and not names_live_session
else ([], b"")
)
connecting: Final = _is_initialize_request(connect_body)
@ -2176,7 +2179,7 @@ if MCP_AVAILABLE:
session_messages, session_body = (
await _read_request_body_for_routing(receive)
if scope.get("method") == "POST" and session_header_present
if scope.get("method") == "POST" and names_live_session
else ([], b"")
)
consumed_messages.extend(session_messages)

View file

@ -3889,6 +3889,85 @@ async def test_stateful_mcp_session_owner_mismatch_is_refused_before_the_body_ar
assert statuses == [403]
@pytest.mark.asyncio
async def test_initialize_naming_a_stale_session_still_meets_the_fail_closed_connect_gate():
"""A client that retries ``initialize`` with a session id this worker no longer knows is connecting, so a
fail-closed sign-in outage answers it 503 at the gate instead of letting the stripped-header retry open a
session."""
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"stale-session-from-a-restarted-worker"),
],
}
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", {}),
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