mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
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:
parent
5fe8415f7c
commit
36147a3baa
4 changed files with 74 additions and 8 deletions
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue