fix(mcp): reject replay of explicitly terminated sessions

This commit is contained in:
Joshua Valluru 2026-09-26 12:27:09 -07:00
parent 2074b360a1
commit e62acdef76
4 changed files with 58 additions and 29 deletions

View file

@ -624,7 +624,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: ...
@ -687,7 +687,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:
"""
@ -1313,22 +1313,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(
@ -1344,7 +1344,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)
@ -1354,7 +1354,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:
@ -1490,12 +1490,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={ # mutable-ok: JSONResponse content must be a plain dict
"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)
@ -2245,7 +2245,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:

View file

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

View file

@ -367,13 +367,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,
@ -386,10 +385,11 @@ async def test_failed_delete_preserves_stateful_session_tracking():
pytest.skip("MCP server not available")
session_id = "delete-failure-session"
user_auth = MagicMock()
user_auth.api_key = "sk-test"
user_auth.user_id = "test-user"
auth_context = MagicMock()
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy._experimental.mcp_server.auth.litellm_auth_handler import MCPAuthenticatedUser
user_auth = UserAPIKeyAuth(api_key="sk-test", user_id="test-user")
auth_context = MCPAuthenticatedUser(user_api_key_auth=user_auth)
session_lock = asyncio.Lock()
mock_instances = {session_id: MagicMock()}
@ -411,6 +411,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(
@ -426,7 +434,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,
@ -435,6 +443,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
@ -443,6 +465,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)

View file

@ -7762,7 +7762,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])