From 26715ba66766b4e98b56c19837538bcde46aa247 Mon Sep 17 00:00:00 2001 From: yucheng Date: Sat, 3 Oct 2026 09:50:50 +0000 Subject: [PATCH] fix(mcp): refuse another caller's torn-down session before peeking at its body under a fail-closed outage Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../proxy/_experimental/mcp_server/server.py | 12 +-- .../test_mcp_server_tool_calls_and_headers.py | 87 +++++++++++++++++++ 2 files changed, 94 insertions(+), 5 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 4d9bdf41696..e3fc403302a 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -2112,8 +2112,13 @@ if MCP_AVAILABLE: names_live_session: Final = ( named_session_id is not None and named_session_id in _stateful_server_instances() ) + request_owner: Final = _owner_fingerprint_for(user_api_key_auth, oauth2_headers, _client_ip) + expected_owner: Final = ( + _stateful_session_owners.get(named_session_id) if named_session_id is not None else None + ) + owner_mismatch: Final = expected_owner is not None and expected_owner != request_owner connect_peek: Final = _ConnectBodyPeek( - receive, peekable=scope.get("method") == "POST" and not names_live_session + receive, peekable=scope.get("method") == "POST" and not names_live_session and not owner_mismatch ) receive = connect_peek.receive @@ -2174,9 +2179,7 @@ if MCP_AVAILABLE: # force-clean another caller's residual tracking entries via a # stale DELETE. if session_id: - expected_owner: Final = _stateful_session_owners.get(session_id) - request_owner = _owner_fingerprint_for(user_api_key_auth, oauth2_headers, _client_ip) - if expected_owner is not None and expected_owner != request_owner: + if owner_mismatch: verbose_logger.warning( "Rejecting MCP request: session '%s' owner mismatch.", session_id, @@ -2220,7 +2223,6 @@ if MCP_AVAILABLE: # session. Cap how many a single caller can hold so an authenticated # client cannot spam `initialize` and exhaust memory. if is_initialize and not session_id: - request_owner = _owner_fingerprint_for(user_api_key_auth, oauth2_headers, _client_ip) if not await _enforce_stateful_session_cap_for_owner(request_owner): verbose_logger.warning( "Rejecting MCP initialize: caller already holds the maximum number of active stateful sessions." 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 85b631251ef..2cc066c6f1b 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 @@ -3893,6 +3893,93 @@ async def test_stateful_mcp_session_owner_mismatch_is_refused_before_the_body_ar assert statuses == [403] +@pytest.mark.asyncio +async def test_owner_mismatch_on_a_torn_down_session_is_refused_before_the_body_under_a_fail_closed_outage(): + """While a session's transport is already gone but its owner binding is still recorded, a POST from another + caller is refused with 403 before its withheld body arrives even when the fail-closed sign-in gate would + otherwise peek at the body to tell an ``initialize`` apart.""" + 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), + ) + session_id = "owned-session-being-torn-down" + owner_auth = UserAPIKeyAuth(api_key="owner-key", user_id="owner") + intruder_auth = UserAPIKeyAuth(api_key="intruder-key", user_id="intruder") + mcp_server._stateful_session_owners[session_id] = mcp_server._owner_fingerprint_for(owner_auth) + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/catalog", + "headers": [ + (b"content-type", b"application/json"), + (b"x-litellm-api-key", b"intruder-key"), + (b"authorization", b"Bearer entra.jwt.token"), + (b"mcp-session-id", session_id.encode()), + ], + } + body_never_sent = asyncio.Event() + + async def withheld_body() -> Message: + await body_never_sent.wait() + return {"type": "http.request", "body": b"", "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=( + intruder_auth, + None, + ["catalog"], + None, + None, + {"x-litellm-api-key": "intruder-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", {}), + ): + await asyncio.wait_for(handle_streamable_http_mcp(scope, withheld_body, capture_send), timeout=2) + finally: + mcp_server._stateful_session_owners.pop(session_id, None) + litellm.logging_callback_manager.remove_callback_from_list_by_object( + litellm.callbacks, guardrail, require_self=False + ) + + assert guardrail.preflight_calls == ["entra.jwt.token"] + handle_request_mock.assert_not_awaited() + statuses = [m["status"] for m in sent_messages if m.get("type") == "http.response.start"] + 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