mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
Fix MCP stateful routing edge cases
This commit is contained in:
parent
6d13264cf3
commit
cf986ff343
3 changed files with 153 additions and 29 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue