diff --git a/litellm/experimental_mcp_client/client.py b/litellm/experimental_mcp_client/client.py index e16d1599d0a..fbfbee2c11e 100644 --- a/litellm/experimental_mcp_client/client.py +++ b/litellm/experimental_mcp_client/client.py @@ -1193,6 +1193,11 @@ _PendingOperation: TypeAlias = tuple[Callable[[ClientSession], Awaitable[object] _MAX_PENDING_OPERATIONS: Final = 64 +class UpstreamSessionClosedError(RuntimeError): + def __init__(self) -> None: + super().__init__("upstream MCP session closed") + + class PersistentMCPSession: """One upstream MCP session kept open across operations. @@ -1245,7 +1250,7 @@ class PersistentMCPSession: 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") if cause is None else cause) + future.set_exception(UpstreamSessionClosedError() if cause is None else cause) def _drained(self) -> tuple[_PendingOperation, ...]: return tuple(self._queue.get_nowait() for _ in range(self._queue.qsize())) @@ -1259,10 +1264,17 @@ class PersistentMCPSession: del quiet_on_error await self._ready if self.closed: - raise RuntimeError("upstream MCP session closed") + raise UpstreamSessionClosedError() future: Final[asyncio.Future[object]] = asyncio.get_running_loop().create_future() await self._queue.put((operation, future)) - return cast(TSessionResult, await future) # cast-ok: one queue serves operations of every result type + try: + _ = await asyncio.wait((future, self._task), return_when=asyncio.FIRST_COMPLETED) + except asyncio.CancelledError: + _ = future.cancel() + raise + if not future.done(): + raise UpstreamSessionClosedError() + return cast(TSessionResult, future.result()) # cast-ok: one queue serves operations of every result type def close(self) -> None: _ = self._task.cancel() diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 84f8769dee5..613f78e0c17 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -66,6 +66,7 @@ from litellm.experimental_mcp_client.client import ( MCPClient, MCPSigV4Auth, PersistentMCPSession, + UpstreamSessionClosedError, strip_auth_scheme, to_basic_credentials, ) @@ -5845,12 +5846,14 @@ class MCPServerManager: persistent_session=persistent_session, ) except Exception as exc: - if _extract_upstream_auth_failure(exc) is None: + auth_failure: Final = _extract_upstream_auth_failure(exc) + if auth_failure is None and not isinstance(exc, UpstreamSessionClosedError): 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) + if auth_failure is not None: + 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) retry_client: Final = await self._create_mcp_client( server=mcp_server, mcp_auth_header=server_auth_header, 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 5c3c24e2b07..2566444d3b1 100644 --- a/tests/test_litellm/experimental_mcp_client/test_mcp_client.py +++ b/tests/test_litellm/experimental_mcp_client/test_mcp_client.py @@ -3034,3 +3034,28 @@ async def test_persistent_session_reports_an_upstream_cancellation_as_a_runtime_ session.close() assert selected.is_error is False, "an upstream cancellation must not be mistaken for a cancelled caller" await asyncio.wait_for(session.wait_closed(), 5) + + +@pytest.mark.asyncio +async def test_closing_persistent_session_fails_a_caller_blocked_on_a_full_queue_instead_of_hanging(): + from litellm.experimental_mcp_client.client import _MAX_PENDING_OPERATIONS + + 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" + + waiters: Final = tuple( + asyncio.ensure_future(session.run(slow_operation)) for _ in range(_MAX_PENDING_OPERATIONS + 2) + ) + await asyncio.wait_for(started.wait(), 5) + await asyncio.sleep(0) + session.close() + 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) 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 d9214447032..cbc3f288a60 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 @@ -10217,6 +10217,32 @@ class TestOBOCallToolRetry: manager._create_mcp_client.assert_awaited_once() assert first.attempts == 1 and retry.attempts == 1 + @pytest.mark.asyncio + async def test_session_closed_by_a_peers_refresh_retries_on_a_fresh_session_without_invalidating(self): + from litellm.experimental_mcp_client.client import UpstreamSessionClosedError + + manager = self._manager() + success = CallToolResult(content=[], isError=False) + first = _RetryFakeClient(raises=UpstreamSessionClosedError()) + retry = _RetryFakeClient(result=success) + manager._create_mcp_client = AsyncMock(return_value=retry) + + result = await manager._obo_call_tool_with_retry( + client=first, + call_tool_params=MagicMock(), + host_progress_callback=None, + mcp_server=_obo_server(), + server_auth_header=None, + extra_headers=None, + stdio_env=None, + subject_token="caller-jwt", + user_api_key_auth=None, + ) + + assert result is success + manager._cred_provider.invalidate_credentials.assert_not_awaited() + assert first.attempts == 1 and retry.attempts == 1 + class TestOBOConcurrencyLimit: """OBO (token_exchange) tool calls must honor the server's max_concurrent_requests.