diff --git a/litellm/experimental_mcp_client/client.py b/litellm/experimental_mcp_client/client.py index c79fe5dfae4..9a812fa5963 100644 --- a/litellm/experimental_mcp_client/client.py +++ b/litellm/experimental_mcp_client/client.py @@ -1205,6 +1205,7 @@ class PersistentMCPSession: self._client: Final = client self._queue: Final[asyncio.Queue[_PendingOperation]] = asyncio.Queue(maxsize=_MAX_PENDING_OPERATIONS) self._ready: Final[asyncio.Future[None]] = loop.create_future() + self._active: asyncio.Future[object] | None = None self._task: Final = loop.create_task(self._serve()) @property @@ -1216,25 +1217,37 @@ class PersistentMCPSession: self._ready.set_result(None) while True: operation, future = await self._queue.get() + if future.done(): + continue + self._active = future try: - future.set_result(await operation(session)) + result: Final = await operation(session) except Exception as e: - future.set_exception(e) + if not future.done(): + future.set_exception(e) if isinstance(e, (ValueError, httpx2.HTTPError, OSError, MCPError)): return + else: + if not future.done(): + future.set_result(result) + self._active = None try: await self._client.run_with_session(drain, quiet_on_error=True) - except BaseException as e: + except Exception as e: if not self._ready.done(): self._ready.set_exception(e) - if not isinstance(e, Exception): - raise finally: - while not self._queue.empty(): - _, future = self._queue.get_nowait() - if not future.done(): - future.set_exception(RuntimeError("upstream MCP session closed")) + self._fail_waiters() + + def _fail_waiters(self) -> None: + pending: Final = (self._ready, self._active, *(future for _, future in self._drained())) + for future in pending: + if future is not None and not future.done(): + future.set_exception(RuntimeError("upstream MCP session closed")) + + def _drained(self) -> tuple[_PendingOperation, ...]: + return tuple(self._queue.get_nowait() for _ in range(self._queue.qsize())) async def run( self, diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 58ec737e13f..84f8769dee5 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -5810,6 +5810,10 @@ class MCPServerManager: for key in tuple(key for key in self._upstream_sessions if key[0] == gateway_session_id): self._upstream_sessions.pop(key).close() + def _drop_upstream_session(self, session: PersistentMCPSession | None) -> None: + for key in tuple(key for key, value in self._upstream_sessions.items() if value is session): + self._upstream_sessions.pop(key).close() + async def _obo_call_tool_with_retry( self, *, @@ -5824,6 +5828,7 @@ class MCPServerManager: user_api_key_auth: UserAPIKeyAuth | None, raw_headers: Mapping[str, str] | None = None, client_ip: str | None = None, + persistent_session: PersistentMCPSession | None = None, ) -> CallToolResult: """Call a token_exchange (OBO) tool; on an upstream 401/403 re-mint the token once and retry. @@ -5834,11 +5839,15 @@ class MCPServerManager: """ try: return await client.call_tool( - call_tool_params, host_progress_callback=host_progress_callback, raise_on_error=True + call_tool_params, + host_progress_callback=host_progress_callback, + raise_on_error=True, + persistent_session=persistent_session, ) except Exception as exc: if _extract_upstream_auth_failure(exc) is None: return MCPClient.error_tool_result(exc) + self._drop_upstream_session(persistent_session) spec: Final = to_server_spec(mcp_server) if spec is not None: await self._cred_provider.invalidate_credentials(to_subject(user_api_key_auth, subject_token), spec) @@ -5852,7 +5861,11 @@ class MCPServerManager: raw_headers=raw_headers, client_ip=client_ip, ) - return await retry_client.call_tool(call_tool_params, host_progress_callback=host_progress_callback) + return await retry_client.call_tool( + call_tool_params, + host_progress_callback=host_progress_callback, + persistent_session=await self._upstream_session_for(retry_client, mcp_server, raw_headers), + ) async def _call_regular_mcp_tool( self, @@ -6041,6 +6054,7 @@ class MCPServerManager: user_api_key_auth=user_api_key_auth, raw_headers=raw_headers, client_ip=client_ip, + persistent_session=persistent_session, ) tool_call_coro = _obo_call_tool_limited() 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 fd49b6d2045..eb806763abd 100644 --- a/tests/test_litellm/experimental_mcp_client/test_mcp_client.py +++ b/tests/test_litellm/experimental_mcp_client/test_mcp_client.py @@ -42,6 +42,7 @@ from pydantic import TypeAdapter, ValidationError import litellm.experimental_mcp_client.client as mcp_client_module from litellm.experimental_mcp_client.client import ( MCPClient, + PersistentMCPSession, _first_non_cancelled_cause, _TransportContext, as_mcp_read_timeout, @@ -2960,3 +2961,55 @@ async def test_persistent_session_keeps_upstream_state_across_tool_calls(): assert created.content[0].text == "a/b" await asyncio.wait_for(session.wait_closed(), 5) assert session.closed + + +def _client_with_session(app) -> tuple[_StatefulUpstreamClient, PersistentMCPSession]: + client: Final = _StatefulUpstreamClient(app, server_url="http://upstream/mcp", transport_type=MCPTransport.http) + return client, client.open_persistent_session() + + +@pytest.mark.asyncio +async def test_persistent_session_survives_a_caller_timeout_on_one_operation(): + app: Final = _stateful_upstream() + async with app.router.lifespan_context(app): + client, session = _client_with_session(app) + try: + with pytest.raises(asyncio.TimeoutError): + await asyncio.wait_for( + client.call_tool( + CallToolRequestParams(name="select_project", arguments={"name": "a"}), + persistent_session=session, + ), + timeout=0, + ) + selected: Final = await client.call_tool( + CallToolRequestParams(name="select_project", arguments={"name": "c"}), persistent_session=session + ) + created: Final = await client.call_tool( + CallToolRequestParams(name="create_feature", arguments={"title": "d"}), persistent_session=session + ) + finally: + session.close() + assert selected.is_error is False, "the session must outlive a timed out call" + assert created.content[0].text == "c/d" + await asyncio.wait_for(session.wait_closed(), 5) + + +@pytest.mark.asyncio +async def test_closing_persistent_session_mid_operation_fails_the_waiter_instead_of_hanging(): + app: Final = _stateful_upstream() + async with app.router.lifespan_context(app): + _, session = _client_with_session(app) + started: Final = asyncio.Event() + + async def slow_operation(_: object) -> str: + started.set() + await asyncio.sleep(30) + return "never" + + waiter: Final = asyncio.ensure_future(session.run(slow_operation)) + await asyncio.wait_for(started.wait(), 5) + session.close() + with pytest.raises(RuntimeError, match="upstream MCP session closed"): + await asyncio.wait_for(waiter, 5) + await asyncio.wait_for(session.wait_closed(), 5) 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 1cd39c65c67..d9214447032 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 @@ -7,6 +7,7 @@ import os import sys from datetime import datetime from pathlib import Path +from types import SimpleNamespace from typing import Any, Dict, Final, Literal, Optional, cast from unittest.mock import AsyncMock, MagicMock, patch @@ -10033,7 +10034,7 @@ class _RetryFakeClient: self._MCPClient = MCPClient self.attempts = 0 - async def call_tool(self, params, host_progress_callback=None, raise_on_error=False): + async def call_tool(self, params, host_progress_callback=None, raise_on_error=False, persistent_session=None): self.attempts += 1 if self._raises is not None: if raise_on_error: @@ -10245,7 +10246,7 @@ class TestOBOConcurrencyLimit: inflight = {"current": 0, "peak": 0} class _ConcurrencyRecordingClient: - async def call_tool(self, params, host_progress_callback=None, raise_on_error=False): + async def call_tool(self, params, host_progress_callback=None, raise_on_error=False, persistent_session=None): inflight["current"] += 1 inflight["peak"] = max(inflight["peak"], inflight["current"]) try: @@ -10293,6 +10294,51 @@ class TestOBOConcurrencyLimit: assert inflight["current"] == 0 assert all(result.is_error is False for result in results) + @pytest.mark.asyncio + async def test_obo_dispatch_reuses_the_gateway_sessions_persistent_upstream_session(self): + server = MCPServer( + server_id="obo-stateful", + name="obo", + url="https://upstream.example/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2_token_exchange, + token_exchange_endpoint="https://idp.example.com/token", + client_id="cid", + client_secret="csec", + ) + sessions_seen = [] + + class _SessionRecordingClient: + async def discovery_auth_fingerprint(self): + return "same-token" + + def open_persistent_session(self): + return SimpleNamespace(closed=False, close=lambda: None) + + async def call_tool(self, params, host_progress_callback=None, raise_on_error=False, persistent_session=None): + sessions_seen.append(persistent_session) + return CallToolResult(content=[], isError=False) + + manager = MCPServerManager() + manager._create_mcp_client = AsyncMock(return_value=_SessionRecordingClient()) + + for tool in ("select_project", "create_feature"): + result = await manager._call_regular_mcp_tool( + mcp_server=server, + original_tool_name=tool, + arguments={}, + tasks=[], + mcp_auth_header=None, + mcp_server_auth_headers=None, + oauth2_headers={"Authorization": "Bearer subject-jwt"}, + raw_headers={"mcp-session-id": "gateway-1"}, + proxy_logging_obj=None, + ) + assert result.is_error is False + + assert len(sessions_seen) == 2 and None not in sessions_seen, sessions_seen + assert sessions_seen[0] is sessions_seen[1], "OBO calls in one gateway session must share one upstream session" + class TestOBOEndpointDiscovery: """An oauth2_token_exchange server with no configured token endpoint discovers it (RFC 9728 ->