fix(mcp): share active catalog snapshots across discovery tasks

This commit is contained in:
Joshua Valluru 2026-10-03 17:21:35 -07:00
parent ed3c941c12
commit 886aadefef
3 changed files with 49 additions and 4 deletions

View file

@ -85,7 +85,7 @@ class TargetCatalog:
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(
self._operation: ContextVar[tuple[CatalogSnapshot, asyncio.Event] | None] = ContextVar(
"mcp_catalog_snapshot", default=None
)
self._staged_routing: ContextVar[tuple[dict[str, str], asyncio.Event] | None] = ContextVar(
@ -161,8 +161,7 @@ class TargetCatalog:
@asynccontextmanager
async def operation(self) -> AsyncGenerator[CatalogSnapshot]:
current: Final = self.current()
scoped: Final = self._operation.get()
if current is not None and scoped is not None and scoped[2] == id(asyncio.current_task()):
if current is not None:
yield current
return
from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry
@ -174,7 +173,7 @@ class TargetCatalog:
routing=dict(shared.routing),
)
closed: Final = asyncio.Event()
token: Final = self._operation.set((snapshot, closed, id(asyncio.current_task())))
token: Final = self._operation.set((snapshot, closed))
try:
with global_mcp_tool_registry.catalog_scope(snapshot.tools):
yield snapshot

View file

@ -81,10 +81,14 @@ def config_only_mcp_manager_factory():
@pytest.fixture(autouse=True)
def _hermetic_mcp_server_registry():
from litellm.proxy._experimental.mcp_server.catalog import TargetCatalog
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
saved_catalog = global_mcp_server_manager.catalog
global_mcp_server_manager.catalog = TargetCatalog(global_mcp_server_manager)
saved_registry = dict(global_mcp_server_manager.registry)
saved_config_servers = dict(global_mcp_server_manager.config_mcp_servers)
saved_tool_mapping = dict(global_mcp_server_manager.tool_name_to_mcp_server_name_mapping)
@ -103,6 +107,7 @@ def _hermetic_mcp_server_registry():
global_mcp_server_manager.tool_name_to_mcp_server_name_mapping.clear()
global_mcp_server_manager.tool_name_to_mcp_server_name_mapping.update(saved_tool_mapping)
global_mcp_server_manager._oauth_discovery_slots = saved_oauth_slots
global_mcp_server_manager.catalog = saved_catalog
@pytest.fixture(autouse=True)

View file

@ -665,3 +665,44 @@ async def test_tools_listing_preserves_explicit_spend_log_policy(log_enabled):
)
assert result.tools == []
assert listing.await_args.kwargs["log_list_tools_to_spendlogs"] is log_enabled
@pytest.mark.asyncio
async def test_discovery_keeps_one_catalog_revision_across_concurrent_listings(monkeypatch):
from mcp.types import (
DiscoverRequest, ListToolsResult, ListPromptsResult, ListResourcesResult,
ListResourceTemplatesResult,
)
from litellm.proxy import proxy_server
from litellm.proxy._experimental.mcp_server import operations
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
manager = MCPServerManager()
original = MCPServer(server_id="catalog-server", name="before", transport=MCPTransport.http)
updated = original.model_copy(update={"name": "after"})
manager.registry = {original.server_id: original}
monkeypatch.setattr(proxy_server, "prisma_client", None)
monkeypatch.setattr(operations, "global_mcp_server_manager", manager)
observed = []
responses = (ListToolsResult(tools=[]), ListPromptsResult(prompts=[]),
ListResourcesResult(resources=[]), ListResourceTemplatesResult(resource_templates=[]))
def listing(index):
async def run(*args, **kwargs):
async with manager.catalog.operation():
observed.append(manager.get_mcp_server_by_id(original.server_id).name)
manager.registry = {updated.server_id: updated}
return responses[index]
return run
with (
patch.object(operations, "_execute_handle_list_tools", side_effect=listing(0)),
patch.object(operations, "_execute_list_prompts", side_effect=listing(1)),
patch.object(operations, "_execute_list_resources", side_effect=listing(2)),
patch.object(operations, "_execute_list_resource_templates", side_effect=listing(3)),
):
result = await GatewayOperations().execute(DiscoverRequest(), prepare_context(UserAPIKeyAuth(user_id="scoped")))
assert result.capabilities.model_dump(exclude_none=True) == {}
assert observed == ["before"] * 4
async with manager.catalog.operation():
assert manager.get_mcp_server_by_id(original.server_id).name == "after"