diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 5456da678e3..ba97dea7332 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -268,6 +268,13 @@ if MCP_AVAILABLE: _stateful_session_locks: Dict[str, asyncio.Lock] = {} _stateful_session_active_request_counts: Dict[str, int] = {} + def _remove_stateful_session_tracking(session_id: str) -> 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) + _stateful_session_locks.pop(session_id, None) + _stateful_session_active_request_counts.pop(session_id, None) + # Keep this alias so existing references to session_manager still work session_manager = session_manager_stateless @@ -302,21 +309,14 @@ if MCP_AVAILABLE: expired_session_ids.append(session_id) for session_id in expired_session_ids: - _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) - _stateful_session_active_request_counts.pop(session_id, None) + _remove_stateful_session_tracking(session_id) transport = server_instances.pop(session_id, None) if transport is not None: await transport.terminate() for session_id in list(_stateful_session_auth_context_last_seen): if session_id not in _stateful_session_auth_contexts: - _stateful_session_auth_context_last_seen.pop(session_id, None) - _stateful_session_owners.pop(session_id, None) - _stateful_session_locks.pop(session_id, None) - _stateful_session_active_request_counts.pop(session_id, None) + _remove_stateful_session_tracking(session_id) async def _cleanup_expired_stateful_session_auth_contexts() -> None: while True: @@ -2821,6 +2821,7 @@ if MCP_AVAILABLE: method = scope.get("method", "").upper() if method == "DELETE": + _remove_stateful_session_tracking(_session_id) verbose_logger.info( "DELETE request for non-existent MCP session '%s'. " "Returning success (idempotent DELETE).", @@ -3120,15 +3121,7 @@ if MCP_AVAILABLE: and session_id and scope.get("method") == "DELETE" ): - _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) - _stateful_session_active_request_counts.pop( - session_id, None - ) + _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 96be37bbd42..20f9c6e3bf3 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 @@ -53,32 +53,55 @@ class TestHandleStaleMcpSession: try: from litellm.proxy._experimental.mcp_server.server import ( _handle_stale_mcp_session, + _stateful_session_active_request_counts, + _stateful_session_auth_context_last_seen, + _stateful_session_auth_contexts, + _stateful_session_locks, + _stateful_session_owners, ) except ImportError: pytest.skip("MCP server not available") + stale_session_id = "stale-id" scope = { "type": "http", "method": "DELETE", "headers": [ (b"content-type", b"application/json"), - (b"mcp-session-id", b"stale-id"), + (b"mcp-session-id", stale_session_id.encode()), ], } receive = AsyncMock() send = AsyncMock() mgr = MagicMock() mgr._server_instances = {} # no active sessions + _stateful_session_auth_contexts[stale_session_id] = MagicMock() + _stateful_session_auth_context_last_seen[stale_session_id] = 1.0 + _stateful_session_owners[stale_session_id] = "owner" + _stateful_session_locks[stale_session_id] = MagicMock() + _stateful_session_active_request_counts[stale_session_id] = 1 - handled = await _handle_stale_mcp_session(scope, receive, send, mgr) + try: + handled = await _handle_stale_mcp_session(scope, receive, send, mgr) - # Should be fully handled (returns True) - assert handled is True - # Should have sent a success response - assert send.called - # Header should NOT be stripped (DELETE needs the session ID) - header_names = [k for k, _ in scope["headers"]] - assert b"mcp-session-id" in header_names + # Should be fully handled (returns True) + assert handled is True + # Should have sent a success response + assert send.called + # Header should NOT be stripped (DELETE needs the session ID) + header_names = [k for k, _ in scope["headers"]] + assert b"mcp-session-id" in header_names + assert stale_session_id not in _stateful_session_auth_contexts + assert stale_session_id not in _stateful_session_auth_context_last_seen + assert stale_session_id not in _stateful_session_owners + assert stale_session_id not in _stateful_session_locks + assert stale_session_id not in _stateful_session_active_request_counts + finally: + _stateful_session_auth_contexts.pop(stale_session_id, None) + _stateful_session_auth_context_last_seen.pop(stale_session_id, None) + _stateful_session_owners.pop(stale_session_id, None) + _stateful_session_locks.pop(stale_session_id, None) + _stateful_session_active_request_counts.pop(stale_session_id, None) @pytest.mark.asyncio async def test_preserves_valid_session_id(self):