From 3a4cd4a8aed41026be7d68edb1b06c951bab0842 Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Tue, 22 Sep 2026 13:40:32 -0700 Subject: [PATCH] fix(mcp): coalesce queued catalog reads without a stale window --- .../proxy/_experimental/mcp_server/catalog.py | 31 +++++-- .../mcp_server/test_mcp_server_manager.py | 86 +++++++++++++++++++ 2 files changed, 110 insertions(+), 7 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/catalog.py b/litellm/proxy/_experimental/mcp_server/catalog.py index a7dabdb0a0e..a3e5508e629 100644 --- a/litellm/proxy/_experimental/mcp_server/catalog.py +++ b/litellm/proxy/_experimental/mcp_server/catalog.py @@ -15,6 +15,7 @@ if TYPE_CHECKING: _P: Final = ParamSpec("_P") _R: Final = TypeVar("_R") +_REFRESH_FAILURE: Final = "MCP server configuration could not be refreshed" @dataclass(slots=True) @@ -29,6 +30,9 @@ class TargetCatalog: self._manager = manager self._reload_lock = asyncio.Lock() self._scope: ContextVar[_CatalogScope | None] = ContextVar("mcp_catalog_scope", default=None) + self._arrival_ticket = 0 + self._completed_ticket = 0 + self._snapshot: Mapping[str, MCPServer] | None = None @property def current(self) -> Mapping[str, MCPServer] | None: @@ -41,8 +45,17 @@ class TargetCatalog: async def _reload(self, *, reuse_unchanged: bool = False) -> None: token: Final = self._scope.set(None) + # A shared read must start after each covered operation arrives. + covered_ticket: Final = self._arrival_ticket try: await self._manager._reload_servers_from_database(reuse_unchanged=reuse_unchanged) # pyright: ignore[reportPrivateUsage] # existing staged loader + except Exception: + self._snapshot = None + self._completed_ticket = covered_ticket + raise + else: + self._snapshot = MappingProxyType(self._manager.config_mcp_servers | self._manager.registry) + self._completed_ticket = covered_ticket finally: self._scope.reset(token) @@ -64,22 +77,26 @@ class TargetCatalog: raise HTTPException(status_code=503, detail="MCP server configuration changed; retry the operation") async def list(self) -> Mapping[str, MCPServer]: + from fastapi import HTTPException # noqa: PLC0415 # optional proxy dependency + from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415 # runtime proxy dependency scope: Final = self._scope.get() if scope is not None and scope.active and scope.owner_task_id == id(asyncio.current_task()): return scope.servers + self._arrival_ticket += 1 + arrival_ticket: Final = self._arrival_ticket async with self._reload_lock: - if prisma_client is not None: + if prisma_client is None: + return MappingProxyType(self._manager.config_mcp_servers | self._manager.registry) + if arrival_ticket > self._completed_ticket: try: await self._reload(reuse_unchanged=True) except Exception as exc: # noqa: BLE001 # never serve an unverified database snapshot - from fastapi import HTTPException # noqa: PLC0415 # optional proxy dependency - - raise HTTPException( - status_code=503, detail="MCP server configuration could not be refreshed" - ) from exc - return MappingProxyType(self._manager.config_mcp_servers | self._manager.registry) + raise HTTPException(status_code=503, detail=_REFRESH_FAILURE) from exc + if self._snapshot is None: + raise HTTPException(status_code=503, detail=_REFRESH_FAILURE) + return self._snapshot @asynccontextmanager async def operation(self) -> AsyncIterator[Mapping[str, MCPServer]]: 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 b031280f975..b0ef6d53c21 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 @@ -14573,6 +14573,92 @@ async def test_catalog_database_failure_keeps_registry_but_rejects_operation(mon assert manager.catalog.current is None +@pytest.mark.asyncio +@pytest.mark.parametrize("failed", [False, True]) +async def test_catalog_coalesces_waiters_only_behind_a_read_started_after_their_arrival(monkeypatch, failed): + started = asyncio.Event() + release = asyncio.Event() + old_row = _catalog_row() + new_row = _catalog_row("updated") + read_rows = AsyncMock(return_value=[old_row]) + _catalog_database(monkeypatch, read_rows) + manager = MCPServerManager() + previous = await manager.catalog.list() + read_rows.reset_mock() + + async def read_current_rows(**kwargs): + first = not started.is_set() + if first: + started.set() + await release.wait() + if failed: + raise RuntimeError("unavailable database") + return [old_row if first else new_row] + + read_rows.side_effect = read_current_rows + first = asyncio.create_task(manager.catalog.list()) + await started.wait() + followers = [asyncio.create_task(manager.catalog.list()) for _ in range(8)] + await asyncio.sleep(0) + release.set() + results = await asyncio.gather(first, *followers, return_exceptions=True) + if failed: + assert all(isinstance(result, HTTPException) and result.status_code == 503 for result in results) + assert manager.registry == dict(previous) + else: + assert results[0]["catalog-server"].name == "initial" + assert all(result["catalog-server"].name == "updated" for result in results[1:]) + assert read_rows.await_count == 2 + + read_rows.side_effect = None + read_rows.return_value = [new_row] + assert (await manager.catalog.list())["catalog-server"].name == "updated" + assert read_rows.await_count == 3 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("interruption", ["background", "cancel_reader", "cancel_waiter"]) +async def test_catalog_queued_refresh_preserves_freshness_through_cancellation(monkeypatch, interruption): + started = asyncio.Event() + release = asyncio.Event() + + async def read_current_rows(**kwargs): + first = not started.is_set() + if first: + started.set() + await release.wait() + return [_catalog_row("initial" if first else "updated")] + + read_rows = AsyncMock(side_effect=read_current_rows) + _catalog_database(monkeypatch, read_rows) + manager = MCPServerManager() + first = asyncio.create_task(manager.catalog.list()) + await started.wait() + background = asyncio.create_task(manager.reload_servers_from_database()) if interruption == "background" else None + waiter = asyncio.create_task(manager.catalog.list()) + survivors = [asyncio.create_task(manager.catalog.list()) for _ in range(4)] + await asyncio.sleep(0) + if interruption == "cancel_reader": + first.cancel() + elif interruption == "cancel_waiter": + waiter.cancel() + release.set() + first_result, waiter_result, *results = await asyncio.gather(first, waiter, *survivors, return_exceptions=True) + if background is not None: + await background + if interruption == "cancel_reader": + assert isinstance(first_result, asyncio.CancelledError) + else: + assert first_result["catalog-server"].name == "initial" + if interruption == "cancel_waiter": + assert isinstance(waiter_result, asyncio.CancelledError) + else: + assert waiter_result["catalog-server"].name == "updated" + assert all(result["catalog-server"].name == "updated" for result in results) + assert read_rows.await_count == 2 + assert manager.catalog.current is None + + @pytest.mark.asyncio async def test_catalog_serializes_background_and_operation_refreshes(monkeypatch): entered = asyncio.Event()