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>
This commit is contained in:
yucheng 2026-10-03 09:50:50 +00:00
parent d7a2d28672
commit 26715ba667
2 changed files with 94 additions and 5 deletions

View file

@ -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."

View file

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