mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
51892184a3
commit
c31a3129e3
13 changed files with 159 additions and 15 deletions
|
|
@ -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();
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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: ...
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)]
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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=[])
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue