From 76b7bdb314a3821047a3a49a7de8c10cb4349d74 Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Sun, 4 Oct 2026 20:05:06 -0700 Subject: [PATCH] fix(mcp): check registered handler freshness and isolate fixtures --- .../_experimental/mcp_server/operations.py | 10 ++++-- .../_experimental/mcp_server/conftest.py | 5 +++ .../mcp_server/test_operations.py | 34 +++++++++++++++++++ 3 files changed, 46 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index 983e56adf9f..36200f658a3 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -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): diff --git a/tests/unit/proxy/_experimental/mcp_server/conftest.py b/tests/unit/proxy/_experimental/mcp_server/conftest.py index 742c96b1fe8..90bb783ad4e 100644 --- a/tests/unit/proxy/_experimental/mcp_server/conftest.py +++ b/tests/unit/proxy/_experimental/mcp_server/conftest.py @@ -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) diff --git a/tests/unit/proxy/_experimental/mcp_server/test_operations.py b/tests/unit/proxy/_experimental/mcp_server/test_operations.py index bdc5769e4b5..38b2e9d3e0e 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_operations.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_operations.py @@ -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()