diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index bf2e1298091..7ee49baf2f5 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -3158,15 +3158,9 @@ if MCP_AVAILABLE: mcp_servers, _client_ip, ): - try: - await target_manager.handle_request(scope, receive, local_send) - finally: - if ( - use_stateful - and session_id - and scope.get("method") == "DELETE" - ): - _remove_stateful_session_tracking(session_id) + await target_manager.handle_request(scope, receive, local_send) + if use_stateful and session_id and scope.get("method") == "DELETE": + _remove_stateful_session_tracking(session_id) try: if session_lock is not None: 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 20f9c6e3bf3..b7b9c984bf9 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 @@ -7,6 +7,7 @@ they may send a stale `mcp-session-id` header. This test verifies that: 2. For DELETE requests: idempotent behavior returns success even if session doesn't exist """ +import asyncio from unittest.mock import AsyncMock, MagicMock, patch from fastapi import HTTPException @@ -375,6 +376,91 @@ async def test_delete_stale_mcp_session_returns_success(): assert send.called, "A response should have been sent" +@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. + """ + try: + from litellm.proxy._experimental.mcp_server.server import ( + _owner_fingerprint_for, + _stateful_session_auth_context_last_seen, + _stateful_session_auth_contexts, + _stateful_session_locks, + _stateful_session_owners, + handle_streamable_http_mcp, + session_manager_stateful, + ) + except ImportError: + 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() + session_lock = asyncio.Lock() + mock_instances = {session_id: MagicMock()} + + scope = { + "type": "http", + "method": "DELETE", + "path": "/mcp", + "headers": [ + (b"content-type", b"application/json"), + (b"mcp-session-id", session_id.encode()), + (b"authorization", b"Bearer sk-test"), + ], + } + receive = AsyncMock() + send = AsyncMock() + + _stateful_session_auth_contexts[session_id] = auth_context + _stateful_session_auth_context_last_seen[session_id] = 1.0 + _stateful_session_owners[session_id] = _owner_fingerprint_for(user_auth) + _stateful_session_locks[session_id] = session_lock + + try: + with ( + patch( + "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", + new_callable=AsyncMock, + return_value=(user_auth, None, None, None, None, None), + ), + patch( + "litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", + True, + ), + patch.object( + session_manager_stateful, + "handle_request", + new_callable=AsyncMock, + side_effect=RuntimeError("delete failed"), + ) as mock_handle_request, + patch.object( + session_manager_stateful, + "_server_instances", + mock_instances, + ), + ): + await handle_streamable_http_mcp(scope, receive, send) + + 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 + assert _stateful_session_owners[session_id] == _owner_fingerprint_for( + user_auth + ) + assert _stateful_session_locks[session_id] is session_lock + assert session_id in mock_instances + finally: + _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) + _stateful_session_locks.pop(session_id, None) + + @pytest.mark.asyncio async def test_valid_mcp_session_id_is_preserved(): """