From 9906303fb40d08e2eb160708e514316dcb63f725 Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Sat, 26 Sep 2026 16:38:36 -0700 Subject: [PATCH] fix(mcp): bound client terminated session retention --- .../proxy/_experimental/mcp_server/server.py | 60 +++++++++++---- .../mcp_server/test_mcp_stale_session.py | 74 +++++++++++++++++-- 2 files changed, 115 insertions(+), 19 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index eb4731de21e..99318af1859 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -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: diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_stale_session.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_stale_session.py index 09dee6025ea..635a9660bbc 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_stale_session.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_stale_session.py @@ -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(): """