fix(mcp): bound client terminated session retention

This commit is contained in:
Joshua Valluru 2026-09-26 16:38:36 -07:00
parent e7b15b8bb6
commit 9906303fb4
2 changed files with 115 additions and 19 deletions

View file

@ -106,6 +106,7 @@ _STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS: Final = 30 * 60
# the cap is still hit (every session in flight), the new `initialize` is
# rejected with 429.
_MAX_STATEFUL_SESSIONS_PER_OWNER: Final = 100
_MAX_CLIENT_TERMINATED_SESSIONS: Final = 10000
# Maximum bytes to peek when sniffing the JSON-RPC method on a POST.
# An `initialize` envelope is a few hundred bytes; capping the peek
# prevents an authenticated client from forcing the proxy to buffer an
@ -629,6 +630,7 @@ if MCP_AVAILABLE:
_stateful_session_active_request_counts: Final[dict[str, int]] = {}
_stateful_session_client_info: Final[dict[str, Implementation]] = {} # mutable-ok: cleared on session teardown
_terminated_session_ids: Final[dict[str, float]] = {} # mutable-ok: explicitly closed id -> last replay
_client_terminated_session_owners: Final[dict[str, str]] = {}
class _TerminableTransport(Protocol):
async def terminate(self) -> None: ...
@ -1322,8 +1324,10 @@ if MCP_AVAILABLE:
session_id
for session_id, last_replayed in _terminated_session_ids.items()
if now - last_replayed >= _STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS
and not _stateful_session_active_request_counts.get(session_id)
]:
del _terminated_session_ids[session_id]
_client_terminated_session_owners.pop(session_id, None)
def _is_terminated_session_id(session_id: str, now: float) -> bool:
last_replayed: Final = _terminated_session_ids.get(session_id)
@ -1331,6 +1335,7 @@ if MCP_AVAILABLE:
return False
if now - last_replayed >= _STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS:
del _terminated_session_ids[session_id]
_client_terminated_session_owners.pop(session_id, None)
return False
_terminated_session_ids[session_id] = now
return True
@ -2231,6 +2236,10 @@ if MCP_AVAILABLE:
_increment_active_request_session(initialized_session_id)
async def _dispatch() -> None:
# Another DELETE may have completed while this request waited for the lock.
if use_stateful and session_id and scope.get("method") == "DELETE":
if await _handle_stale_mcp_session(scope, receive, send, target_manager):
return
_otel_publish_transport_span_on_scope(scope)
_otel_publish_request_destinations_on_scope(scope)
auth_user: Final = _set_or_update_auth_context(
@ -2255,22 +2264,45 @@ if MCP_AVAILABLE:
client_info=_extract_initialize_client_info(body),
)
async with _gateway_initialize_instructions_request_scope(
user_api_key_auth,
mcp_servers,
_client_ip,
scoped_server_endpoint=scoped_server_endpoint,
is_initialize=is_initialize,
):
await target_manager.handle_request(scope, receive, local_send)
deleting_session: Final = session_id if use_stateful and scope.get("method") == "DELETE" else None
if deleting_session:
_forget_expired_terminated_session_ids(time.monotonic())
delete_owner: Final = _owner_fingerprint_for(user_api_key_auth, oauth2_headers, _client_ip)
if (
use_stateful
and session_id
and scope.get("method") == "DELETE"
and session_id not in _stateful_server_instances()
len(_client_terminated_session_owners) >= _MAX_CLIENT_TERMINATED_SESSIONS
or sum(owner == delete_owner for owner in _client_terminated_session_owners.values())
>= _MAX_STATEFUL_SESSIONS_PER_OWNER
):
_terminated_session_ids[session_id] = time.monotonic()
_remove_stateful_session_tracking(session_id)
delete_capacity_response: Final = JSONResponse(
status_code=429,
content={
"error": "Too Many Requests",
"details": "Too many recently terminated MCP sessions. Retry after idle records expire.",
},
)
await delete_capacity_response(scope, receive, local_send)
return
# Reserve before yielding so concurrent DELETEs cannot exceed the cap.
# Never evict fresh records: that would allow terminated IDs to replay.
_client_terminated_session_owners[deleting_session] = delete_owner
_terminated_session_ids[deleting_session] = time.monotonic()
try:
async with _gateway_initialize_instructions_request_scope(
user_api_key_auth,
mcp_servers,
_client_ip,
scoped_server_endpoint=scoped_server_endpoint,
is_initialize=is_initialize,
):
await target_manager.handle_request(scope, receive, local_send)
finally:
if deleting_session:
if deleting_session in _stateful_server_instances():
_client_terminated_session_owners.pop(deleting_session, None)
_terminated_session_ids.pop(deleting_session, None)
else:
_terminated_session_ids[deleting_session] = time.monotonic()
_remove_stateful_session_tracking(deleting_session)
try:
if session_lock is not None:

View file

@ -370,7 +370,7 @@ async def test_delete_stale_mcp_session_returns_success():
@pytest.mark.asyncio
@pytest.mark.parametrize("outcome", ("deleted", "rejected", "exception"))
@pytest.mark.parametrize("outcome", ("deleted", "rejected", "exception", "cancelled", "owner-cap", "worker-cap", "expired", "other-owner", "queued-delete"))
async def test_delete_preserves_tracking_until_the_session_is_terminated(outcome):
try:
from litellm.proxy._experimental.mcp_server.server import (
@ -387,6 +387,8 @@ async def test_delete_preserves_tracking_until_the_session_is_terminated(outcome
except ImportError:
pytest.skip("MCP server not available")
from litellm.proxy._experimental.mcp_server import server
session_id = "delete-failure-session"
from litellm.proxy._experimental.mcp_server.auth.litellm_auth_handler import MCPAuthenticatedUser
@ -415,16 +417,28 @@ async def test_delete_preserves_tracking_until_the_session_is_terminated(outcome
_stateful_session_owners[session_id] = _owner_fingerprint_for(user_auth)
_stateful_session_locks[session_id] = session_lock
queued_deletes = []
queued_send = AsyncMock()
async def dispatch(scope, receive, send):
if outcome == "queued-delete":
queued_deletes.append(asyncio.create_task(handle_streamable_http_mcp(dict(scope), receive, queued_send)))
await asyncio.sleep(0)
if outcome == "cancelled":
raise asyncio.CancelledError
if outcome == "exception":
raise RuntimeError("delete failed")
if outcome == "deleted":
if outcome in ("deleted", "expired", "other-owner", "queued-delete"):
mock_instances.pop(session_id)
await send({"type": "http.response.start", "status": 200 if outcome == "deleted" else 403, "headers": []})
await send({"type": "http.response.start", "status": 200 if outcome in ("deleted", "expired", "other-owner", "queued-delete") else 403, "headers": []})
await send({"type": "http.response.body", "body": b""})
try:
with (
patch.object(server, "_terminated_session_ids", {"retained": 0.0 if outcome == "expired" else server.time.monotonic()}),
patch.object(server, "_client_terminated_session_owners", {"retained": _owner_fingerprint_for(user_auth) if outcome == "owner-cap" else "other-owner"}, create=True),
patch.object(server, "_MAX_STATEFUL_SESSIONS_PER_OWNER", 1 if outcome in ("owner-cap", "expired", "other-owner") else 100),
patch.object(server, "_MAX_CLIENT_TERMINATED_SESSIONS", 1 if outcome == "worker-cap" else 10000, create=True),
patch(
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
new_callable=AsyncMock,
@ -446,8 +460,24 @@ async def test_delete_preserves_tracking_until_the_session_is_terminated(outcome
mock_instances,
),
):
await handle_streamable_http_mcp(scope, receive, send)
if outcome == "deleted":
if outcome == "cancelled":
with pytest.raises(asyncio.CancelledError):
await handle_streamable_http_mcp(scope, receive, send)
else:
await handle_streamable_http_mcp(scope, receive, send)
if queued_deletes:
await asyncio.gather(*queued_deletes)
assert queued_send.await_args_list[0].args[0]["status"] == 200
assert mock_handle_request.await_count == 1
if outcome in ("owner-cap", "worker-cap"):
assert send.await_args_list[0].args[0]["status"] == 429
assert mock_handle_request.await_count == 0
assert session_id in mock_instances
assert session_id not in server._terminated_session_ids
assert _stateful_session_owners[session_id] == _owner_fingerprint_for(user_auth)
return
if outcome in ("deleted", "expired", "other-owner", "queued-delete"):
assert server._client_terminated_session_owners[session_id] == _owner_fingerprint_for(user_auth)
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
@ -462,6 +492,9 @@ async def test_delete_preserves_tracking_until_the_session_is_terminated(outcome
assert repeated_delete.await_args_list[0].args[0]["status"] == 200
return
assert session_id not in server._client_terminated_session_owners
assert session_id not in server._terminated_session_ids
assert mock_handle_request.await_count == 1
assert _stateful_session_auth_contexts[session_id] is auth_context
assert _stateful_session_auth_context_last_seen[session_id] == 1.0
@ -476,6 +509,37 @@ async def test_delete_preserves_tracking_until_the_session_is_terminated(outcome
_stateful_session_locks.pop(session_id, None)
@pytest.mark.parametrize("expiry_path", ("cleanup", "replay"))
def test_expired_client_termination_releases_capacity(expiry_path):
from litellm.proxy._experimental.mcp_server import server
with (
patch.object(server, "_terminated_session_ids", {"expired": 0.0}),
patch.object(server, "_client_terminated_session_owners", {"expired": "owner"}),
patch.object(server, "_stateful_session_active_request_counts", {}),
):
now = server._STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS
if expiry_path == "cleanup":
server._forget_expired_terminated_session_ids(now)
else:
assert server._is_terminated_session_id("expired", now) is False
assert server._terminated_session_ids == {}
assert server._client_terminated_session_owners == {}
def test_inflight_delete_keeps_its_reserved_capacity_until_dispatch_finishes():
from litellm.proxy._experimental.mcp_server import server
with (
patch.object(server, "_terminated_session_ids", {"deleting": 0.0}),
patch.object(server, "_client_terminated_session_owners", {"deleting": "owner"}),
patch.object(server, "_stateful_session_active_request_counts", {"deleting": 1}),
):
server._forget_expired_terminated_session_ids(server._STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS)
assert server._terminated_session_ids == {"deleting": 0.0}
assert server._client_terminated_session_owners == {"deleting": "owner"}
@pytest.mark.asyncio
async def test_valid_mcp_session_id_is_preserved():
"""