fix(mcp): check registered handler freshness and isolate fixtures

This commit is contained in:
Joshua Valluru 2026-10-04 20:05:06 -07:00
parent b0ea4aeca5
commit 76b7bdb314
3 changed files with 46 additions and 3 deletions

View file

@ -2697,12 +2697,16 @@ async def _handle_local_mcp_tool(
"""
import inspect
server: Final = global_mcp_server_manager.server_owning_tool_name_prefix(name)
if server is not None:
global_mcp_server_manager.catalog.assert_current(server)
tool: Final = global_mcp_tool_registry.get_tool(name)
if not tool:
raise HTTPException(status_code=404, detail=f"Tool '{name}' not found")
server: Final = (
global_mcp_server_manager.get_mcp_server_by_id(tool.server_id)
if tool.server_id is not None
else global_mcp_server_manager.server_owning_tool_name_prefix(name)
)
if server is not None:
global_mcp_server_manager.catalog.assert_current(server)
try:
if inspect.iscoroutinefunction(tool.handler):

View file

@ -87,6 +87,10 @@ def _hermetic_mcp_server_registry():
global_mcp_server_manager,
)
from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry
saved_tools = global_mcp_tool_registry.published_tools
global_mcp_tool_registry.published_tools = {}
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)
@ -108,6 +112,7 @@ def _hermetic_mcp_server_registry():
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
global_mcp_tool_registry.published_tools = saved_tools
@pytest.fixture(autouse=True)

View file

@ -819,3 +819,37 @@ async def test_discovery_keeps_one_catalog_revision_across_concurrent_listings(m
assert observed == ["before"] * 4
async with manager.catalog.operation():
assert manager.get_mcp_server_by_id(original.server_id).name == "after"
@pytest.mark.asyncio
@pytest.mark.parametrize("changed_server_id", ["private", "allowed"])
async def test_local_handler_freshness_tracks_registered_owner_with_overlapping_alias(monkeypatch, changed_server_id):
from datetime import datetime, timedelta, timezone
from fastapi import HTTPException
from litellm.proxy import proxy_server
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
monkeypatch.setattr(proxy_server, "prisma_client", None)
manager = MCPServerManager()
now = datetime.now(timezone.utc)
private = MCPServer(server_id="private", name="billing", alias="billing", transport=MCPTransport.http, updated_at=now)
allowed = MCPServer(server_id="allowed", name="billing_admin", alias="billing-admin", transport=MCPTransport.http, updated_at=now)
manager.config_mcp_servers = {server.server_id: server for server in (private, allowed)}
monkeypatch.setattr(operations, "global_mcp_server_manager", manager)
monkeypatch.setattr(global_mcp_tool_registry, "published_tools", {})
handler = AsyncMock(return_value="private result")
global_mcp_tool_registry.register_tool("billing-admin-export", "export", {}, handler, server_id=private.server_id)
async with manager.catalog.operation():
manager.config_mcp_servers[changed_server_id] = manager.config_mcp_servers[changed_server_id].model_copy(update={"updated_at": now + timedelta(seconds=1)})
if changed_server_id == "private":
with pytest.raises(HTTPException) as denied:
await operations._handle_local_mcp_tool("billing-admin-export", {})
assert denied.value.status_code == 503
handler.assert_not_awaited()
else:
result = await operations._handle_local_mcp_tool("billing-admin-export", {})
assert result.is_error is False
handler.assert_awaited_once_with()