mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
fix(mcp): reject replay of explicitly terminated sessions
This commit is contained in:
parent
05ee360a79
commit
def57abac8
4 changed files with 56 additions and 26 deletions
|
|
@ -628,7 +628,7 @@ if MCP_AVAILABLE:
|
|||
_stateful_session_locks: Final[dict[str, asyncio.Lock]] = {}
|
||||
_stateful_session_active_request_counts: Final[dict[str, int]] = {}
|
||||
_stateful_session_client_info: Final[dict[str, Implementation]] = {} # mutable-ok: cleared on session teardown
|
||||
_admin_terminated_session_ids: Final[dict[str, float]] = {} # mutable-ok: admin-closed id -> last replay
|
||||
_terminated_session_ids: Final[dict[str, float]] = {} # mutable-ok: explicitly closed id -> last replay
|
||||
|
||||
class _TerminableTransport(Protocol):
|
||||
async def terminate(self) -> None: ...
|
||||
|
|
@ -691,7 +691,7 @@ if MCP_AVAILABLE:
|
|||
for session_id in list(_stateful_session_auth_context_last_seen):
|
||||
if session_id not in _stateful_session_auth_contexts:
|
||||
_remove_stateful_session_tracking(session_id)
|
||||
_forget_expired_admin_terminated_session_ids(now)
|
||||
_forget_expired_terminated_session_ids(now)
|
||||
|
||||
async def _enforce_stateful_session_cap_for_owner(owner: str) -> bool:
|
||||
"""
|
||||
|
|
@ -1317,22 +1317,22 @@ if MCP_AVAILABLE:
|
|||
key_auth: Final = auth_user.user_api_key_auth
|
||||
return key_auth is not None and key_auth.user_id == user_id
|
||||
|
||||
def _forget_expired_admin_terminated_session_ids(now: float) -> None:
|
||||
def _forget_expired_terminated_session_ids(now: float) -> None:
|
||||
for session_id in [
|
||||
session_id
|
||||
for session_id, last_replayed in _admin_terminated_session_ids.items()
|
||||
for session_id, last_replayed in _terminated_session_ids.items()
|
||||
if now - last_replayed >= _STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS
|
||||
]:
|
||||
del _admin_terminated_session_ids[session_id]
|
||||
del _terminated_session_ids[session_id]
|
||||
|
||||
def _is_admin_terminated_session_id(session_id: str, now: float) -> bool:
|
||||
last_replayed: Final = _admin_terminated_session_ids.get(session_id)
|
||||
def _is_terminated_session_id(session_id: str, now: float) -> bool:
|
||||
last_replayed: Final = _terminated_session_ids.get(session_id)
|
||||
if last_replayed is None:
|
||||
return False
|
||||
if now - last_replayed >= _STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS:
|
||||
del _admin_terminated_session_ids[session_id]
|
||||
del _terminated_session_ids[session_id]
|
||||
return False
|
||||
_admin_terminated_session_ids[session_id] = now
|
||||
_terminated_session_ids[session_id] = now
|
||||
return True
|
||||
|
||||
async def terminate_mcp_gateway_sessions(
|
||||
|
|
@ -1348,7 +1348,7 @@ if MCP_AVAILABLE:
|
|||
admission. Only sessions held by this worker process are affected.
|
||||
"""
|
||||
now: Final = time.monotonic()
|
||||
_forget_expired_admin_terminated_session_ids(now)
|
||||
_forget_expired_terminated_session_ids(now)
|
||||
server_instances: Final = _stateful_server_instances()
|
||||
targets: Final = tuple(
|
||||
(session_id, auth_user)
|
||||
|
|
@ -1358,7 +1358,7 @@ if MCP_AVAILABLE:
|
|||
)
|
||||
terminated: Final = tuple(_gateway_session_for(session_id, auth_user, now) for session_id, auth_user in targets)
|
||||
for session_id, _ in targets:
|
||||
_admin_terminated_session_ids[session_id] = now
|
||||
_terminated_session_ids[session_id] = now
|
||||
transport = server_instances.pop(session_id, None)
|
||||
_remove_stateful_session_tracking(session_id)
|
||||
if transport is not None:
|
||||
|
|
@ -1494,12 +1494,12 @@ if MCP_AVAILABLE:
|
|||
await success_response(scope, receive, send)
|
||||
return True
|
||||
|
||||
if _is_admin_terminated_session_id(_session_id, time.monotonic()):
|
||||
if _is_terminated_session_id(_session_id, time.monotonic()):
|
||||
terminated_response: Final = JSONResponse(
|
||||
status_code=404,
|
||||
content={
|
||||
"error": "Not Found",
|
||||
"details": "mcp-session-id was terminated by an administrator. Send initialize to start a new session.",
|
||||
"details": "mcp-session-id was terminated. Send initialize to start a new session.",
|
||||
},
|
||||
)
|
||||
await terminated_response(scope, receive, send)
|
||||
|
|
@ -2263,7 +2263,13 @@ if MCP_AVAILABLE:
|
|||
is_initialize=is_initialize,
|
||||
):
|
||||
await target_manager.handle_request(scope, receive, local_send)
|
||||
if use_stateful and session_id and scope.get("method") == "DELETE":
|
||||
if (
|
||||
use_stateful
|
||||
and session_id
|
||||
and scope.get("method") == "DELETE"
|
||||
and session_id not in _stateful_server_instances()
|
||||
):
|
||||
_terminated_session_ids[session_id] = time.monotonic()
|
||||
_remove_stateful_session_tracking(session_id)
|
||||
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -3267,7 +3267,7 @@ async def test_admin_terminated_session_id_gets_404_instead_of_a_fresh_stateless
|
|||
)
|
||||
assert [k for k, _ in unknown_scope["headers"]] == [b"content-type"]
|
||||
finally:
|
||||
mcp_server._admin_terminated_session_ids.clear()
|
||||
mcp_server._terminated_session_ids.clear()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -3332,13 +3332,13 @@ async def test_admin_terminated_session_id_stays_refused_while_replayed_and_is_f
|
|||
assert await replay(retrying_id, 1000.0 + elapsed) == (True, [b"content-type", b"mcp-session-id"])
|
||||
|
||||
await mcp_server._purge_expired_stateful_session_auth_contexts(now=1000.0 + idle_timeout)
|
||||
assert set(mcp_server._admin_terminated_session_ids) == {retrying_id}
|
||||
assert set(mcp_server._terminated_session_ids) == {retrying_id}
|
||||
|
||||
assert await replay(silent_id, 1000.0 + idle_timeout) == (False, [b"content-type"])
|
||||
assert await replay(retrying_id, 1000.0 + 4 * idle_timeout) == (False, [b"content-type"])
|
||||
assert mcp_server._admin_terminated_session_ids == {}
|
||||
assert mcp_server._terminated_session_ids == {}
|
||||
finally:
|
||||
mcp_server._admin_terminated_session_ids.clear()
|
||||
mcp_server._terminated_session_ids.clear()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -370,13 +370,12 @@ async def test_delete_stale_mcp_session_returns_success():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_failed_delete_preserves_stateful_session_tracking():
|
||||
"""
|
||||
When the SDK fails to terminate an existing stateful session, keep the
|
||||
owner/auth tracking so the session cannot be hijacked or hidden from cleanup.
|
||||
"""
|
||||
@pytest.mark.parametrize("outcome", ("deleted", "rejected", "exception"))
|
||||
async def test_delete_preserves_tracking_until_the_session_is_terminated(outcome):
|
||||
try:
|
||||
from litellm.proxy._experimental.mcp_server.server import (
|
||||
_handle_stale_mcp_session,
|
||||
_terminated_session_ids,
|
||||
_owner_fingerprint_for,
|
||||
_stateful_session_auth_context_last_seen,
|
||||
_stateful_session_auth_contexts,
|
||||
|
|
@ -389,10 +388,12 @@ async def test_failed_delete_preserves_stateful_session_tracking():
|
|||
pytest.skip("MCP server not available")
|
||||
|
||||
session_id = "delete-failure-session"
|
||||
from litellm.proxy._experimental.mcp_server.auth.litellm_auth_handler import MCPAuthenticatedUser
|
||||
|
||||
user_auth = UserAPIKeyAuth()
|
||||
user_auth.api_key = "sk-test"
|
||||
user_auth.user_id = "test-user"
|
||||
auth_context = MagicMock()
|
||||
auth_context = MCPAuthenticatedUser(user_api_key_auth=user_auth)
|
||||
session_lock = asyncio.Lock()
|
||||
mock_instances = {session_id: MagicMock()}
|
||||
|
||||
|
|
@ -414,6 +415,14 @@ async def test_failed_delete_preserves_stateful_session_tracking():
|
|||
_stateful_session_owners[session_id] = _owner_fingerprint_for(user_auth)
|
||||
_stateful_session_locks[session_id] = session_lock
|
||||
|
||||
async def dispatch(scope, receive, send):
|
||||
if outcome == "exception":
|
||||
raise RuntimeError("delete failed")
|
||||
if outcome == "deleted":
|
||||
mock_instances.pop(session_id)
|
||||
await send({"type": "http.response.start", "status": 200 if outcome == "deleted" else 403, "headers": []})
|
||||
await send({"type": "http.response.body", "body": b""})
|
||||
|
||||
try:
|
||||
with (
|
||||
patch(
|
||||
|
|
@ -429,7 +438,7 @@ async def test_failed_delete_preserves_stateful_session_tracking():
|
|||
session_manager_stateful,
|
||||
"handle_request",
|
||||
new_callable=AsyncMock,
|
||||
side_effect=RuntimeError("delete failed"),
|
||||
side_effect=dispatch,
|
||||
) as mock_handle_request,
|
||||
patch.object(
|
||||
session_manager_stateful,
|
||||
|
|
@ -438,6 +447,20 @@ async def test_failed_delete_preserves_stateful_session_tracking():
|
|||
),
|
||||
):
|
||||
await handle_streamable_http_mcp(scope, receive, send)
|
||||
if outcome == "deleted":
|
||||
assert session_id not in _stateful_session_auth_contexts
|
||||
assert session_id not in _stateful_session_owners
|
||||
assert session_id not in _stateful_session_locks
|
||||
replay_send = AsyncMock()
|
||||
handled = await _handle_stale_mcp_session(
|
||||
{**scope, "method": "POST"}, receive, replay_send, session_manager_stateful
|
||||
)
|
||||
assert handled is True
|
||||
assert replay_send.await_args_list[0].args[0]["status"] == 404
|
||||
repeated_delete = AsyncMock()
|
||||
assert await _handle_stale_mcp_session(scope, receive, repeated_delete, session_manager_stateful)
|
||||
assert repeated_delete.await_args_list[0].args[0]["status"] == 200
|
||||
return
|
||||
|
||||
assert mock_handle_request.await_count == 1
|
||||
assert _stateful_session_auth_contexts[session_id] is auth_context
|
||||
|
|
@ -446,6 +469,7 @@ async def test_failed_delete_preserves_stateful_session_tracking():
|
|||
assert _stateful_session_locks[session_id] is session_lock
|
||||
assert session_id in mock_instances
|
||||
finally:
|
||||
_terminated_session_ids.pop(session_id, None)
|
||||
_stateful_session_auth_contexts.pop(session_id, None)
|
||||
_stateful_session_auth_context_last_seen.pop(session_id, None)
|
||||
_stateful_session_owners.pop(session_id, None)
|
||||
|
|
|
|||
|
|
@ -8103,7 +8103,7 @@ class TestDeleteMCPGatewaySessions:
|
|||
from litellm.proxy._experimental.mcp_server import server as mcp_server
|
||||
|
||||
yield
|
||||
mcp_server._admin_terminated_session_ids.clear()
|
||||
mcp_server._terminated_session_ids.clear()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("role", [LitellmUserRoles.INTERNAL_USER, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY])
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue