mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(mcp): share active catalog snapshots across discovery tasks
This commit is contained in:
parent
ed3c941c12
commit
886aadefef
3 changed files with 49 additions and 4 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue