diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index e043708dd7a..c1b892605b3 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -5292,6 +5292,8 @@ class MCPServerManager: Returns: List of tools with prefixed names """ + from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry + prefixed_tools: Final = [] prefix: Final = get_server_prefix(server) @@ -5304,6 +5306,11 @@ class MCPServerManager: # short ID) so call_tool can resolve regardless of which form a # caller / cached client is using. for spelling in iter_known_tool_name_spellings(original_name, server): + namespace_owner = self.server_owning_tool_name_prefix(spelling) + if namespace_owner is not None and namespace_owner.server_id != server.server_id: + continue + if namespace_owner is None and global_mcp_tool_registry.get_tool(spelling) is not None: + continue self.tool_name_to_mcp_server_name_mapping[spelling] = prefix verbose_logger.info("Successfully fetched %s tools from server %s", len(prefixed_tools), server.name) @@ -6466,6 +6473,13 @@ class MCPServerManager: Returns: MCPServer if found, None otherwise """ + from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry + + # Local handlers belong to their registered namespace, even if a native + # discovery result or an older cached route claims the same spelling. + if global_mcp_tool_registry.get_tool(tool_name) is not None: + return self.server_owning_tool_name_prefix(tool_name) + registry_servers: Final = list(self.get_registry().values()) prefix_to_server: Final = self._known_prefix_to_server() diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py index f3299a03686..ee78c926fc3 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -17031,3 +17031,89 @@ async def test_catalog_cached_revision_retains_newly_discovered_tool_routes(monk if reader is not None: reader.cancel() await asyncio.gather(reader, return_exceptions=True) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("poisoned_route", [False, True]) +async def test_cached_discovery_cannot_authorize_another_servers_local_handler(monkeypatch, poisoned_route): + from datetime import datetime + + from fastapi import HTTPException + from mcp.types import Tool + + from litellm.proxy._experimental.mcp_server import operations + from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry + + _catalog_database(monkeypatch, AsyncMock(return_value=[]), AsyncMock(return_value=_revision_row(7))) + manager = MCPServerManager() + allowed = MCPServer(server_id="allowed", name="allowed", transport=MCPTransport.http) + private = MCPServer(server_id="private", name="private", transport=MCPTransport.http) + manager.config_mcp_servers = {server.server_id: server for server in (allowed, private)} + manager.published_tool_routes = {"getsecret": "private", "private-getsecret": "private"} + handler = AsyncMock(return_value="private result") + monkeypatch.setattr(global_mcp_tool_registry, "published_tools", {}) + global_mcp_tool_registry.register_tool("private-getsecret", "Private tool", {}, handler) + monkeypatch.setattr(operations, "global_mcp_server_manager", manager) + check = AsyncMock(return_value={}) + monkeypatch.setattr(manager, "pre_call_tool_check", check) + async with manager.catalog.operation(): + manager._create_prefixed_tools([Tool(name="private-getsecret", input_schema={})], allowed) + if poisoned_route: + manager.published_tool_routes["private-getsecret"] = "allowed" + async with manager.catalog.operation(): + with pytest.raises(HTTPException) as denied: + await operations._execute_mcp_tool( + name="private-getsecret", arguments={}, allowed_mcp_servers=[allowed], start_time=datetime.now() + ) + assert denied.value.status_code == 403 + handler.assert_not_awaited() + check.assert_not_awaited() + result = await operations._execute_mcp_tool( + name="private-getsecret", arguments={}, allowed_mcp_servers=[private], start_time=datetime.now() + ) + assert result.is_error is False + handler.assert_awaited_once_with() + assert check.await_args.kwargs["server"].server_id == private.server_id + assert manager._get_mcp_server_from_tool_name("allowed-private-getsecret").server_id == allowed.server_id + + +@pytest.mark.parametrize("local_handler", [False, True]) +def test_discovery_preserves_registered_tool_namespaces(monkeypatch, local_handler): + from mcp.types import Tool + + from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry + + manager = MCPServerManager() + allowed = MCPServer(server_id="allowed", name="allowed", transport=MCPTransport.http) + private = MCPServer(server_id="private", name="private", alias="private-alias", transport=MCPTransport.http) + manager.registry = {server.server_id: server for server in (allowed, private)} + manager.published_tool_routes = {"private-getsecret": "private", "private-alias-getsecret": "private"} + monkeypatch.setattr(global_mcp_tool_registry, "published_tools", {}) + if local_handler: + global_mcp_tool_registry.register_tool("orphan", "Orphan local tool", {}, AsyncMock()) + tools = [Tool(name=name, input_schema={}) for name in ("private-getsecret", "private-alias-getsecret", "orphan")] + listed = manager._create_prefixed_tools(tools, allowed) + assert [tool.name for tool in listed] == ["allowed-" + tool.name for tool in tools] + assert manager.published_tool_routes["private-getsecret"] == "private" + assert manager.published_tool_routes["private-alias-getsecret"] == "private" + assert manager._get_mcp_server_from_tool_name("allowed-private-getsecret").server_id == "allowed" + assert ("orphan" in manager.published_tool_routes) is not local_handler + + +@pytest.mark.asyncio +async def test_cached_catalog_excludes_routes_for_servers_outside_its_snapshot(monkeypatch): + read_rows = AsyncMock(return_value=[]) + _catalog_database(monkeypatch, read_rows, AsyncMock(return_value=_revision_row(7))) + manager = MCPServerManager() + pinned = MCPServer(server_id="pinned", name="pinned", transport=MCPTransport.http) + manager.config_mcp_servers = {pinned.server_id: pinned} + async with manager.catalog.operation(): + assert manager.get_mcp_server_by_id("pinned") is not None + published = MCPServer(server_id="published", name="published", transport=MCPTransport.http) + manager.registry[published.server_id] = published + manager.published_tool_routes = {"search": "published", "known": "pinned"} + async with manager.catalog.operation(): + assert manager.get_mcp_server_by_id("published") is None + assert "search" not in manager.tool_name_to_mcp_server_name_mapping + assert manager._get_mcp_server_from_tool_name("known").server_id == "pinned" + read_rows.assert_awaited_once() diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py index 036a1c80042..fe2a2a1efdc 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py @@ -8632,6 +8632,7 @@ async def test_stateful_mcp_tool_call_uses_current_requests_otel_destinations(_m server = MCPServer( server_id="otel-context-test", name="otelcontext", + server_name="otelcontext", transport=MCPTransport.http, allow_all_keys=True, )