fix(mcp): enforce session ownership before reading a session-bearing POST body

A POST that names an existing mcp-session-id skips the connect-time body peek, so another caller's request is refused with 403 before any body byte is awaited. Session-bearing bodies are read after the owner check and replayed to the session manager unchanged

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-10-03 06:27:39 +00:00
parent 7603412660
commit 10c9fd95b2
2 changed files with 143 additions and 10 deletions

View file

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

View file

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