From 6ea404613508bde31c58f9b6444d78816c5caa3d Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Tue, 5 May 2026 19:06:57 +0000 Subject: [PATCH] Fix MCP stateful session cleanup --- .../proxy/_experimental/mcp_server/server.py | 54 +++++++++- .../mcp_server/test_mcp_server.py | 101 ++++++++++++------ 2 files changed, 119 insertions(+), 36 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index ff9c724ad43..e2247e601a3 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -70,6 +70,7 @@ from litellm.utils import Rules, client, function_setup _byok_cred_cache: Dict[Tuple[str, str], Tuple[Optional[str], float]] = {} _BYOK_CRED_CACHE_TTL = 60 # seconds _BYOK_CRED_CACHE_MAX_SIZE = 4096 # cap to prevent unbounded growth +_STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS = 30 * 60 def _invalidate_byok_cred_cache(user_id: str, server_id: str) -> None: @@ -251,6 +252,7 @@ if MCP_AVAILABLE: stateless=False, ) _stateful_session_auth_contexts: Dict[str, MCPAuthenticatedUser] = {} + _stateful_session_auth_context_last_seen: Dict[str, float] = {} # Keep this alias so existing references to session_manager still work session_manager = session_manager_stateless @@ -267,10 +269,40 @@ if MCP_AVAILABLE: _session_manager_cm = None _session_manager_stateful_cm = None _sse_session_manager_cm = None + _stateful_auth_context_cleanup_task: Optional[asyncio.Task] = None + + async def _purge_expired_stateful_session_auth_contexts( + now: Optional[float] = None, + ) -> None: + """Terminate expired stateful sessions and drop their auth contexts.""" + now = now or time.monotonic() + server_instances = getattr(session_manager_stateful, "_server_instances", {}) + expired_session_ids = [ + session_id + for session_id, last_seen in _stateful_session_auth_context_last_seen.items() + if now - last_seen >= _STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS + or session_id not in server_instances + ] + + 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) + 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) + + async def _cleanup_expired_stateful_session_auth_contexts() -> None: + while True: + await asyncio.sleep(_STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS) + await _purge_expired_stateful_session_auth_contexts() async def initialize_session_managers(): """Initialize the session managers. Can be called from main app lifespan.""" - global _SESSION_MANAGERS_INITIALIZED, _session_manager_cm, _session_manager_stateful_cm, _sse_session_manager_cm + global _SESSION_MANAGERS_INITIALIZED, _session_manager_cm, _session_manager_stateful_cm, _sse_session_manager_cm, _stateful_auth_context_cleanup_task # Use async lock to prevent concurrent initialization async with _INITIALIZATION_LOCK: @@ -288,6 +320,9 @@ if MCP_AVAILABLE: await _session_manager_cm.__aenter__() await _session_manager_stateful_cm.__aenter__() await _sse_session_manager_cm.__aenter__() + _stateful_auth_context_cleanup_task = asyncio.create_task( + _cleanup_expired_stateful_session_auth_contexts() + ) _SESSION_MANAGERS_INITIALIZED = True verbose_logger.info( @@ -296,12 +331,16 @@ if MCP_AVAILABLE: async def shutdown_session_managers(): """Shutdown the session managers.""" - global _SESSION_MANAGERS_INITIALIZED, _session_manager_cm, _session_manager_stateful_cm, _sse_session_manager_cm + global _SESSION_MANAGERS_INITIALIZED, _session_manager_cm, _session_manager_stateful_cm, _sse_session_manager_cm, _stateful_auth_context_cleanup_task if _SESSION_MANAGERS_INITIALIZED: verbose_logger.info("Shutting down MCP session managers...") try: + if _stateful_auth_context_cleanup_task: + _stateful_auth_context_cleanup_task.cancel() + with contextlib.suppress(asyncio.CancelledError): + await _stateful_auth_context_cleanup_task if _session_manager_cm: await _session_manager_cm.__aexit__(None, None, None) if _session_manager_stateful_cm: @@ -314,6 +353,7 @@ if MCP_AVAILABLE: _session_manager_cm = None _session_manager_stateful_cm = None _sse_session_manager_cm = None + _stateful_auth_context_cleanup_task = None _SESSION_MANAGERS_INITIALIZED = False @contextlib.asynccontextmanager @@ -2958,6 +2998,7 @@ if MCP_AVAILABLE: finally: if use_stateful 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) except HTTPException: # Re-raise HTTP exceptions to preserve status codes and details raise @@ -3133,7 +3174,8 @@ if MCP_AVAILABLE: auth_user = ( _stateful_session_auth_contexts.get(session_id) if session_id else None ) - if auth_user is not None: + if auth_user is not None and session_id is not None: + _stateful_session_auth_context_last_seen[session_id] = time.monotonic() _update_auth_context( auth_user=auth_user, user_api_key_auth=user_api_key_auth, @@ -3164,7 +3206,11 @@ if MCP_AVAILABLE: if message.get("type") == "http.response.start": for key, value in message.get("headers", []): if key.lower() == b"mcp-session-id": - _stateful_session_auth_contexts[value.decode()] = auth_user + session_id = value.decode() + _stateful_session_auth_contexts[session_id] = auth_user + _stateful_session_auth_context_last_seen[session_id] = ( + time.monotonic() + ) break await send(message) 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 e26a3d9050e..240f08b9dfb 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 @@ -1154,31 +1154,39 @@ async def test_mcp_routing_initialize_to_stateful_no_session_to_stateless(): async def stateful_handle(s, r, se): stateful_called.append(1) - with patch( - "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", - new_callable=AsyncMock, - return_value=(MagicMock(), None, ["progress_test"], None, None, None), - ), patch( - "litellm.proxy._experimental.mcp_server.server.set_auth_context", - ), patch( - "litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", - True, - ), patch.object( - session_manager_stateless, - "handle_request", - side_effect=stateless_handle, - ), patch.object( - session_manager_stateful, - "handle_request", - side_effect=stateful_handle, - ), patch.object( - session_manager_stateless, - "_server_instances", - {}, - ), patch.object( - session_manager_stateful, - "_server_instances", - {}, + with ( + patch( + "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", + new_callable=AsyncMock, + return_value=(MagicMock(), None, ["progress_test"], None, None, None), + ), + patch( + "litellm.proxy._experimental.mcp_server.server.set_auth_context", + ), + patch( + "litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", + True, + ), + patch.object( + session_manager_stateless, + "handle_request", + side_effect=stateless_handle, + ), + patch.object( + session_manager_stateful, + "handle_request", + side_effect=stateful_handle, + ), + patch.object( + session_manager_stateless, + "_server_instances", + {}, + ), + patch.object( + session_manager_stateful, + "_server_instances", + {}, + ), ): await handle_streamable_http_mcp(scope, receive, send) @@ -1187,16 +1195,16 @@ async def test_mcp_routing_initialize_to_stateful_no_session_to_stateless(): # initialize → stateful init_body = b'{"jsonrpc":"2.0","id":1,"method":"initialize","params":{"protocolVersion":"2024-11-05"}}' stateless_called, stateful_called = await make_request(init_body) - assert stateful_called and not stateless_called, ( - "initialize (no session) should route to stateful, not stateless" - ) + assert ( + stateful_called and not stateless_called + ), "initialize (no session) should route to stateful, not stateless" # tools/list → stateless tools_body = b'{"jsonrpc":"2.0","id":2,"method":"tools/list","params":{}}' stateless_called, stateful_called = await make_request(tools_body) - assert stateless_called and not stateful_called, ( - "tools/list (no session) should route to stateless, not stateful" - ) + assert ( + stateless_called and not stateful_called + ), "tools/list (no session) should route to stateless, not stateful" @pytest.mark.asyncio @@ -1379,6 +1387,36 @@ async def test_stateful_mcp_requests_refresh_session_auth_context(): mcp_server._stateful_session_auth_contexts.pop(session_id, None) +@pytest.mark.asyncio +async def test_stateful_mcp_auth_contexts_expire_with_idle_sessions(): + """Expired session auth contexts should not remain in memory indefinitely.""" + try: + from litellm.proxy._experimental.mcp_server import server as mcp_server + except ImportError: + pytest.skip("MCP server not available") + + session_id = "expired-stateful-session" + auth_user = UserAPIKeyAuth(api_key="expired-key", user_id="expired-user") + transport = MagicMock() + now = 1000.0 + + mcp_server._stateful_session_auth_contexts[session_id] = auth_user + mcp_server._stateful_session_auth_context_last_seen[session_id] = ( + now - mcp_server._STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS + ) + + with patch.object( + mcp_server.session_manager_stateful, + "_server_instances", + {session_id: transport}, + ): + await mcp_server._purge_expired_stateful_session_auth_contexts(now=now) + + assert session_id not in mcp_server._stateful_session_auth_contexts + assert session_id not in mcp_server._stateful_session_auth_context_last_seen + transport.terminate.assert_awaited_once() + + @pytest.mark.asyncio @pytest.mark.no_parallel async def test_mcp_routing_with_conflicting_alias_and_group_name(): @@ -3405,4 +3443,3 @@ async def test_call_tool_empty_extra_headers_returns_none(): "P2 API consistency issue: expected None for empty extra_headers, got: " + str(captured_extra_headers) ) -