diff --git a/litellm/proxy/_experimental/mcp_server/catalog.py b/litellm/proxy/_experimental/mcp_server/catalog.py index 1395c4b6b08..606ceb22444 100644 --- a/litellm/proxy/_experimental/mcp_server/catalog.py +++ b/litellm/proxy/_experimental/mcp_server/catalog.py @@ -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 diff --git a/tests/unit/proxy/_experimental/mcp_server/conftest.py b/tests/unit/proxy/_experimental/mcp_server/conftest.py index 51cab559797..742c96b1fe8 100644 --- a/tests/unit/proxy/_experimental/mcp_server/conftest.py +++ b/tests/unit/proxy/_experimental/mcp_server/conftest.py @@ -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) diff --git a/tests/unit/proxy/_experimental/mcp_server/test_operations.py b/tests/unit/proxy/_experimental/mcp_server/test_operations.py index b16b27ac919..06211f9a0d2 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_operations.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_operations.py @@ -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"