mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
7603412660
commit
10c9fd95b2
2 changed files with 143 additions and 10 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue