fix(mcp): read catalog revisions and snapshots from the writer

This commit is contained in:
Joshua Valluru 2026-10-05 12:31:45 -07:00
parent 76b7bdb314
commit 2b13ca365e
9 changed files with 117 additions and 38 deletions

View file

@ -120,22 +120,20 @@ class TargetCatalog:
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 or revision != self._applied_revision:
try:
try:
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 or revision != self._applied_revision:
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"
) from exc
if self._shared_snapshot is None:
raise HTTPException(status_code=503, detail="MCP server configuration could not be refreshed")
return self._shared_snapshot
if self._shared_snapshot is None:
raise HTTPException(status_code=503, detail="MCP server configuration could not be refreshed")
return self._shared_snapshot
except Exception as exc:
raise HTTPException(status_code=503, detail="MCP server configuration could not be refreshed") from exc
async def list(self) -> Mapping[str, MCPServer]:
async with self.operation() as snapshot:

View file

@ -648,7 +648,7 @@ 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(
row: Final = await ConfigRepository(prisma_client, use_writer=True).table.find_unique(
where={"param_name": MCP_CATALOG_REVISION_PARAM_NAME}
)
if row is None:
@ -792,7 +792,7 @@ async def get_runtime_mcp_server_rows(
where: Final[prisma_db_types.LiteLLM_MCPServerTableWhereInput] = {
"OR": [{"approval_status": None}, {"approval_status": {"in": ["active", "approved"]}}]
}
return await _db_find_mcp_server_rows(prisma_client, where)
return await MCPServerRepository(prisma_client, use_writer=True).table.find_many(where=where)
async def get_mcp_server(prisma_client: PrismaClient, server_id: str) -> LiteLLM_MCPServerTable | None:

View file

@ -283,38 +283,40 @@ def test_access_group_membership_follows_edits(gateway: Gateway) -> None:
assert tool_calls(peer.drain()) == ()
def test_peer_worker_observes_create_edit_and_delete_without_restart(gateway: Gateway, peer: Gateway) -> None:
with mcp_peer() as first, mcp_peer() as second, gateway.scenario() as scenario:
def test_peer_worker_observes_create_edit_and_delete_without_restart(gateway: Gateway, tmp_path: Path) -> None:
config: Final = tmp_path / "peer.yaml"
config.write_text(yaml.safe_dump({
"model_list": [],
"general_settings": {"master_key": gateway.key, "store_model_in_db": True, "proxy_config_reload_interval_seconds": 3600},
}))
with (
owned_proxy(gateway, tmp_path, {}, config=config, database_setup=()) as peer,
mcp_peer() as first,
mcp_peer() as second,
gateway.scenario() as scenario,
):
alias: Final = "mgmt" + uuid.uuid4().hex[:8]
identity: Final = register_mcp(scenario, first, alias)
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
eventually(
lambda: peer.client.get(
"/mcp-rest/tools/list", headers={"x-litellm-api-key": key}, params={"server_id": identity}
),
lambda value: value.status_code == 200 and value.json() != [],
seconds=40,
listing: Final = peer.client.get(
"/mcp-rest/tools/list", headers={"x-litellm-api-key": key}, params={"server_id": identity}
)
assert listing.status_code == 200 and listing.json() != [], listing.text
names: Final = tool_names(peer, key, identity)
assert call_tool(peer, key, identity, names["add"], ADD).status_code == 200
assert len(tool_calls(first.drain())) == 1
moved: Final = gateway.request("PUT", "/v1/mcp/server", {"server_id": identity, "url": second.url})
assert moved.status_code == 202, moved.text
eventually(
lambda: call_tool(peer, key, identity, names["add"], ADD),
lambda value: value.status_code == 200 and len(tool_calls(second.drain())) == 1,
seconds=40,
)
called: Final = call_tool(peer, key, identity, names["add"], ADD)
assert called.status_code == 200, called.text
assert tool_calls(first.drain()) == ()
assert len(tool_calls(second.drain())) == 1
delete_mcp(gateway, identity)
eventually(
lambda: peer.client.get(
"/mcp-rest/tools/list", headers={"x-litellm-api-key": key}, params={"server_id": identity}
),
lambda value: value.status_code >= 400 or value.json() == [],
seconds=40,
deleted: Final = peer.client.get(
"/mcp-rest/tools/list", headers={"x-litellm-api-key": key}, params={"server_id": identity}
)
second.drain()
assert call_tool(peer, key, identity, names["add"], ADD).status_code >= 400
assert deleted.status_code in (403, 404) or (deleted.status_code == 200 and deleted.json() == []), deleted.text
assert call_tool(peer, key, identity, names["add"], ADD).status_code in (403, 404)
assert tool_calls(second.drain()) == ()

View file

@ -9348,6 +9348,7 @@ async def test_reload_servers_from_database_hydrates_dcr_clients():
)
prisma = MagicMock()
prisma.writer_db = prisma.db
prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[])
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
@ -12654,6 +12655,7 @@ async def test_authorize_observes_committed_peer_server_changes(monkeypatch, cha
)
read_rows = AsyncMock(return_value=[] if change == "delete" else [row])
prisma = SimpleNamespace(db=SimpleNamespace(litellm_mcpservertable=SimpleNamespace(find_many=read_rows), litellm_config=SimpleNamespace(find_unique=AsyncMock(return_value=None))))
prisma.writer_db = prisma.db
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

@ -1279,6 +1279,7 @@ class TestMCPServerManager:
find_many=AsyncMock(return_value=[_row(cached.server_id, corrupted), _row("healthy-sibling", stored)])
)
prisma = SimpleNamespace(db=SimpleNamespace(litellm_config=SimpleNamespace(find_unique=AsyncMock(return_value=None)), litellm_mcpservertable=table))
prisma.writer_db = prisma.db
monkeypatch.setattr(proxy_server, "prisma_client", prisma)
with caplog.at_level(logging.DEBUG, logger="LiteLLM"):
@ -16488,6 +16489,7 @@ def _catalog_database(monkeypatch, read_rows, read_revision=None):
litellm_config=SimpleNamespace(find_unique=read_revision),
)
)
client.writer_db = client.db
monkeypatch.setattr(proxy_server, "prisma_client", client)
@ -16737,6 +16739,7 @@ async def test_catalog_reload_preserves_concurrent_config_discovery_and_routes(a
return []
prisma: Final = MagicMock()
prisma.writer_db = prisma.db
prisma.db.litellm_mcpservertable.find_many = AsyncMock(side_effect=read_rows)
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
with (
@ -16769,6 +16772,7 @@ async def test_catalog_reload_does_not_restore_replaced_config_credentials_or_ro
return []
prisma: Final = MagicMock()
prisma.writer_db = prisma.db
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):
@ -16850,6 +16854,7 @@ async def test_catalog_reload_keeps_new_route_owner_over_earlier_route(change, m
return []
prisma: Final = MagicMock()
prisma.writer_db = prisma.db
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):
@ -16875,6 +16880,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.writer_db = prisma.db
prisma.db.litellm_mcpservertable.find_many = AsyncMock(side_effect=([row], [updated], []))
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
manager: Final = MCPServerManager()
@ -16902,6 +16908,7 @@ async def test_catalog_rebuilt_unchanged_server_keeps_discovered_tool_routes():
url="https://upstream.example/mcp", transport=MCPTransport.http)
manager: Final = MCPServerManager()
prisma: Final = MagicMock()
prisma.writer_db = prisma.db
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):
@ -16952,6 +16959,7 @@ async def test_catalog_reload_retains_routes_discovered_for_a_late_server(change
"updated_at": row.updated_at + timedelta(seconds=1)})]
prisma: Final = MagicMock()
prisma.writer_db = prisma.db
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())
@ -16990,6 +16998,7 @@ async def test_catalog_openapi_refresh_does_not_restore_removed_operations(tmp_p
row: Final = LiteLLM_MCPServerTable(server_id="spec-refresh", alias="spec_refresh",
url="https://upstream.example", transport=MCPTransport.http, spec_path=str(spec_path))
prisma: Final = MagicMock()
prisma.writer_db = prisma.db
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"})
@ -17032,6 +17041,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.writer_db = prisma.db
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()
@ -17058,6 +17068,7 @@ async def test_catalog_failed_reload_preserves_published_discovery_state():
manager.registry = {server.server_id: server}
manager._upstream_initialize_instructions_by_server_id[server.server_id] = "healthy instructions"
prisma: Final = MagicMock()
prisma.writer_db = prisma.db
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):
@ -17082,6 +17093,7 @@ async def test_catalog_cancellation_retains_state_and_releases_refresh_lock():
manager.registry = {server.server_id: server}
manager._upstream_initialize_instructions_by_server_id[server.server_id] = "healthy instructions"
prisma: Final = MagicMock()
prisma.writer_db = prisma.db
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):
@ -17136,6 +17148,7 @@ async def test_catalog_cancelled_openapi_refresh_retains_tools_and_discovery(mon
row: Final = LiteLLM_MCPServerTable(server_id=server.server_id, alias="staged", transport=MCPTransport.http,
url="https://after.example/mcp", spec_path="after.json", updated_at=stamp + timedelta(seconds=1), auth_type=MCPAuth.oauth2)
prisma: Final = MagicMock()
prisma.writer_db = prisma.db
prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[row])
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
@ -17164,6 +17177,7 @@ async def test_catalog_snapshot_identity_is_independent_of_worker_oauth_discover
row: Final = LiteLLM_MCPServerTable(server_id="identity-server", alias="identity_server", transport=MCPTransport.http,
url="https://upstream.example/mcp", auth_type=MCPAuth.oauth2, updated_at=datetime.now())
prisma: Final = MagicMock()
prisma.writer_db = prisma.db
prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[row])
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
manager: Final = MCPServerManager()
@ -17192,6 +17206,7 @@ async def test_catalog_failed_openapi_row_does_not_publish_partial_handlers(monk
row: Final = LiteLLM_MCPServerTable(server_id="broken", alias="broken", transport=MCPTransport.http,
url="https://upstream.example/mcp", spec_path="broken.json")
prisma: Final = MagicMock()
prisma.writer_db = prisma.db
prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[row])
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
@ -17446,6 +17461,7 @@ async def test_catalog_fresh_lookup_does_not_fall_back_to_stale_grants_when_data
allow_all_keys=True)
manager.registry = {server.server_id: server}
prisma: Final = MagicMock()
prisma.writer_db = prisma.db
prisma.db.litellm_mcpservertable.find_many = AsyncMock(side_effect=RuntimeError("unavailable"))
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
with (
@ -17627,6 +17643,7 @@ class TestSharedIdentifierPrefixWarning:
raw_rows = [MagicMock(model_dump=lambda row=row, **kwargs: row.model_dump(**kwargs)) for row in rows]
repository = MagicMock()
prisma = MagicMock()
prisma.writer_db = prisma.db
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
repository.table.find_many = AsyncMock(return_value=raw_rows)
@ -17684,6 +17701,7 @@ async def test_reload_warns_once_about_a_blocked_stdio_row_that_is_rebuilt_every
manager = MCPServerManager()
repository = MagicMock()
prisma = MagicMock()
prisma.writer_db = prisma.db
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
async def build_from_table(table, **_kwargs):
@ -18539,3 +18557,53 @@ async def test_overlapping_server_prefix_cannot_authorize_registered_openapi_han
assert result.is_error is False
handler.assert_awaited_once_with()
assert check.await_args.kwargs["server"].server_id == private.server_id
@pytest.mark.asyncio
@pytest.mark.parametrize("change", ["update", "delete"])
async def test_catalog_observes_writer_changes_while_read_replica_lags(monkeypatch, change):
from types import SimpleNamespace
from litellm.proxy import proxy_server
from litellm.proxy.db.routing_prisma_wrapper import _RoutedActions
original: Final = _catalog_row()
changed: Final = _catalog_row("updated")
reader: Final = SimpleNamespace(
litellm_mcpservertable=SimpleNamespace(find_many=AsyncMock(return_value=[original])),
litellm_config=SimpleNamespace(find_unique=AsyncMock(return_value=_revision_row(7))),
)
writer: Final = SimpleNamespace(
litellm_mcpservertable=SimpleNamespace(find_many=AsyncMock(return_value=[original])),
litellm_config=SimpleNamespace(find_unique=AsyncMock(return_value=_revision_row(7))),
)
routed: Final = SimpleNamespace(**{
name: _RoutedActions(getattr(writer, name), getattr(reader, name), lambda: True)
for name in ("litellm_mcpservertable", "litellm_config")
})
monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=routed, writer_db=writer))
manager: Final = MCPServerManager()
assert (await manager.catalog.resolve(original.server_id)).name == "initial"
writer.litellm_mcpservertable.find_many.return_value = [changed] if change == "update" else []
writer.litellm_config.find_unique.return_value = _revision_row(8)
selected: Final = await manager.catalog.resolve(original.server_id)
assert (selected.name if selected is not None else None) == ("updated" if change == "update" else None)
writer.litellm_config.find_unique.side_effect = RuntimeError("writer unavailable")
with pytest.raises(HTTPException) as unavailable:
await manager.catalog.resolve(original.server_id)
assert unavailable.value.status_code == 503
@pytest.mark.asyncio
async def test_catalog_revision_failure_preserves_state_and_recovers(monkeypatch):
reader: Final = AsyncMock(side_effect=([_catalog_row()], [_catalog_row("updated")]))
revisions: Final = AsyncMock(side_effect=[_revision_row(7), RuntimeError("private database detail"), _revision_row(8)])
_catalog_database(monkeypatch, reader, revisions)
manager: Final = MCPServerManager()
assert (await manager.catalog.resolve("catalog-server")).name == "initial"
with pytest.raises(HTTPException) as unavailable:
await manager.catalog.resolve("catalog-server")
assert unavailable.value.status_code == 503
assert unavailable.value.detail == "MCP server configuration could not be refreshed"
assert manager.registry["catalog-server"].name == "initial"
assert (await manager.catalog.resolve("catalog-server")).name == "updated"

View file

@ -5405,6 +5405,7 @@ class TestMCPServerManagerReload:
db_row = _make_db_mcp_server("server-1", timestamp)
mock_prisma = MagicMock()
mock_prisma.writer_db = mock_prisma.db
mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[db_row])
mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
with (
@ -5448,6 +5449,7 @@ class TestMCPServerManagerReload:
)
mock_prisma = MagicMock()
mock_prisma.writer_db = mock_prisma.db
mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[db_row])
mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
with (
@ -5502,6 +5504,7 @@ class TestMCPServerManagerReload:
return another_healthy_server
mock_prisma = MagicMock()
mock_prisma.writer_db = mock_prisma.db
mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(
return_value=[healthy_row, bad_row, another_healthy_row]
)
@ -5575,6 +5578,7 @@ class TestMCPServerManagerReload:
raise RuntimeError("blocked address")
mock_prisma = MagicMock()
mock_prisma.writer_db = mock_prisma.db
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 (

View file

@ -218,6 +218,7 @@ async def test_guardrail_observes_saved_server_creation_and_deletion_on_another_
row = LiteLLM_MCPServerTable(server_id="peer-server", alias="peer_server", transport="http",
url="https://upstream.example/mcp")
prisma = MagicMock()
prisma.writer_db = prisma.db
prisma.db.litellm_mcpservertable.find_many = AsyncMock(side_effect=([row], []))
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
data = {"tools": [{"type": "mcp", "server_url": "litellm_proxy/mcp/peer-server"}],

View file

@ -2954,6 +2954,7 @@ class TestTemporaryMCPSessionEndpoints:
)
manager: Final = MCPServerManager()
prisma: Final = MagicMock()
prisma.writer_db = prisma.db
async def persisted_rows(*, where: Mapping[str, object]) -> list[LiteLLM_MCPServerTable]:
return [] if where.get("approval_status") == "draft" else [row]
@ -8380,6 +8381,7 @@ async def test_saved_server_authorize_denial_does_not_dispatch_upstream():
auth_type=MCPAuth.oauth2, url="https://upstream.example/mcp",
authorization_url="https://upstream.example/authorize", token_url="https://upstream.example/token")
prisma: Final = MagicMock()
prisma.writer_db = prisma.db
prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[row])
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
user: Final = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.INTERNAL_USER)
@ -8896,6 +8898,7 @@ def _mock_mcp_resolution_prisma_client(
object_permission: LiteLLM_ObjectPermissionTable | None = None,
) -> MagicMock:
prisma: Final = MagicMock()
prisma.writer_db = prisma.db
prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
prisma.db.litellm_verificationtoken.find_unique = AsyncMock(
return_value=SimpleNamespace(object_permission=key_permission)

View file

@ -646,6 +646,7 @@ async def test_dynamic_route_observes_committed_peer_catalog_changes(monkeypatch
litellm_mcptoolsettable=SimpleNamespace(find_first=AsyncMock(return_value=None)),
litellm_config=SimpleNamespace(find_unique=AsyncMock(return_value=None)),
))
prisma.writer_db = prisma.db
manager = mcp_server_manager.MCPServerManager()
if change != "create":
manager.registry = {old_row.server_id: await manager.build_mcp_server_from_table(old_row)}