From cf986ff3436943030ee382005da1510b17a36d50 Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Tue, 5 May 2026 18:37:53 +0000 Subject: [PATCH] Fix MCP stateful routing edge cases --- .../proxy/_experimental/mcp_server/server.py | 67 +++++++++----- .../mcp_server/test_mcp_server.py | 87 ++++++++++++++++++- .../mcp_server/test_mcp_stale_session.py | 28 ++++-- 3 files changed, 153 insertions(+), 29 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 4cd49267a54..fbb5aea7478 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -28,7 +28,7 @@ from fastapi import FastAPI, HTTPException from pydantic import AnyUrl, ConfigDict from starlette.requests import Request as StarletteRequest from starlette.responses import JSONResponse -from starlette.types import Receive, Scope, Send +from starlette.types import Message, Receive, Scope, Send from litellm._logging import verbose_logger from litellm.constants import MAXIMUM_TRACEBACK_LINES_TO_LOG @@ -2617,9 +2617,6 @@ if MCP_AVAILABLE: """ Check if the request body is a JSON-RPC initialize method. Returns True if method is "initialize", False otherwise or on parse error. - - Note: Assumes the body fits in a single ASGI receive() chunk. MCP initialize - requests are typically small; large chunked bodies may not be fully parsed. """ if not body: return False @@ -2629,6 +2626,31 @@ if MCP_AVAILABLE: except (json.JSONDecodeError, TypeError): return False + async def _read_request_body_for_routing( + receive: Receive, + ) -> Tuple[List[Message], bytes]: + """ + Consume request body messages for routing, returning them for replay. + """ + consumed_messages: List[Message] = [] + body_chunks: List[bytes] = [] + + while True: + message = await receive() + consumed_messages.append(message) + + if message.get("type") != "http.request": + break + + body = message.get("body", b"") or b"" + if body: + body_chunks.append(body) + + if not message.get("more_body", False): + break + + return consumed_messages, b"".join(body_chunks) + async def _handle_stale_mcp_session( scope: Scope, receive: Receive, @@ -2883,12 +2905,22 @@ if MCP_AVAILABLE: # - No session ID + other → stateless (curl, Inspector, Notion) session_id = _get_session_id_from_scope(scope) is_initialize = False - first_msg = None + consumed_messages: List[Message] = [] + + # Handle stale session IDs before choosing a target manager. Stale + # non-DELETE requests have their session header stripped and should + # be routed as no-session requests. + if session_id: + handled = await _handle_stale_mcp_session( + scope, receive, send, session_manager_stateful + ) + if handled: + # Request was fully handled (e.g., DELETE on non-existent session) + return + session_id = _get_session_id_from_scope(scope) if scope.get("method") == "POST" and not session_id: - # Peek at first chunk to detect initialize (assumes body fits in one chunk) - first_msg = await receive() - body = first_msg.get("body", b"") or b"" + consumed_messages, body = await _read_request_body_for_routing(receive) is_initialize = _is_initialize_request(body) use_stateful = bool(session_id or is_initialize) @@ -2902,28 +2934,17 @@ if MCP_AVAILABLE: + (" (initialize)" if is_initialize else "") ) - # Replay first message if we consumed it for peeking + # Replay body messages if we consumed them for peeking original_receive = receive - if first_msg is not None: + if consumed_messages: async def wrapped_receive(): - nonlocal first_msg - if first_msg is not None: - msg, first_msg = first_msg, None - return msg + if consumed_messages: + return consumed_messages.pop(0) return await original_receive() receive = wrapped_receive - # Handle stale session IDs - either strip them for reconnection - # or return success for idempotent DELETE operations - handled = await _handle_stale_mcp_session( - scope, receive, send, target_manager - ) - if handled: - # Request was fully handled (e.g., DELETE on non-existent session) - return - async with _gateway_initialize_instructions_request_scope( user_api_key_auth, mcp_servers, diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 672d37a3238..96cdbffad75 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -1135,7 +1135,13 @@ async def test_mcp_routing_initialize_to_stateful_no_session_to_stateless(): (b"authorization", b"Bearer test-key"), ], } - receive = AsyncMock(return_value={"type": "http.request", "body": method_body, "more_body": False}) + receive = AsyncMock( + return_value={ + "type": "http.request", + "body": method_body, + "more_body": False, + } + ) send = AsyncMock() stateless_called = [] @@ -1192,6 +1198,85 @@ async def test_mcp_routing_initialize_to_stateful_no_session_to_stateless(): ) +@pytest.mark.asyncio +async def test_mcp_routing_chunked_initialize_to_stateful(): + """ + Test that chunked initialize requests route to the stateful manager. + """ + try: + from litellm.proxy._experimental.mcp_server.server import ( + handle_streamable_http_mcp, + session_manager_stateful, + session_manager_stateless, + ) + except ImportError: + pytest.skip("MCP server not available") + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/progress_test", + "headers": [ + (b"content-type", b"application/json"), + (b"authorization", b"Bearer test-key"), + ], + } + messages = [ + { + "type": "http.request", + "body": b'{"jsonrpc":"2.0","id":1,', + "more_body": True, + }, + { + "type": "http.request", + "body": b'"method":"initialize","params":{}}', + "more_body": False, + }, + ] + receive = AsyncMock(side_effect=messages) + send = AsyncMock() + stateless_called = [] + stateful_called = [] + + async def stateless_handle(s, r, se): + stateless_called.append(1) + + async def stateful_handle(s, r, se): + stateful_called.append(1) + + with patch( + "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", + new_callable=AsyncMock, + return_value=(MagicMock(), None, ["progress_test"], None, None, None), + ), patch( + "litellm.proxy._experimental.mcp_server.server.set_auth_context", + ), patch( + "litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", + True, + ), patch.object( + session_manager_stateless, + "handle_request", + side_effect=stateless_handle, + ), patch.object( + session_manager_stateful, + "handle_request", + side_effect=stateful_handle, + ), patch.object( + session_manager_stateless, + "_server_instances", + {}, + ), patch.object( + session_manager_stateful, + "_server_instances", + {}, + ): + await handle_streamable_http_mcp(scope, receive, send) + + assert stateful_called and not stateless_called, ( + "chunked initialize (no session) should route to stateful, not stateless" + ) + + @pytest.mark.asyncio @pytest.mark.no_parallel async def test_mcp_routing_with_conflicting_alias_and_group_name(): diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py index f4a9b118095..8cba3dd2842 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py @@ -202,7 +202,8 @@ async def test_stale_mcp_session_id_is_stripped(): try: from litellm.proxy._experimental.mcp_server.server import ( handle_streamable_http_mcp, - session_manager, + session_manager_stateful, + session_manager_stateless, ) except ImportError: pytest.skip("MCP server not available") @@ -225,8 +226,9 @@ async def test_stale_mcp_session_id_is_stripped(): # Simulate: session manager has NO sessions (the stale one was cleaned up) captured_scope = {} + stateful_handle_request = AsyncMock() - async def mock_handle_request(s, r, se): + async def stateless_handle_request(s, r, se): # Capture the scope that was actually passed captured_scope.update(s) @@ -244,12 +246,22 @@ async def test_stale_mcp_session_id_is_stripped(): True, ), patch.object( - session_manager, + session_manager_stateless, "handle_request", - side_effect=mock_handle_request, + side_effect=stateless_handle_request, ), patch.object( - session_manager, + session_manager_stateless, + "_server_instances", + {}, + ), + patch.object( + session_manager_stateful, + "handle_request", + side_effect=stateful_handle_request, + ), + patch.object( + session_manager_stateful, "_server_instances", {}, # Empty dict = no active sessions ), @@ -261,6 +273,12 @@ async def test_stale_mcp_session_id_is_stripped(): assert ( b"mcp-session-id" not in header_names ), "Stale mcp-session-id header should have been stripped from the scope" + assert ( + stateless_handle_request.called + ), "Stale non-initialize requests should route stateless" + assert ( + not stateful_handle_request.called + ), "Stale non-initialize requests should not route stateful" @pytest.mark.asyncio