mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(mcp): coalesce queued catalog reads without a stale window
This commit is contained in:
parent
180acefb3d
commit
3a4cd4a8ae
2 changed files with 110 additions and 7 deletions
|
|
@ -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]]:
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue