From 8b598907891a15800b9aeddc2c6a3b7552f61b69 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 23 Sep 2026 23:33:08 +0000 Subject: [PATCH] fix(mcp): only keep upstream sessions for gateway-issued session ids and end them on dead streams Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/experimental_mcp_client/client.py | 11 ++++++++++- .../mcp_server/mcp_server_manager.py | 7 ++++++- litellm/proxy/_experimental/mcp_server/server.py | 1 + .../experimental_mcp_client/test_mcp_client.py | 15 +++++++++++++++ .../mcp_server/test_mcp_server_manager.py | 10 ++++++++++ 5 files changed, 42 insertions(+), 2 deletions(-) diff --git a/litellm/experimental_mcp_client/client.py b/litellm/experimental_mcp_client/client.py index 2dab8757100..0af1164c7dd 100644 --- a/litellm/experimental_mcp_client/client.py +++ b/litellm/experimental_mcp_client/client.py @@ -1191,6 +1191,15 @@ class MCPClient: _PendingOperation: TypeAlias = "tuple[Callable[[ClientSession], Awaitable[object]], asyncio.Future[object]]" _MAX_PENDING_OPERATIONS: Final = 64 +_SESSION_ENDING_ERRORS: Final = ( + ValueError, + httpx2.HTTPError, + OSError, + MCPError, + anyio.BrokenResourceError, + anyio.ClosedResourceError, + anyio.EndOfStream, +) class UpstreamSessionClosedError(RuntimeError): @@ -1233,7 +1242,7 @@ class PersistentMCPSession: future.set_exception(RuntimeError("upstream MCP operation was cancelled")) else: future.set_result(outcome) - if isinstance(outcome, (ValueError, httpx2.HTTPError, OSError, MCPError)): + if isinstance(outcome, _SESSION_ENDING_ERRORS): return self._active = None diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index de40030c95a..60702c2cb67 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -1869,6 +1869,7 @@ class MCPServerManager: ) self.registry: dict[str, MCPServer] = {} self._upstream_sessions: dict[tuple[str, str, str], PersistentMCPSession] = {} # mutable-ok: session registry + self._live_gateway_sessions: frozenset[str] = frozenset() self._openapi_health_probes: Callable[[str], _OpenAPIHealthProbe] = lru_cache(maxsize=128)(_OpenAPIHealthProbe) self.config_mcp_servers: dict[str, MCPServer] = {} """ @@ -5798,7 +5799,7 @@ class MCPServerManager: ), None, ) - if gateway_session_id is None or mcp_server.transport == MCPTransport.stdio: + if gateway_session_id not in self._live_gateway_sessions or mcp_server.transport == MCPTransport.stdio: return None key: Final = (gateway_session_id, mcp_server.server_id, await client.discovery_auth_fingerprint()) existing: Final = self._upstream_sessions.get(key) @@ -5808,7 +5809,11 @@ class MCPServerManager: self._upstream_sessions[key] = opened return opened + def track_gateway_session(self, gateway_session_id: str) -> None: + self._live_gateway_sessions = self._live_gateway_sessions | frozenset((gateway_session_id,)) + def release_upstream_sessions(self, gateway_session_id: str) -> None: + self._live_gateway_sessions = self._live_gateway_sessions - frozenset((gateway_session_id,)) for key in tuple(key for key in self._upstream_sessions if key[0] == gateway_session_id): self._upstream_sessions.pop(key).close() diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 815eb923863..78b813a5366 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -2569,6 +2569,7 @@ if MCP_AVAILABLE: _stateful_session_auth_contexts[session_id] = auth_user _stateful_session_auth_context_last_seen[session_id] = time.monotonic() _stateful_session_owners[session_id] = owner_fingerprint + operations.global_mcp_server_manager.track_gateway_session(session_id) if client_info is not None: _stateful_session_client_info[session_id] = client_info break diff --git a/tests/test_litellm/experimental_mcp_client/test_mcp_client.py b/tests/test_litellm/experimental_mcp_client/test_mcp_client.py index 7a1a796bb7b..8edfa899fa1 100644 --- a/tests/test_litellm/experimental_mcp_client/test_mcp_client.py +++ b/tests/test_litellm/experimental_mcp_client/test_mcp_client.py @@ -3060,3 +3060,18 @@ async def test_closing_persistent_session_fails_a_caller_blocked_on_a_full_queue outcomes: Final = await asyncio.wait_for(asyncio.gather(*waiters, return_exceptions=True), 5) assert all(isinstance(outcome, RuntimeError) for outcome in outcomes), outcomes await asyncio.wait_for(session.wait_closed(), 5) + + +@pytest.mark.asyncio +async def test_persistent_session_ends_after_a_broken_stream_so_the_next_call_gets_a_fresh_one(): + app: Final = _stateful_upstream() + async with app.router.lifespan_context(app): + _, session = _client_with_session(app) + + async def broken_stream(_: object) -> str: + raise anyio.BrokenResourceError() + + with pytest.raises(anyio.BrokenResourceError): + await asyncio.wait_for(session.run(broken_stream), 5) + await asyncio.wait_for(session.wait_closed(), 5) + assert session.closed, "a dead transport must end the session instead of being reused" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 2db8f946475..7c9bafdca23 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -10383,6 +10383,7 @@ class TestOBOConcurrencyLimit: manager = MCPServerManager() manager._create_mcp_client = AsyncMock(return_value=_SessionRecordingClient()) + manager.track_gateway_session("gateway-1") for tool in ("select_project", "create_feature"): result = await manager._call_regular_mcp_tool( @@ -14643,7 +14644,12 @@ async def test_upstream_session_is_shared_per_gateway_session_and_released_with_ try: assert await manager._upstream_session_for(client, server, None) is None assert await manager._upstream_session_for(client, server, {"accept": "application/json"}) is None + assert await manager._upstream_session_for(client, server, {"mcp-session-id": "forged"}) is None, ( + "an mcp-session-id the gateway never issued must not open a long-lived upstream session" + ) + manager.track_gateway_session("gw-1") + manager.track_gateway_session("gw-2") first: Final = await manager._upstream_session_for(client, server, {"Mcp-Session-Id": "gw-1"}) second: Final = await manager._upstream_session_for(client, server, {"mcp-session-id": "gw-1"}) other: Final = await manager._upstream_session_for(client, server, {"mcp-session-id": "gw-2"}) @@ -14653,6 +14659,10 @@ async def test_upstream_session_is_shared_per_gateway_session_and_released_with_ manager.release_upstream_sessions("gw-1") await asyncio.wait_for(first.wait_closed(), 5) assert first.closed and not other.closed + assert await manager._upstream_session_for(client, server, {"mcp-session-id": "gw-1"}) is None, ( + "a released gateway session must not reopen upstream sessions" + ) + manager.track_gateway_session("gw-1") replacement: Final = await manager._upstream_session_for(client, server, {"mcp-session-id": "gw-1"}) assert replacement is not first and not replacement.closed