diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index f17a488e7a7..4cc7446fd32 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -72,8 +72,8 @@ _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 -# Maximum bytes to peek when sniffing the JSON-RPC method on a no-session-id -# POST. An `initialize` envelope is a few hundred bytes; capping the peek +# Maximum bytes to peek when sniffing the JSON-RPC method on a POST. +# An `initialize` envelope is a few hundred bytes; capping the peek # prevents an authenticated client from forcing the proxy to buffer an # arbitrarily large body just to make a routing decision. _MCP_ROUTING_PEEK_MAX_BYTES = 4096 @@ -3083,7 +3083,7 @@ if MCP_AVAILABLE: return session_id = _get_session_id_from_scope(scope) - if scope.get("method") == "POST" and not session_id: + if scope.get("method") == "POST": consumed_messages, body = await _read_request_body_for_routing(receive) is_initialize = _is_initialize_request(body) @@ -3148,30 +3148,24 @@ if MCP_AVAILABLE: session_id, asyncio.Lock() ) - 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 - ) + active_request_session_ids: List[str] = [] + + def _increment_active_request_session(session_id_to_track: str) -> None: + if session_id_to_track in active_request_session_ids: + return + active_request_session_ids.append(session_id_to_track) + _stateful_session_active_request_counts[session_id_to_track] = ( + _stateful_session_active_request_counts.get(session_id_to_track, 0) + 1 ) + if use_stateful and session_id: + _increment_active_request_session(session_id) + 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 - ) + _increment_active_request_session(initialized_session_id) async def _dispatch() -> None: auth_user = _set_or_update_auth_context( @@ -3212,7 +3206,7 @@ if MCP_AVAILABLE: else: await _dispatch() finally: - if active_request_session_id: + for active_request_session_id in active_request_session_ids: active_request_count = ( _stateful_session_active_request_counts.get( active_request_session_id, 0 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 790f2829654..cf6fdb04a23 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 @@ -1639,6 +1639,110 @@ async def test_initialize_request_tracks_active_session_after_response_header(): mcp_server._remove_stateful_session_tracking(session_id) +@pytest.mark.asyncio +async def test_initialize_request_with_existing_session_tracks_new_session(): + 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") + + existing_session_id = "existing-initialize-session" + new_session_id = "reinitialized-session" + owner_auth = UserAPIKeyAuth(api_key="initialize-key", user_id="user-a") + owner_fingerprint = mcp_server._owner_fingerprint_for(owner_auth) + initialize_body = b'{"jsonrpc":"2.0","id":1,"method":"initialize","params":{}}' + scope = { + "type": "http", + "method": "POST", + "path": "/mcp", + "headers": [ + (b"content-type", b"application/json"), + (b"authorization", b"Bearer initialize-key"), + (b"mcp-session-id", existing_session_id.encode()), + ], + } + receive = AsyncMock( + return_value={ + "type": "http.request", + "body": initialize_body, + "more_body": False, + } + ) + stateful_called = [] + + async def stateful_handle(s, r, se): + stateful_called.append(1) + message = await r() + assert message["body"] == initialize_body + await se( + { + "type": "http.response.start", + "headers": [(b"mcp-session-id", new_session_id.encode())], + } + ) + assert mcp_server._stateful_session_auth_contexts[new_session_id] + assert mcp_server._stateful_session_owners[new_session_id] == owner_fingerprint + assert mcp_server._stateful_session_active_request_counts[new_session_id] == 1 + now = ( + mcp_server._stateful_session_auth_context_last_seen[new_session_id] + + mcp_server._STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS + ) + await mcp_server._purge_expired_stateful_session_auth_contexts(now=now) + assert new_session_id in mcp_server._stateful_session_auth_contexts + + async def stateless_handle(s, r, se): + raise AssertionError( + "initialize request with session should use stateful manager" + ) + + try: + mcp_server._stateful_session_auth_contexts[existing_session_id] = ( + mcp_server.MCPAuthenticatedUser(user_api_key_auth=owner_auth) + ) + mcp_server._stateful_session_auth_context_last_seen[existing_session_id] = 1.0 + mcp_server._stateful_session_owners[existing_session_id] = owner_fingerprint + + 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", + {existing_session_id: MagicMock()}, + ), + ): + await handle_streamable_http_mcp(scope, receive, AsyncMock()) + + assert stateful_called + assert new_session_id not in mcp_server._stateful_session_active_request_counts + assert new_session_id in mcp_server._stateful_session_auth_contexts + finally: + mcp_server._remove_stateful_session_tracking(existing_session_id) + mcp_server._remove_stateful_session_tracking(new_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."""