feat(mcp): gate catalog refresh on a database revision marker

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
joshua 2026-09-23 05:35:23 +00:00
parent 51892184a3
commit c31a3129e3
13 changed files with 159 additions and 15 deletions

View file

@ -0,0 +1,13 @@
CREATE OR REPLACE FUNCTION litellm_bump_mcp_catalog_revision() RETURNS TRIGGER AS $$
BEGIN
INSERT INTO "LiteLLM_Config" ("param_name", "reload_revision")
VALUES ('mcp_catalog', 1)
ON CONFLICT ("param_name") DO UPDATE SET "reload_revision" = "LiteLLM_Config"."reload_revision" + 1;
RETURN NULL;
END;
$$ LANGUAGE plpgsql;
DROP TRIGGER IF EXISTS "litellm_mcp_catalog_revision" ON "LiteLLM_MCPServerTable";
CREATE TRIGGER "litellm_mcp_catalog_revision"
AFTER INSERT OR UPDATE OR DELETE ON "LiteLLM_MCPServerTable"
FOR EACH STATEMENT EXECUTE PROCEDURE litellm_bump_mcp_catalog_revision();

View file

@ -81,6 +81,7 @@ class TargetCatalog:
self._arrival_ticket = 0
self._completed_ticket = 0
self._shared_snapshot: CatalogSnapshot | None = None
self._applied_revision: int | None = None
self._warned_shadowed_config_server_ids: frozenset[str] = frozenset()
self._warned_capturing_config_server_ids: frozenset[str] = frozenset()
self._operation: ContextVar[tuple[CatalogSnapshot, asyncio.Event, int] | None] = ContextVar(
@ -116,12 +117,17 @@ class TargetCatalog:
if not should_load_db_object("mcp"):
return _snapshot(self.manager, self._database_identity)
from litellm.proxy._experimental.mcp_server.db import get_mcp_catalog_revision
revision: Final = await get_mcp_catalog_revision(prisma_client)
if revision is not None and revision == self._applied_revision and self._shared_snapshot is not None:
return self._shared_snapshot
self._arrival_ticket += 1
arrival: Final = self._arrival_ticket
async with self._refresh_lock:
if arrival > self._completed_ticket:
try:
await self._publish_refresh(reuse_unchanged=True)
await self._publish_refresh(revision, reuse_unchanged=True)
except Exception as exc:
raise HTTPException(
status_code=503, detail="MCP server configuration could not be refreshed"
@ -227,10 +233,14 @@ class TargetCatalog:
return resolved
async def reload(self) -> None:
async with self._refresh_lock:
await self._publish_refresh()
from litellm.proxy._experimental.mcp_server.db import get_mcp_catalog_revision
from litellm.proxy.proxy_server import prisma_client
async def _publish_refresh(self, *, reuse_unchanged: bool = False) -> None:
revision: Final = await get_mcp_catalog_revision(prisma_client) if prisma_client is not None else None
async with self._refresh_lock:
await self._publish_refresh(revision)
async def _publish_refresh(self, revision: int | None, *, reuse_unchanged: bool = False) -> None:
covered: Final = self._arrival_ticket
token: Final = self._operation.set(None)
try:
@ -241,6 +251,7 @@ class TargetCatalog:
raise
else:
self._shared_snapshot = _snapshot(self.manager, self._database_identity)
self._applied_revision = revision
self._completed_ticket = covered
finally:
self._operation.reset(token)
@ -522,6 +533,5 @@ def global_manager() -> MCPServerManager:
return global_mcp_server_manager
public_catalog_operation: Final[Callable[[Callable[_P, Awaitable[_R]]], Callable[_P, Awaitable[_R]]]] = (
catalog_operation(global_manager)
)
def public_catalog_operation(function: Callable[_P, Awaitable[_R]]) -> Callable[_P, Awaitable[_R]]:
return catalog_operation(global_manager)(function)

View file

@ -40,6 +40,7 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import (
encrypt_value_helper,
)
from litellm.proxy.utils import PrismaClient
from litellm.repositories.config_repository import ConfigRepository
from litellm.repositories.object_permission_repository import ObjectPermissionRepository
from litellm.repositories.prisma_protocols import TableActions
from litellm.repositories.table_repositories import (
@ -513,6 +514,20 @@ async def _db_find_mcp_server_row(
return await _mcp_server_table_actions(prisma_client).find_unique(where={"server_id": server_id})
MCP_CATALOG_REVISION_PARAM_NAME: Final = "mcp_catalog"
async def get_mcp_catalog_revision(prisma_client: PrismaClient) -> int | None:
"""The ``LiteLLM_Config`` revision the MCP server table trigger last published, or None when
no write has ever bumped it (or the trigger is not installed), meaning always reload."""
row: Final = await ConfigRepository(prisma_client).table.find_unique(
where={"param_name": MCP_CATALOG_REVISION_PARAM_NAME}
)
if row is None:
return None
return int(row.reload_revision or 0)
async def _db_update_mcp_server_row(
prisma_client: PrismaClient,
server_id: str,

View file

@ -22,6 +22,9 @@ class _ConfigRow(Protocol):
@property
def param_value(self) -> object: ...
@property
def reload_revision(self) -> int | None: ...
class _ConfigTable(Protocol):
async def find_unique(self, *, where: Mapping[str, str]) -> _ConfigRow | None: ...

View file

@ -7732,6 +7732,7 @@ class TestGatewaySessionAdmission:
)
prisma = MagicMock()
prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[])
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
with (
patch("litellm.proxy.auth.auth_checks.get_user_object", get_user_object),
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
@ -9824,6 +9825,7 @@ class TestScopedSessionAdmission:
)
prisma = MagicMock()
prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[])
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
with (
patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY),
patch("litellm.proxy.auth.auth_checks.get_user_object", get_user_object),

View file

@ -634,6 +634,7 @@ async def test_execute_byok_tool_missing_credential_advertises_api_key_flow(monk
monkeypatch.setattr(proxy_server, "should_load_db_object", lambda _kind: False)
prisma = MagicMock()
prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[])
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
prisma.db.litellm_mcpusercredentials.find_unique = AsyncMock(return_value=None)
monkeypatch.setattr(proxy_server, "prisma_client", prisma)
with pytest.raises(HTTPException) as exc_info:

View file

@ -464,7 +464,7 @@ class _MapTable:
@pytest.mark.parametrize("quoted", [False, True])
async def test_secret_maps_create_update_round_trip(map_algorithm: str, field: str, quoted: bool) -> None:
table: Final = _MapTable(quoted=quoted)
prisma: Final = SimpleNamespace(db=SimpleNamespace(litellm_mcpservertable=table))
prisma: Final = SimpleNamespace(db=SimpleNamespace(litellm_mcpservertable=table, litellm_config=SimpleNamespace(find_unique=AsyncMock(return_value=None))))
original: Final = {"TOKEN": " sensitive-secret\n", "PREFIX": "v2:gcm:literal", "TEMPLATE": "Bearer ${TOKEN}"}
create: Final = NewMCPServerRequest.model_validate({
"server_id": "srv-map", "transport": "http", "url": "https://up.example.com/mcp", field: original,
@ -537,7 +537,7 @@ async def test_secret_map_rotation_migrates_rekeys_and_preserves_corrupt(
{"server_id": "encrypted", field: old, other: None},
)
prisma: Final = SimpleNamespace(db=SimpleNamespace(
litellm_mcpservertable=table, litellm_mcpserveroauthclient=SimpleNamespace(find_many=AsyncMock(return_value=[]))
litellm_mcpservertable=table, litellm_mcpserveroauthclient=SimpleNamespace(find_many=AsyncMock(return_value=[])), litellm_config=SimpleNamespace(find_unique=AsyncMock(return_value=None))
))
await rotate_mcp_server_credentials_master_key(prisma, touched_by="test", new_master_key="rotated-map-key")
assert table.rows["broken"][field] == corrupt
@ -565,7 +565,7 @@ async def test_bulk_reads_isolate_corrupt_secret_maps(reader, field, map_algorit
]
snapshot = [row.model_dump() for row in rows]
table = SimpleNamespace(find_many=AsyncMock(return_value=rows))
prisma = SimpleNamespace(db=SimpleNamespace(litellm_mcpservertable=table))
prisma = SimpleNamespace(db=SimpleNamespace(litellm_mcpservertable=table, litellm_config=SimpleNamespace(find_unique=AsyncMock(return_value=None))))
result = await reader(prisma, ["broken", "healthy"]) if reader is get_mcp_servers else await reader(prisma)
items = result.items if reader is get_mcp_submissions else result
assert [row.server_id for row in items] == ["healthy"]
@ -584,7 +584,7 @@ async def test_bulk_reads_do_not_swallow_unrelated_validation_errors(reader):
row = _prisma_map_row({"server_id": "invalid", "transport": "unsupported"})
table = SimpleNamespace(find_many=AsyncMock(return_value=[row]))
prisma = SimpleNamespace(db=SimpleNamespace(litellm_mcpservertable=table))
prisma = SimpleNamespace(db=SimpleNamespace(litellm_mcpservertable=table, litellm_config=SimpleNamespace(find_unique=AsyncMock(return_value=None))))
request = reader(prisma, ["invalid"]) if reader is get_mcp_servers else reader(prisma)
with pytest.raises(ValidationError, match="transport"):
await request
@ -1502,6 +1502,7 @@ async def test_master_key_rotation_reencrypts_oauth_client_store(monkeypatch):
prisma = MagicMock()
prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[])
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
prisma.db.litellm_mcpserveroauthclient.find_many = AsyncMock(
return_value=[SimpleNamespace(server_id="config_faros", credentials=blob_old)]
)

View file

@ -9349,6 +9349,7 @@ async def test_reload_servers_from_database_hydrates_dcr_clients():
prisma = MagicMock()
prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[])
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
hydrate_spy = AsyncMock()
with (
@ -12652,7 +12653,7 @@ async def test_authorize_observes_committed_peer_server_changes(monkeypatch, cha
updated_at=stamp + timedelta(seconds=1),
)
read_rows = AsyncMock(return_value=[] if change == "delete" else [row])
prisma = SimpleNamespace(db=SimpleNamespace(litellm_mcpservertable=SimpleNamespace(find_many=read_rows)))
prisma = SimpleNamespace(db=SimpleNamespace(litellm_mcpservertable=SimpleNamespace(find_many=read_rows), litellm_config=SimpleNamespace(find_unique=AsyncMock(return_value=None))))
monkeypatch.setattr(proxy_server, "prisma_client", prisma)
monkeypatch.setattr(
global_mcp_server_manager, "registry", {} if change == "create" else {old_server.server_id: old_server}

View file

@ -5390,6 +5390,7 @@ class TestMCPServerManagerReload:
mock_prisma = MagicMock()
mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[db_row])
mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
@ -5432,6 +5433,7 @@ class TestMCPServerManagerReload:
mock_prisma = MagicMock()
mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[db_row])
mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
@ -5487,6 +5489,7 @@ class TestMCPServerManagerReload:
mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(
return_value=[healthy_row, bad_row, another_healthy_row]
)
mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
@ -5557,6 +5560,7 @@ class TestMCPServerManagerReload:
mock_prisma = MagicMock()
mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[healthy_row, bad_openapi_row])
mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
@ -8718,6 +8722,7 @@ async def test_get_active_submitted_mcp_server_ids_for_user_queries_active_rows(
row.server_id = "submitted-1"
prisma_client = MagicMock()
prisma_client.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[row])
prisma_client.db.litellm_config.find_unique = AsyncMock(return_value=None)
result = await get_active_submitted_mcp_server_ids_for_user(prisma_client, "submitter-user")
@ -8738,6 +8743,7 @@ async def test_get_active_submitted_mcp_server_ids_for_user_empty_user_id_skips_
prisma_client = MagicMock()
prisma_client.db.litellm_mcpservertable.find_many = AsyncMock()
prisma_client.db.litellm_config.find_unique = AsyncMock(return_value=None)
assert await get_active_submitted_mcp_server_ids_for_user(prisma_client, "") == []
prisma_client.db.litellm_mcpservertable.find_many.assert_not_awaited()

View file

@ -1095,7 +1095,7 @@ class TestMCPServerManager:
table = SimpleNamespace(
find_many=AsyncMock(return_value=[_row(cached.server_id, corrupted), _row("healthy-sibling", stored)])
)
prisma = SimpleNamespace(db=SimpleNamespace(litellm_mcpservertable=table))
prisma = SimpleNamespace(db=SimpleNamespace(litellm_config=SimpleNamespace(find_unique=AsyncMock(return_value=None)), litellm_mcpservertable=table))
monkeypatch.setattr(proxy_server, "prisma_client", prisma)
with caplog.at_level(logging.DEBUG, logger="LiteLLM"):
@ -14530,12 +14530,19 @@ def _catalog_row(name="initial"):
)
def _catalog_database(monkeypatch, read_rows):
def _catalog_database(monkeypatch, read_rows, read_revision=None):
from types import SimpleNamespace
from litellm.proxy import proxy_server
client = SimpleNamespace(db=SimpleNamespace(litellm_mcpservertable=SimpleNamespace(find_many=read_rows)))
if read_revision is None:
read_revision = AsyncMock(return_value=None)
client = SimpleNamespace(
db=SimpleNamespace(
litellm_mcpservertable=SimpleNamespace(find_many=read_rows),
litellm_config=SimpleNamespace(find_unique=read_revision),
)
)
monkeypatch.setattr(proxy_server, "prisma_client", client)
@ -14786,6 +14793,7 @@ async def test_catalog_reload_preserves_concurrent_config_discovery_and_routes(a
prisma: Final = MagicMock()
prisma.db.litellm_mcpservertable.find_many = AsyncMock(side_effect=read_rows)
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma),
patch("litellm.proxy._experimental.mcp_server.discoverable_endpoints.hydrate_config_server_dcr_client", side_effect=hydrate),
@ -14817,6 +14825,7 @@ async def test_catalog_reload_does_not_restore_replaced_config_credentials_or_ro
prisma: Final = MagicMock()
prisma.db.litellm_mcpservertable.find_many = AsyncMock(side_effect=read_rows)
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
await manager.reload_servers_from_database()
current: Final = manager.get_mcp_server_by_id(server.server_id)
@ -14897,6 +14906,7 @@ async def test_catalog_reload_keeps_new_route_owner_over_earlier_route(change, m
prisma: Final = MagicMock()
prisma.db.litellm_mcpservertable.find_many = AsyncMock(side_effect=read_rows)
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
await manager.reload_servers_from_database()
if change == "deleted" or not mapped:
@ -14921,6 +14931,7 @@ async def test_catalog_observes_committed_update_and_delete_without_background_r
updated: Final = row.model_copy(update={"url": "https://second.example.com/mcp", "updated_at": timestamp + timedelta(seconds=1)})
prisma: Final = MagicMock()
prisma.db.litellm_mcpservertable.find_many = AsyncMock(side_effect=([row], [updated], []))
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
manager: Final = MCPServerManager()
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma),
@ -14947,6 +14958,7 @@ async def test_catalog_rebuilt_unchanged_server_keeps_discovered_tool_routes():
manager: Final = MCPServerManager()
prisma: Final = MagicMock()
prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[row])
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
await manager.reload_servers_from_database()
server: Final = manager.get_mcp_server_by_id(row.server_id)
@ -14996,6 +15008,7 @@ async def test_catalog_reload_retains_routes_discovered_for_a_late_server(change
prisma: Final = MagicMock()
prisma.db.litellm_mcpservertable.find_many = AsyncMock(side_effect=read_rows)
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
task: Final = asyncio.create_task(publish())
try:
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
@ -15033,6 +15046,7 @@ async def test_catalog_openapi_refresh_does_not_restore_removed_operations(tmp_p
url="https://upstream.example", transport=MCPTransport.http, spec_path=str(spec_path))
prisma: Final = MagicMock()
prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[row])
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
upstream: Final = respx_mock.get("https://upstream.example/retained").respond(200, json={"value": "retained"})
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
await manager.reload_servers_from_database()
@ -15070,6 +15084,7 @@ async def test_catalog_lookup_uses_one_snapshot_until_operation_finishes():
updated: Final = row.model_copy(update={"url": "https://second.example.com/mcp", "updated_at": row.updated_at + timedelta(seconds=1)})
prisma: Final = MagicMock()
prisma.db.litellm_mcpservertable.find_many = AsyncMock(side_effect=([row], [updated], [updated]))
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
manager: Final = MCPServerManager()
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma),
@ -15095,6 +15110,7 @@ async def test_catalog_failed_reload_preserves_published_discovery_state():
manager._upstream_initialize_instructions_by_server_id[server.server_id] = "healthy instructions"
prisma: Final = MagicMock()
prisma.db.litellm_mcpservertable.find_many = AsyncMock(side_effect=RuntimeError("database unavailable"))
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
with pytest.raises(RuntimeError, match="database unavailable"):
await manager.reload_servers_from_database()
@ -15118,6 +15134,7 @@ async def test_catalog_cancellation_retains_state_and_releases_refresh_lock():
manager._upstream_initialize_instructions_by_server_id[server.server_id] = "healthy instructions"
prisma: Final = MagicMock()
prisma.db.litellm_mcpservertable.find_many = AsyncMock(side_effect=blocked_read)
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
with patch("litellm.proxy.proxy_server.prisma_client", prisma):
pending = asyncio.create_task(manager.reload_servers_from_database())
await asyncio.wait_for(entered.wait(), timeout=2)
@ -15171,6 +15188,7 @@ async def test_catalog_cancelled_openapi_refresh_retains_tools_and_discovery(mon
url="https://after.example/mcp", spec_path="after.json", updated_at=stamp + timedelta(seconds=1), auth_type=MCPAuth.oauth2)
prisma: Final = MagicMock()
prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[row])
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
async def cancelled_registration(*args: object, **kwargs: object) -> None:
registry.register_tool("staged-new", "new", {}, lambda: "new")
@ -15198,6 +15216,7 @@ async def test_catalog_snapshot_identity_is_independent_of_worker_oauth_discover
url="https://upstream.example/mcp", auth_type=MCPAuth.oauth2, updated_at=datetime.now())
prisma: Final = MagicMock()
prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[row])
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
manager: Final = MCPServerManager()
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma),
@ -15225,6 +15244,7 @@ async def test_catalog_failed_openapi_row_does_not_publish_partial_handlers(monk
url="https://upstream.example/mcp", spec_path="broken.json")
prisma: Final = MagicMock()
prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[row])
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
async def broken_registration(server: MCPServer, **kwargs: object) -> None:
registry.register_tool("broken-partial", "partial", {}, lambda: "must not run")
@ -15478,6 +15498,7 @@ async def test_catalog_fresh_lookup_does_not_fall_back_to_stale_grants_when_data
manager.registry = {server.server_id: server}
prisma: Final = MagicMock()
prisma.db.litellm_mcpservertable.find_many = AsyncMock(side_effect=RuntimeError("unavailable"))
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma),
patch("litellm.proxy.proxy_server.general_settings", {}),
@ -15509,3 +15530,70 @@ async def test_catalog_cancelled_registration_does_not_publish_partial_handlers(
assert [tool.name for tool in registry.list_tools()] == ["existing"]
assert manager.registry == {}
def _revision_row(revision):
from types import SimpleNamespace
return SimpleNamespace(reload_revision=revision)
@pytest.mark.asyncio
async def test_catalog_unchanged_revision_skips_table_reload(monkeypatch):
read_rows = AsyncMock(side_effect=([_catalog_row()], [_catalog_row("updated")]))
read_revision = AsyncMock(return_value=_revision_row(7))
_catalog_database(monkeypatch, read_rows, read_revision)
manager = MCPServerManager()
assert (await manager.catalog.list())["catalog-server"].name == "initial"
assert (await manager.catalog.list())["catalog-server"].name == "initial"
read_rows.assert_awaited_once()
assert read_revision.await_count == 2
@pytest.mark.asyncio
async def test_catalog_changed_revision_reloads_table(monkeypatch):
read_rows = AsyncMock(side_effect=([_catalog_row()], [_catalog_row("updated")]))
read_revision = AsyncMock(side_effect=[_revision_row(7), _revision_row(8)])
_catalog_database(monkeypatch, read_rows, read_revision)
manager = MCPServerManager()
assert (await manager.catalog.list())["catalog-server"].name == "initial"
assert (await manager.catalog.list())["catalog-server"].name == "updated"
assert read_rows.await_count == 2
assert read_revision.await_count == 2
@pytest.mark.asyncio
async def test_catalog_absent_revision_row_reloads_table_every_operation(monkeypatch):
read_rows = AsyncMock(side_effect=([_catalog_row()], [_catalog_row("updated")]))
read_revision = AsyncMock(return_value=None)
_catalog_database(monkeypatch, read_rows, read_revision)
manager = MCPServerManager()
assert (await manager.catalog.list())["catalog-server"].name == "initial"
assert (await manager.catalog.list())["catalog-server"].name == "updated"
assert read_rows.await_count == 2
@pytest.mark.asyncio
async def test_catalog_failed_refresh_does_not_apply_the_read_revision(monkeypatch):
read_rows = AsyncMock(side_effect=([_catalog_row()], RuntimeError("unavailable"), [_catalog_row("updated")]))
read_revision = AsyncMock(side_effect=[_revision_row(7), _revision_row(8), _revision_row(8)])
_catalog_database(monkeypatch, read_rows, read_revision)
manager = MCPServerManager()
assert (await manager.catalog.list())["catalog-server"].name == "initial"
with pytest.raises(HTTPException) as error:
await manager.catalog.list()
assert error.value.status_code == 503
assert (await manager.catalog.list())["catalog-server"].name == "updated"
assert read_rows.await_count == 3
@pytest.mark.asyncio
async def test_catalog_reload_applies_revision_so_next_operation_skips_read(monkeypatch):
read_rows = AsyncMock(return_value=[_catalog_row()])
read_revision = AsyncMock(return_value=_revision_row(7))
_catalog_database(monkeypatch, read_rows, read_revision)
manager = MCPServerManager()
await manager.reload_servers_from_database()
assert (await manager.catalog.list())["catalog-server"].name == "initial"
read_rows.assert_awaited_once()
assert read_revision.await_count == 2

View file

@ -993,6 +993,7 @@ class TestRotateCredentials:
mock_prisma = MagicMock()
mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[server])
mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
mock_prisma.db.litellm_mcpservertable.update = AsyncMock()
mock_prisma.db.litellm_mcpserveroauthclient.find_many = AsyncMock(return_value=[])
@ -1041,6 +1042,7 @@ class TestRotateCredentials:
mock_prisma = MagicMock()
mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[server])
mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
mock_prisma.db.litellm_mcpservertable.update = AsyncMock()
mock_prisma.db.litellm_mcpserveroauthclient.find_many = AsyncMock(return_value=[])

View file

@ -134,6 +134,7 @@ def _byok_key_row(server_id):
def _mock_prisma(null_rows, token_rows):
mock_prisma = MagicMock()
mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=null_rows)
mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
mock_prisma.db.litellm_mcpservertable.update_many = AsyncMock(return_value=MagicMock())
mock_prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(return_value=token_rows)
return mock_prisma

View file

@ -29,6 +29,7 @@ def _row(**overrides):
def _prisma(rows):
prisma_client = MagicMock()
prisma_client.db.litellm_mcpservertable.find_many = AsyncMock(return_value=rows)
prisma_client.db.litellm_config.find_unique = AsyncMock(return_value=None)
prisma_client.db.litellm_mcpservertable.update = AsyncMock()
return prisma_client