From eaf99edfe30dd824f0efc6326c376f5df1a4095a Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Wed, 13 May 2026 17:32:23 +0000 Subject: [PATCH] Fix MCP initialize session active tracking Co-authored-by: Yassin Kortam --- .../proxy/_experimental/mcp_server/server.py | 61 ++++++++++---- .../mcp_server/test_mcp_server.py | 80 +++++++++++++++++++ 2 files changed, 125 insertions(+), 16 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 67b71d72e0a..f17a488e7a7 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -3148,10 +3148,29 @@ if MCP_AVAILABLE: session_id, asyncio.Lock() ) - track_active_stateful_request = bool(use_stateful and session_id) - if track_active_stateful_request and session_id: - _stateful_session_active_request_counts[session_id] = ( - _stateful_session_active_request_counts.get(session_id, 0) + 1 + active_request_session_id = ( + session_id if use_stateful and session_id else None + ) + if active_request_session_id: + _stateful_session_active_request_counts[active_request_session_id] = ( + _stateful_session_active_request_counts.get( + active_request_session_id, 0 + ) + + 1 + ) + + def _track_initialized_stateful_session( + initialized_session_id: str, + ) -> None: + nonlocal active_request_session_id + if active_request_session_id is not None: + return + active_request_session_id = initialized_session_id + _stateful_session_active_request_counts[initialized_session_id] = ( + _stateful_session_active_request_counts.get( + initialized_session_id, 0 + ) + + 1 ) async def _dispatch() -> None: @@ -3174,6 +3193,7 @@ if MCP_AVAILABLE: _owner_fingerprint_for( user_api_key_auth, oauth2_headers, _client_ip ), + _track_initialized_stateful_session, ) async with _gateway_initialize_instructions_request_scope( @@ -3192,32 +3212,38 @@ if MCP_AVAILABLE: else: await _dispatch() finally: - if track_active_stateful_request and session_id: + if active_request_session_id: active_request_count = ( - _stateful_session_active_request_counts.get(session_id, 0) - 1 + _stateful_session_active_request_counts.get( + active_request_session_id, 0 + ) + - 1 ) if active_request_count > 0: - _stateful_session_active_request_counts[session_id] = ( - active_request_count - ) + _stateful_session_active_request_counts[ + active_request_session_id + ] = active_request_count else: - _stateful_session_active_request_counts.pop(session_id, None) + _stateful_session_active_request_counts.pop( + active_request_session_id, None + ) if ( scope.get("method") != "DELETE" - and session_id in _stateful_session_auth_contexts + and active_request_session_id in _stateful_session_auth_contexts ): - _stateful_session_auth_context_last_seen[session_id] = ( - time.monotonic() - ) + _stateful_session_auth_context_last_seen[ + active_request_session_id + ] = time.monotonic() # Periodic cleanup iterates _stateful_session_auth_context_last_seen, # so locks for untracked sessions must be dropped here. if ( active_request_count <= 0 - and session_id not in _stateful_session_auth_contexts + and active_request_session_id + not in _stateful_session_auth_contexts ): - _stateful_session_locks.pop(session_id, None) + _stateful_session_locks.pop(active_request_session_id, None) except HTTPException: # Re-raise HTTP exceptions to preserve status codes and details raise @@ -3422,6 +3448,7 @@ if MCP_AVAILABLE: send: Send, auth_user: MCPAuthenticatedUser, owner_fingerprint: str, + on_session_registered: Optional[Callable[[str], None]] = None, ) -> Send: async def wrapped_send(message: Message) -> None: if message.get("type") == "http.response.start": @@ -3431,6 +3458,8 @@ if MCP_AVAILABLE: session_id = ( value.decode() if isinstance(value, bytes) else str(value) ) + if on_session_registered is not None: + on_session_registered(session_id) auth_context_var.set(auth_user) _stateful_session_auth_contexts[session_id] = auth_user _stateful_session_auth_context_last_seen[session_id] = ( 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 b5d255093f0..790f2829654 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 @@ -1559,6 +1559,86 @@ async def test_initialize_response_capture_accepts_str_headers_and_sets_auth_con mcp_server._remove_stateful_session_tracking(session_id) +@pytest.mark.asyncio +async def test_initialize_request_tracks_active_session_after_response_header(): + try: + from litellm.proxy._experimental.mcp_server import server as mcp_server + from litellm.proxy._experimental.mcp_server.server import ( + handle_streamable_http_mcp, + session_manager_stateful, + session_manager_stateless, + ) + except ImportError: + pytest.skip("MCP server not available") + + session_id = "initialize-active-session-1" + owner_auth = UserAPIKeyAuth(api_key="initialize-key", user_id="user-a") + scope = { + "type": "http", + "method": "POST", + "path": "/mcp", + "headers": [ + (b"content-type", b"application/json"), + (b"authorization", b"Bearer initialize-key"), + ], + } + receive = AsyncMock( + return_value={ + "type": "http.request", + "body": b'{"jsonrpc":"2.0","id":1,"method":"initialize","params":{}}', + "more_body": False, + } + ) + + async def stateful_handle(s, r, se): + await se( + { + "type": "http.response.start", + "headers": [(b"mcp-session-id", session_id.encode())], + } + ) + assert mcp_server._stateful_session_active_request_counts[session_id] == 1 + now = ( + mcp_server._stateful_session_auth_context_last_seen[session_id] + + mcp_server._STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS + ) + await mcp_server._purge_expired_stateful_session_auth_contexts(now=now) + assert session_id in mcp_server._stateful_session_auth_contexts + + async def stateless_handle(s, r, se): + raise AssertionError("initialize request should use stateful manager") + + try: + with ( + patch( + "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", + new_callable=AsyncMock, + return_value=(owner_auth, None, None, None, None, None), + ), + patch( + "litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", + True, + ), + patch.object( + session_manager_stateful, + "handle_request", + side_effect=stateful_handle, + ), + patch.object( + session_manager_stateless, + "handle_request", + side_effect=stateless_handle, + ), + patch.object(session_manager_stateful, "_server_instances", {}), + ): + await handle_streamable_http_mcp(scope, receive, AsyncMock()) + + assert session_id not in mcp_server._stateful_session_active_request_counts + assert session_id in mcp_server._stateful_session_auth_contexts + finally: + mcp_server._remove_stateful_session_tracking(session_id) + + @pytest.mark.asyncio async def test_stateful_mcp_auth_contexts_expire_with_idle_sessions(): """Expired session auth contexts should not remain in memory indefinitely."""