fix(mcp): coalesce queued catalog reads without a stale window

This commit is contained in:
Joshua Valluru 2026-09-22 13:40:32 -07:00
parent 180acefb3d
commit 3a4cd4a8ae
2 changed files with 110 additions and 7 deletions

View file

@ -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]]:

View file

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