diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 42ef852c932..9ec3eed32b5 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -2072,21 +2072,23 @@ 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) - consumed_messages, body = ( - await _read_request_body_for_routing(receive) if scope.get("method") == "POST" else ([], b"") + session_header_present: Final = _get_session_id_from_scope(scope) is not None + consumed_messages, connect_body = ( + await _read_request_body_for_routing(receive) + if scope.get("method") == "POST" and not session_header_present + else ([], b"") ) - is_initialize: Final = _is_initialize_request(body) + connecting: Final = _is_initialize_request(connect_body) # Replay body messages if we consumed them for peeking original_receive: Final = receive - if consumed_messages: - async def wrapped_receive(): - if consumed_messages: - return consumed_messages.pop(0) - return await original_receive() + async def wrapped_receive(): + if consumed_messages: + return consumed_messages.pop(0) + return await original_receive() - receive = wrapped_receive + receive = wrapped_receive # https://datatracker.ietf.org/doc/html/rfc9728#name-www-authenticate-response # Must run after toolset scoping so the challenge set is derived @@ -2100,7 +2102,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=is_initialize, + connecting=connecting, allowed_server_ids=toolset_allowed_server_ids, raw_headers=raw_headers, ) @@ -2172,6 +2174,15 @@ 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 session_header_present + else ([], b"") + ) + consumed_messages.extend(session_messages) + body: Final = connect_body or session_body + is_initialize: Final = _is_initialize_request(body) + use_stateful: Final = bool(session_id or is_initialize) target_manager: Final = session_manager_stateful if use_stateful else session_manager_stateless 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 2d510e70bc6..5dc3e9898db 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 @@ -3754,6 +3754,128 @@ async def test_stateful_mcp_session_owner_mismatch_returns_403(): mcp_server._stateful_session_owners.pop(session_id, None) +@pytest.mark.asyncio +async def test_stateful_mcp_session_owner_mismatch_is_rejected_before_the_body_is_read(): + """A POST carrying another caller's mcp-session-id is refused before any body chunk is awaited, so a + slow sender cannot hold the request open past the owner check.""" + 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") + + session_id = "owned-session-slow-body" + owner_auth = UserAPIKeyAuth(api_key="owner-key", user_id="owner") + intruder_auth = UserAPIKeyAuth(api_key="intruder-key", user_id="intruder") + mcp_server._stateful_session_auth_contexts[session_id] = MagicMock() + mcp_server._stateful_session_owners[session_id] = mcp_server._owner_fingerprint_for(owner_auth) + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp", + "headers": [ + (b"content-type", b"application/json"), + (b"authorization", b"Bearer intruder-key"), + (b"mcp-session-id", session_id.encode()), + ], + } + body_never_arrives = asyncio.Event() + + async def stalled_receive(): + await body_never_arrives.wait() + return {"type": "http.request", "body": b"", "more_body": False} + + sent_messages: list = [] + + async def capture_send(message): + sent_messages.append(message) + + try: + with ( + patch( + "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", + new_callable=AsyncMock, + return_value=(intruder_auth, None, None, None, None, None), + ), + patch("litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", True), + patch.object(session_manager_stateful, "handle_request", new_callable=AsyncMock), + patch.object(session_manager_stateful, "_server_instances", {session_id: MagicMock()}), + ): + await asyncio.wait_for(handle_streamable_http_mcp(scope, stalled_receive, capture_send), timeout=2) + finally: + body_never_arrives.set() + mcp_server._stateful_session_auth_contexts.pop(session_id, None) + mcp_server._stateful_session_owners.pop(session_id, None) + + statuses = [m["status"] for m in sent_messages if m.get("type") == "http.response.start"] + assert statuses == [403], 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 + after the routing peek.""" + 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") + + session_id = "owned-session-replay" + owner_auth = UserAPIKeyAuth(api_key="owner-key", user_id="owner") + mcp_server._stateful_session_auth_contexts[session_id] = MagicMock() + mcp_server._stateful_session_owners[session_id] = mcp_server._owner_fingerprint_for(owner_auth) + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp", + "headers": [ + (b"content-type", b"application/json"), + (b"authorization", b"Bearer owner-key"), + (b"mcp-session-id", session_id.encode()), + ], + } + chunks = [ + {"type": "http.request", "body": b'{"jsonrpc":"2.0","id":7,"method":"tools/list",', "more_body": True}, + {"type": "http.request", "body": b'"params":{}}', "more_body": False}, + ] + receive = AsyncMock(side_effect=list(chunks)) + delivered: list[bytes] = [] + + async def drain_body(scope_, receive_, send_): + while True: + message = await receive_() + delivered.append(message.get("body", b"")) + if not message.get("more_body", False): + return + + try: + with ( + patch( + "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", + new_callable=AsyncMock, + return_value=(owner_auth, None, None, None, None, None), + ), + patch("litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", True), + patch.object(session_manager_stateful, "handle_request", side_effect=drain_body), + patch.object(session_manager_stateful, "_server_instances", {session_id: MagicMock()}), + ): + await handle_streamable_http_mcp(scope, receive, AsyncMock()) + finally: + mcp_server._stateful_session_auth_contexts.pop(session_id, None) + mcp_server._stateful_session_owners.pop(session_id, None) + + assert b"".join(delivered) == b'{"jsonrpc":"2.0","id":7,"method":"tools/list","params":{}}' + + @pytest.mark.asyncio async def test_stateful_mcp_session_serializes_concurrent_requests(): """