fix(mcp): preserve local tool ownership across discovery
Some checks failed
LiteLLM Rust / rust-lint (push) Waiting to run
LiteLLM Rust / rust-test (push) Waiting to run
LiteLLM Rust / rust-wheel (push) Waiting to run
Terraform Provider / gofmt, vet, build, test (push) Waiting to run
Terraform Provider / Provider endpoints vs proxy OpenAPI schema (push) Waiting to run
Terraform Modules / fmt, validate, test (aws) (push) Has been cancelled
Terraform Modules / fmt, validate, test (gcp) (push) Has been cancelled

This commit is contained in:
Joshua Valluru 2026-10-03 18:28:08 -07:00
parent 57dfed6e39
commit b63edce2dc
3 changed files with 101 additions and 0 deletions

View file

@ -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()

View file

@ -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()

View file

@ -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,
)