Fix MCP stateful routing edge cases

This commit is contained in:
Cursor Agent 2026-05-05 18:37:53 +00:00
parent 6d13264cf3
commit cf986ff343
No known key found for this signature in database
3 changed files with 153 additions and 29 deletions

View file

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

View file

@ -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():

View file

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