fix(mcp): fail producers blocked on a full queue at teardown and retry peers evicted by an OBO refresh

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
Devin AI 2026-09-23 16:21:49 +00:00
parent 5fe8415f7c
commit 36147a3baa
4 changed files with 74 additions and 8 deletions

View file

@ -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()

View file

@ -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,

View file

@ -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)

View file

@ -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.