diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 9b41472be6e..c7035a8ccf9 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -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: diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index b64fb1a4888..79f401cc1bb 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -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 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py index ec6fdef69ee..1bbd49b73bc 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py @@ -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) diff --git a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py index a9ec575e99b..d5640a8cfb0 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -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])