diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260923000000_add_mcp_catalog_revision_trigger/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260923000000_add_mcp_catalog_revision_trigger/migration.sql new file mode 100644 index 00000000000..bc43353063e --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260923000000_add_mcp_catalog_revision_trigger/migration.sql @@ -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(); diff --git a/litellm/proxy/_experimental/mcp_server/catalog.py b/litellm/proxy/_experimental/mcp_server/catalog.py index ecf71089731..44cbb3c8203 100644 --- a/litellm/proxy/_experimental/mcp_server/catalog.py +++ b/litellm/proxy/_experimental/mcp_server/catalog.py @@ -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) diff --git a/litellm/proxy/_experimental/mcp_server/db.py b/litellm/proxy/_experimental/mcp_server/db.py index ee7298fe146..96921056a29 100644 --- a/litellm/proxy/_experimental/mcp_server/db.py +++ b/litellm/proxy/_experimental/mcp_server/db.py @@ -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, diff --git a/litellm/repositories/config_repository.py b/litellm/repositories/config_repository.py index 8b8280622fd..bbc28936234 100644 --- a/litellm/repositories/config_repository.py +++ b/litellm/repositories/config_repository.py @@ -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: ... diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index b7959959688..4fc7b0b6bcc 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py @@ -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), diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_byok_oauth_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_byok_oauth_endpoints.py index 204b50d0372..eed9caa72ee 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_byok_oauth_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_byok_oauth_endpoints.py @@ -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: diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_db_credentials.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_db_credentials.py index cfcff73b857..68625dda3a2 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_db_credentials.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_db_credentials.py @@ -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)] ) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 921f144b468..872e3a63ef5 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -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} diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index ba9a509ccf5..c019c226736 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -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() 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 ba26d4024a4..eb0e8cba9f8 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 @@ -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 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_sigv4_auth.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_sigv4_auth.py index 8cf3bc6fcc7..42f8bc4a9ff 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_sigv4_auth.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_sigv4_auth.py @@ -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=[]) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth2_flow_backfill.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth2_flow_backfill.py index c1239c228aa..ce079878328 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth2_flow_backfill.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth2_flow_backfill.py @@ -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 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth_issuer_stamp_backfill.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth_issuer_stamp_backfill.py index b6c946b95fa..a7e9ac80a1d 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth_issuer_stamp_backfill.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth_issuer_stamp_backfill.py @@ -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