From 393853834fed8a23a02968e04c65efa11f26890c Mon Sep 17 00:00:00 2001 From: yucheng Date: Sat, 3 Oct 2026 07:04:46 +0000 Subject: [PATCH] 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> --- .../proxy/_experimental/mcp_server/server.py | 9 ++- .../test_mcp_server_tool_calls_and_headers.py | 79 +++++++++++++++++++ 2 files changed, 85 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 9ec3eed32b5..47d41325e74 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -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) 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 b58206701da..bfd759bc2a1 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 @@ -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