fix(mcp): route tool calls by server_id so unprefixed names cannot pick an arbitrary server

tool_name_to_mcp_server_name_mapping mapped a tool name to a server name and was
written with an unqualified, global, last-writer-wins key. Two server entries sharing
an upstream URL expose the same tool names, so that key always collided and the winner
was whichever server registered last, which meant config declaration order decided
which upstream credential left the gateway.

server_name is not a sound identity: schema.prisma has no unique constraint on
server_name or alias, and the create endpoint only checks server_id. The registry is
already keyed by server_id and config servers derive a deterministic one, so server_id
was always the real identity.

The map now holds frozensets of server ids and accumulates owners instead of
overwriting them, so a name served by several reachable servers stays visibly
ambiguous and is rejected rather than dispatched to an arbitrary server. Ambiguity is
judged against the servers the caller can reach, so a session scoped with
x-mcp-servers still resolves its unprefixed names. All three writers now go through one
registration chokepoint; they previously disagreed on what the value was. Cleanup
withdraws only the departing server's id from each row instead of matching by name,
which stopped removing one server from deleting a same-named server's routes.
execute_mcp_tool now dispatches the server it resolved rather than passing a name that
call_tool re-resolved, closing the same drop on the responses API path.
This commit is contained in:
Tin Chi Lo 2026-07-16 23:19:06 -07:00
parent 3459956fd2
commit 7fe0d9501e
8 changed files with 651 additions and 139 deletions

View file

@ -14,6 +14,7 @@ import os
import re
import time
from contextlib import asynccontextmanager
from dataclasses import dataclass
from typing import Any, AsyncIterator, Callable, Literal, Optional, Union, cast
from urllib.parse import urlparse
@ -186,6 +187,31 @@ _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES: tuple[MCPAuth, ...] = (
)
@dataclass(frozen=True, slots=True)
class MCPToolRouteResolved:
"""Exactly one MCP server serves the tool name."""
server: MCPServer
@dataclass(frozen=True, slots=True)
class MCPToolRouteNotFound:
"""No MCP server serves the tool name."""
tool_name: str
@dataclass(frozen=True, slots=True)
class MCPToolRouteAmbiguous:
"""Several MCP servers serve the tool name, so it cannot address one of them."""
tool_name: str
server_ids: frozenset[str]
MCPToolRoute = Union[MCPToolRouteResolved, MCPToolRouteNotFound, MCPToolRouteAmbiguous]
def _blank_to_none(value: str | None) -> str | None:
"""Collapse an absent, empty, or whitespace-only string to ``None``.
@ -1035,7 +1061,7 @@ class MCPServerManager:
# semaphore so an edited limit rebuilds it instead of keeping the old cap
# until restart.
self._server_call_semaphores: dict[str, tuple[int, asyncio.Semaphore]] = {}
self.tool_name_to_mcp_server_name_mapping: dict[str, str] = {}
self.tool_name_to_mcp_server_ids_mapping: dict[str, frozenset[str]] = {}
"""
{
"gmail_send_email": "zapier_mcp_server",
@ -1367,7 +1393,7 @@ class MCPServerManager:
verbose_logger.debug(f"Loaded MCP Servers: {json.dumps(self.config_mcp_servers, indent=4, default=str)}")
self.initialize_tool_name_to_mcp_server_name_mapping()
self.initialize_tool_name_to_mcp_server_ids_mapping()
async def _register_openapi_tools(self, spec_path: str, server: MCPServer, base_url: str):
"""
@ -1462,7 +1488,7 @@ class MCPServerManager:
resolved_operation = resolve_operation_params(operation, path_item, components)
# Generate tool name (without prefix initially)
operation_id = operation.get("operationId", f"{method}_{path.replace('/', '_')}")
operation_id = str(operation.get("operationId", f"{method}_{path.replace('/', '_')}"))
base_tool_name = operation_id.replace(" ", "_").lower()
# Add server prefix to tool name
@ -1490,9 +1516,9 @@ class MCPServerManager:
handler=tool_func,
)
# Update tool name to server name mapping (for both prefixed and base names)
self.tool_name_to_mcp_server_name_mapping[base_tool_name] = server_prefix
self.tool_name_to_mcp_server_name_mapping[prefixed_tool_name] = server_prefix
# Update tool name to server id mapping (for both prefixed and base names)
self._register_tool_route(base_tool_name, server.server_id)
self._register_tool_route(prefixed_tool_name, server.server_id)
registered_count += 1
verbose_logger.debug(f"Registered OpenAPI tool: {prefixed_tool_name} for server {server.name}")
@ -1508,9 +1534,12 @@ class MCPServerManager:
When a server leaves ``self.registry`` (eviction, ``remove_server``, etc.),
OpenAPI tools remain in ``global_mcp_tool_registry`` and
``tool_name_to_mcp_server_name_mapping`` unless removed here. Stale
mappings make ``_get_mcp_server_from_tool_name`` resolve to a prefix that
``tool_name_to_mcp_server_ids_mapping`` unless removed here. Stale
mappings make ``_get_mcp_server_from_tool_name`` resolve to a server that
no longer exists in the live registry.
Only ``server``'s own id is withdrawn from each routing row, so a tool name
shared with another server keeps routing to that other server.
"""
from litellm.proxy._experimental.mcp_server.tool_registry import (
global_mcp_tool_registry,
@ -1521,24 +1550,11 @@ class MCPServerManager:
openapi_key_prefix = prefix_root + MCP_TOOL_PREFIX_SEPARATOR
global_mcp_tool_registry.unregister_tools_with_prefix(openapi_key_prefix)
owned_raw: set[str] = set()
for p in iter_known_server_prefixes(server):
if p:
owned_raw.add(p)
if server.name:
owned_raw.add(server.name)
owned_normalized = {normalize_server_name(x) for x in owned_raw}
stale_mapping_keys: list[str] = []
for tool_name, mapped_server in list(self.tool_name_to_mcp_server_name_mapping.items()):
if mapped_server in owned_raw:
stale_mapping_keys.append(tool_name)
elif normalize_server_name(str(mapped_server)) in owned_normalized:
stale_mapping_keys.append(tool_name)
for key in stale_mapping_keys:
del self.tool_name_to_mcp_server_name_mapping[key]
self.tool_name_to_mcp_server_ids_mapping = {
tool_name: remaining
for tool_name, owner_ids in self.tool_name_to_mcp_server_ids_mapping.items()
if (remaining := owner_ids - {server.server_id})
}
def remove_server(self, mcp_server: LiteLLM_MCPServerTable):
"""
@ -1946,7 +1962,7 @@ class MCPServerManager:
base_url=server.url or "",
)
if initialize_mapping:
self.initialize_tool_name_to_mcp_server_name_mapping()
self.initialize_tool_name_to_mcp_server_ids_mapping()
async def add_server(self, mcp_server: LiteLLM_MCPServerTable):
# The runtime registry is the allowlist for tool calls and health
@ -3818,10 +3834,10 @@ class MCPServerManager:
# Register every known prefix form (alias, server_name, server_id,
# short ID) so call_tool can resolve regardless of which form a
# caller / cached client is using.
self.tool_name_to_mcp_server_name_mapping[original_name] = prefix
self._register_tool_route(original_name, server.server_id)
for known_prefix in iter_known_server_prefixes(server):
qualified = add_server_prefix_to_name(original_name, known_prefix)
self.tool_name_to_mcp_server_name_mapping[qualified] = prefix
self._register_tool_route(qualified, server.server_id)
verbose_logger.info(f"Successfully fetched {len(prefixed_tools)} tools from server {server.name}")
return prefixed_tools
@ -4546,8 +4562,8 @@ class MCPServerManager:
if resolved_by_server_name_only:
tool_known = (
name in self.tool_name_to_mcp_server_name_mapping
or prefixed_tool_name in self.tool_name_to_mcp_server_name_mapping
name in self.tool_name_to_mcp_server_ids_mapping
or prefixed_tool_name in self.tool_name_to_mcp_server_ids_mapping
)
if not tool_known:
raise ValueError(f"Tool {name} not found")
@ -4658,6 +4674,7 @@ class MCPServerManager:
oauth2_headers: Optional[dict[str, str]] = None,
raw_headers: Optional[dict[str, str]] = None,
host_progress_callback: Optional[Callable] = None,
resolved_server: MCPServer | None = None,
) -> CallToolResult:
"""
Call a tool with the given name and arguments
@ -4670,13 +4687,21 @@ class MCPServerManager:
mcp_auth_header: MCP auth header (deprecated)
mcp_server_auth_headers: Optional dict of server-specific auth headers {server_alias: auth_value}
proxy_logging_obj: Optional ProxyLogging object for hook integration
resolved_server: Server the caller already resolved. When supplied it is
dispatched to verbatim, so identity never round-trips through the
non-unique ``server_name``. Callers without one fall back to
resolving by name.
Returns:
CallToolResult from the MCP server
"""
start_time = datetime.datetime.now()
mcp_server = self._resolve_mcp_server_for_tool_call(server_name, name)
mcp_server = (
resolved_server
if resolved_server is not None
else self._resolve_mcp_server_for_tool_call(server_name, name)
)
# Resolved before any hook runs so a missing BYOK credential (401) never
# leaves during-hook side effects (audit logging, rate-limit bookkeeping)
@ -4776,19 +4801,19 @@ class MCPServerManager:
# End of Methods that call the upstream MCP servers
#########################################################
def initialize_tool_name_to_mcp_server_name_mapping(self):
def initialize_tool_name_to_mcp_server_ids_mapping(self):
"""
On startup, initialize the tool name to MCP server name mapping
"""
try:
if asyncio.get_running_loop():
asyncio.create_task(self._initialize_tool_name_to_mcp_server_name_mapping())
asyncio.create_task(self._initialize_tool_name_to_mcp_server_ids_mapping())
except RuntimeError as e: # no running event loop
verbose_logger.exception(
f"No running event loop - skipping tool name to MCP server name mapping initialization: {str(e)}"
)
async def _initialize_tool_name_to_mcp_server_name_mapping(self):
async def _initialize_tool_name_to_mcp_server_ids_mapping(self):
"""
Call list_tools for each server and update the tool name to MCP server name mapping
Note: This now handles prefixed tool names
@ -4816,8 +4841,47 @@ class MCPServerManager:
# The tool.name here is already prefixed from _get_tools_from_server
# Extract original name for mapping
original_name, _ = split_server_prefix_from_name(tool.name)
self.tool_name_to_mcp_server_name_mapping[original_name] = server.name
self.tool_name_to_mcp_server_name_mapping[tool.name] = server.name
self._register_tool_route(original_name, server.server_id)
self._register_tool_route(tool.name, server.server_id)
def _register_tool_route(self, tool_name: str, server_id: str) -> None:
"""Record that ``server_id`` serves ``tool_name``, preserving any other owners.
Routing rows accumulate rather than overwrite so that a tool name served by
more than one server stays resolvable as ambiguous instead of silently
collapsing to whichever server registered last.
"""
owners = self.tool_name_to_mcp_server_ids_mapping.get(tool_name, frozenset())
self.tool_name_to_mcp_server_ids_mapping[tool_name] = owners | {server_id}
def resolve_tool_route(
self,
tool_name: str,
allowed_server_ids: frozenset[str] | None = None,
) -> "MCPToolRoute":
"""Resolve ``tool_name`` to the single MCP server that serves it.
Ambiguity is judged against ``allowed_server_ids`` when given, because an
unprefixed name is only ambiguous relative to the servers the caller can
actually reach; a session scoped to one server (``x-mcp-servers``) is served
unprefixed names precisely because nothing else is in scope for it.
Returns ``MCPToolRouteAmbiguous`` when several reachable servers expose the
name, so callers fail loudly instead of dispatching to an arbitrary one.
"""
owner_ids = self.tool_name_to_mcp_server_ids_mapping.get(tool_name)
if owner_ids:
in_scope = owner_ids if allowed_server_ids is None else owner_ids & allowed_server_ids
if len(in_scope) > 1:
return MCPToolRouteAmbiguous(tool_name=tool_name, server_ids=in_scope)
if len(in_scope) == 1:
scoped_server = self.get_mcp_server_by_id(next(iter(in_scope)))
if scoped_server is not None:
return MCPToolRouteResolved(server=scoped_server)
server = self._get_mcp_server_from_tool_name(tool_name)
if server is None:
return MCPToolRouteNotFound(tool_name=tool_name)
return MCPToolRouteResolved(server=server)
def _get_mcp_server_from_tool_name(self, tool_name: str) -> Optional[MCPServer]:
"""
@ -4827,7 +4891,9 @@ class MCPServerManager:
tool_name: Tool name (can be prefixed or non-prefixed)
Returns:
MCPServer if found, None otherwise
MCPServer if found, None otherwise. None is also returned when several
servers expose ``tool_name``; use ``resolve_tool_route`` to tell the two
cases apart.
"""
registry_servers = list(self.get_registry().values())
@ -4841,14 +4907,11 @@ class MCPServerManager:
prefix_to_server.setdefault(normalised, server)
# First try with the original tool name
if tool_name in self.tool_name_to_mcp_server_name_mapping:
server_name = self.tool_name_to_mcp_server_name_mapping[tool_name]
normalised_lookup = normalize_server_name(server_name)
if normalised_lookup in prefix_to_server:
return prefix_to_server[normalised_lookup]
for server in registry_servers:
if normalize_server_name(server.name) == normalised_lookup:
return server
owner_ids = self.tool_name_to_mcp_server_ids_mapping.get(tool_name)
if owner_ids is not None and len(owner_ids) == 1:
owned = self.get_mcp_server_by_id(next(iter(owner_ids)))
if owned is not None:
return owned
# If not found and tool name is prefixed, extract the prefix and
# match against any known form.
@ -4860,8 +4923,8 @@ class MCPServerManager:
normalised_prefix = normalize_server_name(server_name_from_prefix)
matched_server = prefix_to_server.get(normalised_prefix)
if matched_server is not None and (
original_tool_name in self.tool_name_to_mcp_server_name_mapping
or tool_name in self.tool_name_to_mcp_server_name_mapping
original_tool_name in self.tool_name_to_mcp_server_ids_mapping
or tool_name in self.tool_name_to_mcp_server_ids_mapping
):
return matched_server
@ -4964,7 +5027,7 @@ class MCPServerManager:
self.registry = registered_registry
if registered_openapi_tools:
self.initialize_tool_name_to_mcp_server_name_mapping()
self.initialize_tool_name_to_mcp_server_ids_mapping()
verbose_logger.debug("MCP registry refreshed (%s servers in registry)", len(registered_registry))

View file

@ -358,6 +358,8 @@ if MCP_AVAILABLE:
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
MCPServerManager,
MCPToolRouteAmbiguous,
MCPToolRouteResolved,
_caller_authorization_fans_out,
_client_forwarded_authorization_headers,
_should_strip_caller_authorization,
@ -2622,7 +2624,33 @@ if MCP_AVAILABLE:
original_tool_name = name
else:
# Resolve from tool name (MCP JSON-RPC or prefixed REST tool names).
mcp_server = global_mcp_server_manager._get_mcp_server_from_tool_name(name)
if requested_server is None:
route = global_mcp_server_manager.resolve_tool_route(
name,
allowed_server_ids=frozenset(server.server_id for server in allowed_mcp_servers),
)
if isinstance(route, MCPToolRouteAmbiguous):
candidates = sorted(
server.name
for server in (
global_mcp_server_manager.get_mcp_server_by_id(server_id) for server_id in route.server_ids
)
if server is not None
)
raise HTTPException(
status_code=409,
detail={
"error": "ambiguous_tool_name",
"message": (
f"Tool '{name}' is served by more than one MCP server "
f"({', '.join(candidates)}). Call it by its server-prefixed name to "
f"select one."
),
},
)
mcp_server = route.server if isinstance(route, MCPToolRouteResolved) else None
else:
mcp_server = global_mcp_server_manager._get_mcp_server_from_tool_name(name)
if mcp_server is None and requested_server is not None:
for known_prefix in iter_known_server_prefixes(requested_server):
candidate = global_mcp_server_manager._get_mcp_server_from_tool_name(
@ -2820,6 +2848,7 @@ if MCP_AVAILABLE:
raw_headers=raw_headers,
litellm_logging_obj=litellm_logging_obj,
host_progress_callback=host_progress_callback,
resolved_server=mcp_server,
)
# Fall back to local tool registry with original name (legacy support)
@ -3132,6 +3161,7 @@ if MCP_AVAILABLE:
raw_headers: Optional[Dict[str, str]] = None,
litellm_logging_obj: Optional[Any] = None,
host_progress_callback: Optional[Callable] = None,
resolved_server: MCPServer | None = None,
) -> CallToolResult:
"""Handle tool execution for managed server tools"""
# Import here to avoid circular import
@ -3148,6 +3178,7 @@ if MCP_AVAILABLE:
raw_headers=raw_headers,
proxy_logging_obj=proxy_logging_obj,
host_progress_callback=host_progress_callback,
resolved_server=resolved_server,
)
verbose_logger.debug("CALL TOOL RESULT: %s", call_tool_result)
return call_tool_result

View file

@ -467,17 +467,20 @@ class PanwPrismaAirsHandler(CustomGuardrail):
)
server_id = request_data.get("server_id")
if server_id:
server = global_mcp_server_manager.get_mcp_server_by_id(server_id)
if server:
return (
getattr(server, "alias", None)
or getattr(server, "server_name", None)
or getattr(server, "name", None)
or getattr(server, "server_id", None)
or "unknown"
)
return global_mcp_server_manager.tool_name_to_mcp_server_name_mapping.get(mcp_tool_name, "unknown")
server = (
global_mcp_server_manager.get_mcp_server_by_id(server_id)
if server_id
else global_mcp_server_manager._get_mcp_server_from_tool_name(mcp_tool_name)
)
if server:
return (
getattr(server, "alias", None)
or getattr(server, "server_name", None)
or getattr(server, "name", None)
or getattr(server, "server_id", None)
or "unknown"
)
return "unknown"
except ImportError:
return "unknown"
except Exception:

View file

@ -779,6 +779,7 @@ class LiteLLM_Proxy_MCP_Handler:
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
proxy_logging_obj=proxy_logging_obj,
resolved_server=mcp_server,
)
if litellm_logging_obj:

View file

@ -104,7 +104,7 @@ async def test_mcp_cost_tracking():
litellm.callbacks = [test_logger]
# Initialize the tool mapping
await local_mcp_server_manager._initialize_tool_name_to_mcp_server_name_mapping()
await local_mcp_server_manager._initialize_tool_name_to_mcp_server_ids_mapping()
# Patch the global manager in both modules where it's used
with (
@ -121,17 +121,18 @@ async def test_mcp_cost_tracking():
_set_authorized_user(local_mcp_server_manager.get_all_mcp_server_ids())
print(
"tool_name_to_mcp_server_name_mapping",
local_mcp_server_manager.tool_name_to_mcp_server_name_mapping,
"tool_name_to_mcp_server_ids_mapping",
local_mcp_server_manager.tool_name_to_mcp_server_ids_mapping,
)
# Manually add the tool mapping to ensure it's available (since mocking might not capture it properly)
local_mcp_server_manager.tool_name_to_mcp_server_name_mapping[
zapier_server_ids = frozenset(local_mcp_server_manager.get_all_mcp_server_ids())
local_mcp_server_manager.tool_name_to_mcp_server_ids_mapping[
"add_tools"
] = "zapier_gmail_server"
local_mcp_server_manager.tool_name_to_mcp_server_name_mapping[
] = zapier_server_ids
local_mcp_server_manager.tool_name_to_mcp_server_ids_mapping[
"zapier_gmail_server-add_tools"
] = "zapier_gmail_server"
] = zapier_server_ids
# Call mcp tool
response = await mcp_server_tool_call(
@ -237,21 +238,22 @@ async def test_mcp_cost_tracking_per_tool():
litellm.callbacks = [test_logger]
# Initialize the tool mapping
await local_mcp_server_manager._initialize_tool_name_to_mcp_server_name_mapping()
await local_mcp_server_manager._initialize_tool_name_to_mcp_server_ids_mapping()
# Manually add the tool mapping to ensure it's available (since mocking might not capture it properly)
local_mcp_server_manager.tool_name_to_mcp_server_name_mapping[
test_server_ids = frozenset(local_mcp_server_manager.get_all_mcp_server_ids())
local_mcp_server_manager.tool_name_to_mcp_server_ids_mapping[
"expensive_tool"
] = "test_server"
local_mcp_server_manager.tool_name_to_mcp_server_name_mapping[
] = test_server_ids
local_mcp_server_manager.tool_name_to_mcp_server_ids_mapping[
"test_server-expensive_tool"
] = "test_server"
local_mcp_server_manager.tool_name_to_mcp_server_name_mapping["cheap_tool"] = (
"test_server"
] = test_server_ids
local_mcp_server_manager.tool_name_to_mcp_server_ids_mapping["cheap_tool"] = (
test_server_ids
)
local_mcp_server_manager.tool_name_to_mcp_server_name_mapping[
local_mcp_server_manager.tool_name_to_mcp_server_ids_mapping[
"test_server-cheap_tool"
] = "test_server"
] = test_server_ids
# Patch the global manager in both modules where it's used
with (
@ -268,8 +270,8 @@ async def test_mcp_cost_tracking_per_tool():
_set_authorized_user(local_mcp_server_manager.get_all_mcp_server_ids())
print(
"tool_name_to_mcp_server_name_mapping",
local_mcp_server_manager.tool_name_to_mcp_server_name_mapping,
"tool_name_to_mcp_server_ids_mapping",
local_mcp_server_manager.tool_name_to_mcp_server_ids_mapping,
)
# Test 1: Call expensive_tool - should cost 5.0
@ -401,15 +403,16 @@ async def test_mcp_tool_call_hook():
litellm.callbacks = [test_logger]
# Initialize the tool mapping
await local_mcp_server_manager._initialize_tool_name_to_mcp_server_name_mapping()
await local_mcp_server_manager._initialize_tool_name_to_mcp_server_ids_mapping()
# Manually add the tool mapping to ensure it's available (since mocking might not capture it properly)
local_mcp_server_manager.tool_name_to_mcp_server_name_mapping["add_tools"] = (
"zapier_gmail_server"
zapier_server_ids = frozenset(local_mcp_server_manager.get_all_mcp_server_ids())
local_mcp_server_manager.tool_name_to_mcp_server_ids_mapping["add_tools"] = (
zapier_server_ids
)
local_mcp_server_manager.tool_name_to_mcp_server_name_mapping[
local_mcp_server_manager.tool_name_to_mcp_server_ids_mapping[
"zapier_gmail_server-add_tools"
] = "zapier_gmail_server"
] = zapier_server_ids
# Patch the global manager in both modules where it's used
with (
@ -426,8 +429,8 @@ async def test_mcp_tool_call_hook():
_set_authorized_user(local_mcp_server_manager.get_all_mcp_server_ids())
print(
"tool_name_to_mcp_server_name_mapping",
local_mcp_server_manager.tool_name_to_mcp_server_name_mapping,
"tool_name_to_mcp_server_ids_mapping",
local_mcp_server_manager.tool_name_to_mcp_server_ids_mapping,
)
# Call mcp tool using the correct separator format (- not /)

View file

@ -105,12 +105,18 @@ async def test_mcp_server_manager_https_server():
assert tools[0].name == f"{expected_prefix}-gmail_send_email"
# Manually set up the tool mapping for the call_tool test
mcp_server_manager.tool_name_to_mcp_server_name_mapping["gmail_send_email"] = (
expected_prefix
zapier_server_ids = frozenset(
server_id
for server_id, server in mcp_server_manager.get_registry().items()
if server.server_name == "zapier_mcp_server"
)
mcp_server_manager.tool_name_to_mcp_server_name_mapping[
assert len(zapier_server_ids) == 1, "Expected exactly one configured zapier server"
mcp_server_manager.tool_name_to_mcp_server_ids_mapping["gmail_send_email"] = (
zapier_server_ids
)
mcp_server_manager.tool_name_to_mcp_server_ids_mapping[
f"{expected_prefix}-gmail_send_email"
] = expected_prefix
] = zapier_server_ids
result = await mcp_server_manager.call_tool(
server_name="zapier_mcp_server",
@ -221,16 +227,16 @@ async def test_mcp_http_transport_list_tools_mock():
# Verify tool mapping was updated
expected_prefix = "test_http_server"
assert (
test_manager.tool_name_to_mcp_server_name_mapping[
test_manager.tool_name_to_mcp_server_ids_mapping[
f"{expected_prefix}-gmail_send_email"
]
== expected_prefix
== frozenset(allowed_server_ids)
)
assert (
test_manager.tool_name_to_mcp_server_name_mapping[
test_manager.tool_name_to_mcp_server_ids_mapping[
f"{expected_prefix}-calendar_create_event"
]
== expected_prefix
== frozenset(allowed_server_ids)
)
@ -275,8 +281,8 @@ async def test_mcp_http_transport_call_tool_mock():
)
# Manually set up tool mapping (normally done by list_tools)
test_manager.tool_name_to_mcp_server_name_mapping["gmail_send_email"] = (
"test_http_server"
test_manager.tool_name_to_mcp_server_ids_mapping["gmail_send_email"] = frozenset(
test_manager.get_registry().keys()
)
# Call the tool
@ -341,8 +347,8 @@ async def test_mcp_http_transport_call_tool_error_mock():
)
# Manually set up tool mapping
test_manager.tool_name_to_mcp_server_name_mapping["gmail_send_email"] = (
"test_http_server"
test_manager.tool_name_to_mcp_server_ids_mapping["gmail_send_email"] = frozenset(
test_manager.get_registry().keys()
)
# Call the tool with invalid data
@ -383,8 +389,8 @@ async def test_mcp_http_transport_tool_not_found():
)
# Mapping populated for this server but not for the requested tool
test_manager.tool_name_to_mcp_server_name_mapping["gmail_send_email"] = (
"test_http_server"
test_manager.tool_name_to_mcp_server_ids_mapping["gmail_send_email"] = frozenset(
test_manager.get_registry().keys()
)
# Try to call a tool that doesn't exist in mapping
@ -1395,12 +1401,12 @@ async def test_mcp_server_manager_alias_tool_prefixing():
# Verify mapping is updated correctly
assert (
test_manager.tool_name_to_mcp_server_name_mapping["send_email"]
== "my_alias"
test_manager.tool_name_to_mcp_server_ids_mapping["send_email"]
== frozenset({"test-server-123"})
)
assert (
test_manager.tool_name_to_mcp_server_name_mapping["my_alias-send_email"]
== "my_alias"
test_manager.tool_name_to_mcp_server_ids_mapping["my_alias-send_email"]
== frozenset({"test-server-123"})
)
@ -1455,12 +1461,12 @@ async def test_mcp_server_manager_server_name_tool_prefixing():
# Verify mapping is updated correctly
assert (
test_manager.tool_name_to_mcp_server_name_mapping["send_email"]
== "Test Server"
test_manager.tool_name_to_mcp_server_ids_mapping["send_email"]
== frozenset({"test-server-123"})
)
assert (
test_manager.tool_name_to_mcp_server_name_mapping["Test_Server-send_email"]
== "Test Server"
test_manager.tool_name_to_mcp_server_ids_mapping["Test_Server-send_email"]
== frozenset({"test-server-123"})
)
@ -1515,14 +1521,14 @@ async def test_mcp_server_manager_server_id_tool_prefixing():
# Verify mapping is updated correctly
assert (
test_manager.tool_name_to_mcp_server_name_mapping["send_email"]
== "test-server-123"
test_manager.tool_name_to_mcp_server_ids_mapping["send_email"]
== frozenset({"test-server-123"})
)
assert (
test_manager.tool_name_to_mcp_server_name_mapping[
test_manager.tool_name_to_mcp_server_ids_mapping[
"test-server-123-send_email"
]
== "test-server-123"
== frozenset({"test-server-123"})
)

View file

@ -39,11 +39,11 @@ def cleanup_mcp_global_state():
# Clear before test
global_mcp_server_manager.registry.clear()
global_mcp_server_manager.tool_name_to_mcp_server_name_mapping.clear()
global_mcp_server_manager.tool_name_to_mcp_server_ids_mapping.clear()
yield
# Clear after test
global_mcp_server_manager.registry.clear()
global_mcp_server_manager.tool_name_to_mcp_server_name_mapping.clear()
global_mcp_server_manager.tool_name_to_mcp_server_ids_mapping.clear()
except ImportError:
# MCP not available, skip cleanup
yield
@ -5957,8 +5957,8 @@ async def test_execute_mcp_tool_rest_server_id_authoritative_for_unprefixed_tool
with (
patch.dict(
mcp_module.global_mcp_server_manager.tool_name_to_mcp_server_name_mapping,
{"echo": oauth_server.name},
mcp_module.global_mcp_server_manager.tool_name_to_mcp_server_ids_mapping,
{"echo": frozenset({oauth_server.server_id})},
),
patch.object(
mcp_module.global_mcp_server_manager,
@ -6044,8 +6044,8 @@ async def test_execute_mcp_tool_rest_server_id_injects_requested_server_credenti
with (
patch.dict(
mcp_module.global_mcp_server_manager.tool_name_to_mcp_server_name_mapping,
{"echo": collision_server.name},
mcp_module.global_mcp_server_manager.tool_name_to_mcp_server_ids_mapping,
{"echo": frozenset({collision_server.server_id})},
),
patch.object(
mcp_module.global_mcp_server_manager,
@ -6087,6 +6087,309 @@ async def test_execute_mcp_tool_rest_server_id_injects_requested_server_credenti
assert routed.authentication_token != collision_server.authentication_token
@pytest.mark.asyncio
async def test_execute_mcp_tool_jsonrpc_unprefixed_ambiguous_tool_is_rejected():
"""MCP JSON-RPC carries no server_id, so an unprefixed name owned by two servers must 409.
The REST path disambiguates with requested_server_id, but the MCP protocol has no such
field. Without this the call silently dispatched to whichever server registered the tool
name last, sending that server's upstream credentials.
"""
from litellm.proxy._experimental.mcp_server import server as mcp_module
alpha = MCPServer(
server_id="alpha-server-id",
name="echo_alpha",
server_name="echo_alpha",
url="http://127.0.0.1:5115/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.api_key,
authentication_token="alpha-secret",
)
zulu = MCPServer(
server_id="zulu-server-id",
name="echo_zulu",
server_name="echo_zulu",
url="http://127.0.0.1:5115/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.api_key,
authentication_token="zulu-secret",
)
injected: dict = {}
async def fake_create_mcp_client(server, **kwargs):
injected["server"] = server
raise AssertionError("upstream must not be called for an ambiguous tool name")
with (
patch.dict(
mcp_module.global_mcp_server_manager.tool_name_to_mcp_server_ids_mapping,
{"echo": frozenset({alpha.server_id, zulu.server_id})},
),
patch.object(
mcp_module.global_mcp_server_manager,
"get_registry",
return_value={alpha.server_id: alpha, zulu.server_id: zulu},
),
patch.object(
mcp_module.global_mcp_server_manager,
"_create_mcp_client",
new=fake_create_mcp_client,
),
patch.object(mcp_module.MCPRequestHandler, "is_tool_allowed", return_value=True),
patch.object(mcp_module.global_mcp_tool_registry, "get_tool", return_value=None),
patch("litellm.proxy.proxy_server.proxy_logging_obj", None),
pytest.raises(HTTPException) as exc_info,
):
await mcp_module.execute_mcp_tool(
name="echo",
arguments={"message": "hello"},
allowed_mcp_servers=[alpha, zulu],
start_time=datetime.now(),
)
assert exc_info.value.status_code == 409
assert exc_info.value.detail["error"] == "ambiguous_tool_name"
assert "echo_alpha" in exc_info.value.detail["message"]
assert "echo_zulu" in exc_info.value.detail["message"]
assert "server" not in injected
@pytest.mark.asyncio
async def test_execute_mcp_tool_jsonrpc_unprefixed_resolves_when_caller_reaches_one_server():
"""A caller scoped to one server is served unprefixed names and must still route.
Ambiguity is relative to what the caller can reach. Judging it against the whole
registry would reject an x-mcp-servers scoped session, whose unprefixed tool names
only ever had one candidate.
"""
from mcp.types import TextContent
from litellm.proxy._experimental.mcp_server import server as mcp_module
alpha = MCPServer(
server_id="alpha-server-id",
name="echo_alpha",
server_name="echo_alpha",
url="http://127.0.0.1:5115/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.api_key,
authentication_token="alpha-secret",
)
zulu = MCPServer(
server_id="zulu-server-id",
name="echo_zulu",
server_name="echo_zulu",
url="http://127.0.0.1:5115/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.api_key,
authentication_token="zulu-secret",
)
fake_client = MagicMock()
fake_client._last_initialize_instructions = None
fake_client.call_tool = AsyncMock(
return_value=mcp_module.CallToolResult(
content=[TextContent(type="text", text="ok")],
isError=False,
)
)
injected: dict = {}
async def fake_create_mcp_client(server, **kwargs):
injected["server"] = server
return fake_client
with (
patch.dict(
mcp_module.global_mcp_server_manager.tool_name_to_mcp_server_ids_mapping,
{"echo": frozenset({alpha.server_id, zulu.server_id})},
),
patch.object(
mcp_module.global_mcp_server_manager,
"get_registry",
return_value={alpha.server_id: alpha, zulu.server_id: zulu},
),
patch.object(
mcp_module.global_mcp_server_manager,
"_create_mcp_client",
new=fake_create_mcp_client,
),
patch.object(mcp_module.MCPRequestHandler, "is_tool_allowed", return_value=True),
patch.object(mcp_module.global_mcp_tool_registry, "get_tool", return_value=None),
patch("litellm.proxy.proxy_server.proxy_logging_obj", None),
):
await mcp_module.execute_mcp_tool(
name="echo",
arguments={"message": "hello"},
allowed_mcp_servers=[zulu],
start_time=datetime.now(),
)
routed = injected["server"]
assert routed.server_id == zulu.server_id
assert routed.authentication_token == "zulu-secret"
@pytest.mark.asyncio
async def test_execute_mcp_tool_dispatches_resolved_server_not_a_name_lookup():
"""The server execute_mcp_tool resolved must be the server dispatched to.
server_name is not unique, so re-deriving the server from it inside call_tool picks the
first registry entry with that name, which defeats server_id resolution when two servers
share a name.
"""
from mcp.types import TextContent
from litellm.proxy._experimental.mcp_server import server as mcp_module
first_named_shared = MCPServer(
server_id="first-id",
name="shared",
server_name="shared",
url="http://127.0.0.1:5115/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.api_key,
authentication_token="first-secret",
)
second_named_shared = MCPServer(
server_id="second-id",
name="shared",
server_name="shared",
url="http://127.0.0.1:5115/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.api_key,
authentication_token="second-secret",
)
fake_client = MagicMock()
fake_client._last_initialize_instructions = None
fake_client.call_tool = AsyncMock(
return_value=mcp_module.CallToolResult(
content=[TextContent(type="text", text="ok")],
isError=False,
)
)
injected: dict = {}
async def fake_create_mcp_client(server, **kwargs):
injected["server"] = server
return fake_client
with (
patch.dict(
mcp_module.global_mcp_server_manager.tool_name_to_mcp_server_ids_mapping,
{"echo": frozenset({second_named_shared.server_id})},
),
patch.object(
mcp_module.global_mcp_server_manager,
"get_registry",
return_value={
first_named_shared.server_id: first_named_shared,
second_named_shared.server_id: second_named_shared,
},
),
patch.object(
mcp_module.global_mcp_server_manager,
"_create_mcp_client",
new=fake_create_mcp_client,
),
patch.object(mcp_module.MCPRequestHandler, "is_tool_allowed", return_value=True),
patch.object(mcp_module.global_mcp_tool_registry, "get_tool", return_value=None),
patch("litellm.proxy.proxy_server.proxy_logging_obj", None),
):
await mcp_module.execute_mcp_tool(
name="echo",
arguments={"message": "hello"},
allowed_mcp_servers=[first_named_shared, second_named_shared],
start_time=datetime.now(),
requested_server_id=second_named_shared.server_id,
)
routed = injected["server"]
assert routed.server_id == second_named_shared.server_id
assert routed.authentication_token == "second-secret"
@pytest.mark.asyncio
async def test_execute_mcp_tool_jsonrpc_prefixed_tool_routes_to_its_own_server():
"""The prefixed name stays unambiguous and must still reach exactly its own server."""
from mcp.types import TextContent
from litellm.proxy._experimental.mcp_server import server as mcp_module
alpha = MCPServer(
server_id="alpha-server-id",
name="echo_alpha",
server_name="echo_alpha",
url="http://127.0.0.1:5115/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.api_key,
authentication_token="alpha-secret",
)
zulu = MCPServer(
server_id="zulu-server-id",
name="echo_zulu",
server_name="echo_zulu",
url="http://127.0.0.1:5115/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.api_key,
authentication_token="zulu-secret",
)
fake_client = MagicMock()
fake_client._last_initialize_instructions = None
fake_client.call_tool = AsyncMock(
return_value=mcp_module.CallToolResult(
content=[TextContent(type="text", text="ok")],
isError=False,
)
)
injected: dict = {}
async def fake_create_mcp_client(server, **kwargs):
injected["server"] = server
return fake_client
with (
patch.dict(
mcp_module.global_mcp_server_manager.tool_name_to_mcp_server_ids_mapping,
{
"echo": frozenset({alpha.server_id, zulu.server_id}),
"echo_alpha-echo": frozenset({alpha.server_id}),
"echo_zulu-echo": frozenset({zulu.server_id}),
},
),
patch.object(
mcp_module.global_mcp_server_manager,
"get_registry",
return_value={alpha.server_id: alpha, zulu.server_id: zulu},
),
patch.object(
mcp_module.global_mcp_server_manager,
"_create_mcp_client",
new=fake_create_mcp_client,
),
patch.object(mcp_module.MCPRequestHandler, "is_tool_allowed", return_value=True),
patch.object(mcp_module.global_mcp_tool_registry, "get_tool", return_value=None),
patch("litellm.proxy.proxy_server.proxy_logging_obj", None),
):
await mcp_module.execute_mcp_tool(
name="echo_zulu-echo",
arguments={"message": "hello"},
allowed_mcp_servers=[alpha, zulu],
start_time=datetime.now(),
)
routed = injected["server"]
assert routed.server_id == zulu.server_id
assert routed.authentication_token == "zulu-secret"
@pytest.mark.asyncio
async def test_execute_mcp_tool_rest_prefixed_tool_still_validates_server_id():
"""Prefixed REST tool names must still match the requested server_id."""

View file

@ -3005,7 +3005,7 @@ class TestMCPServerManager:
assert minimal_spec.config.profile == "rfc8693"
@pytest.mark.asyncio
async def test_config_oauth_initialize_tool_name_to_mcp_server_name_mapping(self):
async def test_config_oauth_initialize_tool_name_to_mcp_server_ids_mapping(self):
manager = MCPServerManager()
config = {
@ -3022,8 +3022,8 @@ class TestMCPServerManager:
await manager.load_servers_from_config(config)
# Initialize the tool mapping
await manager._initialize_tool_name_to_mcp_server_name_mapping()
assert manager.tool_name_to_mcp_server_name_mapping == {}
await manager._initialize_tool_name_to_mcp_server_ids_mapping()
assert manager.tool_name_to_mcp_server_ids_mapping == {}
@pytest.mark.asyncio
async def test_list_tools_handles_missing_server_alias(self):
@ -3869,8 +3869,8 @@ class TestMCPServerManager:
transport=MCPTransport.http,
)
manager.registry = {"jira": server}
manager.tool_name_to_mcp_server_name_mapping["jira-search_issues"] = "jira"
manager.tool_name_to_mcp_server_name_mapping["search_issues"] = "jira"
manager.tool_name_to_mcp_server_ids_mapping["jira-search_issues"] = frozenset({"jira"})
manager.tool_name_to_mcp_server_ids_mapping["search_issues"] = frozenset({"jira"})
resolved = manager._resolve_mcp_server_for_tool_call("jira", "search_issues")
assert resolved is server
@ -3885,11 +3885,113 @@ class TestMCPServerManager:
transport=MCPTransport.http,
)
manager.registry = {"srv-uuid-123": server}
manager.tool_name_to_mcp_server_name_mapping["create_zap"] = "zapier"
manager.tool_name_to_mcp_server_ids_mapping["create_zap"] = frozenset({"srv-uuid-123"})
resolved = manager._resolve_mcp_server_for_tool_call("zapier-alias", "create_zap")
assert resolved is server
def test_register_tool_route_accumulates_owners_across_servers(self):
"""Two servers exposing one tool name are both recorded, not last-writer-wins."""
manager = MCPServerManager()
manager._register_tool_route("echo", "id-alpha")
manager._register_tool_route("echo", "id-zulu")
assert manager.tool_name_to_mcp_server_ids_mapping["echo"] == frozenset({"id-alpha", "id-zulu"})
def test_register_tool_route_is_idempotent_for_one_server(self):
"""Re-listing the same server must not make its own tool look ambiguous."""
manager = MCPServerManager()
manager._register_tool_route("echo", "id-alpha")
manager._register_tool_route("echo", "id-alpha")
assert manager.tool_name_to_mcp_server_ids_mapping["echo"] == frozenset({"id-alpha"})
def test_get_mcp_server_from_tool_name_refuses_ambiguous_unprefixed_name(self):
"""An unprefixed name owned by several servers resolves to nothing, never to one of them."""
manager = MCPServerManager()
alpha = MCPServer(server_id="id-alpha", name="echo_alpha", transport=MCPTransport.http)
zulu = MCPServer(server_id="id-zulu", name="echo_zulu", transport=MCPTransport.http)
manager.registry = {"id-alpha": alpha, "id-zulu": zulu}
manager._register_tool_route("echo", "id-alpha")
manager._register_tool_route("echo", "id-zulu")
assert manager._get_mcp_server_from_tool_name("echo") is None
def test_get_mcp_server_from_tool_name_resolves_sole_owner_by_id(self):
"""A tool name owned by exactly one server still resolves, addressed by server_id."""
manager = MCPServerManager()
alpha = MCPServer(server_id="id-alpha", name="echo_alpha", transport=MCPTransport.http)
manager.registry = {"id-alpha": alpha}
manager._register_tool_route("echo", "id-alpha")
assert manager._get_mcp_server_from_tool_name("echo") is alpha
def test_get_mcp_server_from_tool_name_refuses_prefixed_name_of_duplicate_named_servers(self):
"""server_name is not unique, so a prefix shared by two servers must not silently pick one."""
manager = MCPServerManager()
first = MCPServer(server_id="id-first", name="shared", transport=MCPTransport.http)
second = MCPServer(server_id="id-second", name="shared", transport=MCPTransport.http)
manager.registry = {"id-first": first, "id-second": second}
manager._register_tool_route("shared-echo", "id-first")
manager._register_tool_route("shared-echo", "id-second")
assert manager._get_mcp_server_from_tool_name("shared-echo") is None
def test_resolve_tool_route_names_every_ambiguous_owner(self):
"""The ambiguous route carries all owners so callers can report the real candidates."""
manager = MCPServerManager()
alpha = MCPServer(server_id="id-alpha", name="echo_alpha", transport=MCPTransport.http)
zulu = MCPServer(server_id="id-zulu", name="echo_zulu", transport=MCPTransport.http)
manager.registry = {"id-alpha": alpha, "id-zulu": zulu}
manager._register_tool_route("echo", "id-alpha")
manager._register_tool_route("echo", "id-zulu")
route = manager.resolve_tool_route("echo")
assert type(route).__name__ == "MCPToolRouteAmbiguous"
assert route.server_ids == frozenset({"id-alpha", "id-zulu"})
def test_resolve_tool_route_resolves_sole_owner(self):
manager = MCPServerManager()
alpha = MCPServer(server_id="id-alpha", name="echo_alpha", transport=MCPTransport.http)
manager.registry = {"id-alpha": alpha}
manager._register_tool_route("echo", "id-alpha")
route = manager.resolve_tool_route("echo")
assert type(route).__name__ == "MCPToolRouteResolved"
assert route.server is alpha
def test_resolve_tool_route_reports_unknown_tool(self):
manager = MCPServerManager()
assert type(manager.resolve_tool_route("nothing_serves_this")).__name__ == "MCPToolRouteNotFound"
def test_cleanup_withdraws_only_departing_server_from_shared_tool_name(self):
"""Removing one server must leave a co-owned tool name routable to the survivor."""
manager = MCPServerManager()
alpha = MCPServer(server_id="id-alpha", name="echo_alpha", transport=MCPTransport.http)
zulu = MCPServer(server_id="id-zulu", name="echo_zulu", transport=MCPTransport.http)
manager.registry = {"id-alpha": alpha, "id-zulu": zulu}
manager._register_tool_route("echo", "id-alpha")
manager._register_tool_route("echo", "id-zulu")
manager._cleanup_server_tool_routing_artifacts(zulu)
del manager.registry["id-zulu"]
assert manager.tool_name_to_mcp_server_ids_mapping["echo"] == frozenset({"id-alpha"})
assert manager._get_mcp_server_from_tool_name("echo") is alpha
def test_cleanup_drops_the_route_when_its_last_owner_leaves(self):
manager = MCPServerManager()
alpha = MCPServer(server_id="id-alpha", name="echo_alpha", transport=MCPTransport.http)
manager.registry = {"id-alpha": alpha}
manager._register_tool_route("echo", "id-alpha")
manager._cleanup_server_tool_routing_artifacts(alpha)
assert "echo" not in manager.tool_name_to_mcp_server_ids_mapping
def test_resolve_mcp_server_for_tool_call_unknown_tool_with_empty_mapping(self):
"""Server-name match alone must not let unknown tools through when the
mapping has no entries for that server (e.g. listing has not completed
@ -3916,7 +4018,7 @@ class TestMCPServerManager:
transport=MCPTransport.http,
)
manager.registry = {"linear": server}
manager.tool_name_to_mcp_server_name_mapping["create_issue"] = "linear"
manager.tool_name_to_mcp_server_ids_mapping["create_issue"] = frozenset({"linear"})
# server_name is empty so the fallback unprefixed lookup runs and matches.
resolved = manager._resolve_mcp_server_for_tool_call("", "create_issue")
@ -3943,8 +4045,8 @@ class TestMCPServerManager:
)
manager.registry = {"github": server}
# Mapping has *some* tools for github but not "missing_tool".
manager.tool_name_to_mcp_server_name_mapping["github-list_repos"] = "github"
manager.tool_name_to_mcp_server_name_mapping["list_repos"] = "github"
manager.tool_name_to_mcp_server_ids_mapping["github-list_repos"] = frozenset({"github"})
manager.tool_name_to_mcp_server_ids_mapping["list_repos"] = frozenset({"github"})
with pytest.raises(ValueError, match="Tool missing_tool not found"):
manager._resolve_mcp_server_for_tool_call("github", "missing_tool")
@ -4254,10 +4356,10 @@ class TestMCPServerManager:
assert names == ["close_issue", "create_issue"]
# Mapping should include both original and prefixed names -> resolves calls either way
assert manager.tool_name_to_mcp_server_name_mapping["create_issue"] == "jira"
assert manager.tool_name_to_mcp_server_name_mapping["jira-create_issue"] == "jira"
assert manager.tool_name_to_mcp_server_name_mapping["close_issue"] == "jira"
assert manager.tool_name_to_mcp_server_name_mapping["jira-close_issue"] == "jira"
assert manager.tool_name_to_mcp_server_ids_mapping["create_issue"] == frozenset({"jira"})
assert manager.tool_name_to_mcp_server_ids_mapping["jira-create_issue"] == frozenset({"jira"})
assert manager.tool_name_to_mcp_server_ids_mapping["close_issue"] == frozenset({"jira"})
assert manager.tool_name_to_mcp_server_ids_mapping["jira-close_issue"] == frozenset({"jira"})
def test_get_mcp_server_from_tool_name_with_prefixed_and_unprefixed(self):
"""After mapping is populated, manager resolves both prefixed and unprefixed tool names to the same server."""
@ -4746,8 +4848,8 @@ class TestMCPServerManager:
# Register the server and map a tool to it
manager.registry = {"test-server": server}
manager.tool_name_to_mcp_server_name_mapping["test_tool"] = "test-server"
manager.tool_name_to_mcp_server_name_mapping["test-server-test_tool"] = "test-server"
manager.tool_name_to_mcp_server_ids_mapping["test_tool"] = frozenset({"test-server"})
manager.tool_name_to_mcp_server_ids_mapping["test-server-test_tool"] = frozenset({"test-server"})
# Create mock client that tracks call_tool usage
mock_client = AsyncMock()
@ -5198,8 +5300,8 @@ class TestMCPServerManager:
prefixed_tool_name = add_server_prefix_to_name(tool_name, "test_server")
# Populate the mapping with the original tool name
manager.tool_name_to_mcp_server_name_mapping[tool_name] = "test_server"
manager.tool_name_to_mcp_server_name_mapping[prefixed_tool_name] = "test_server"
manager.tool_name_to_mcp_server_ids_mapping[tool_name] = frozenset({server.server_id})
manager.tool_name_to_mcp_server_ids_mapping[prefixed_tool_name] = frozenset({server.server_id})
# Test: _get_mcp_server_from_tool_name should find the server using server.server_name
# even when server.name is different
@ -7250,15 +7352,15 @@ class TestApprovalStatusGate:
input_schema={"type": "object"},
handler=_noop_handler,
)
manager.tool_name_to_mcp_server_name_mapping["demo_tool"] = prefix
manager.tool_name_to_mcp_server_name_mapping[prefixed] = prefix
manager.tool_name_to_mcp_server_ids_mapping["demo_tool"] = frozenset({server.server_id})
manager.tool_name_to_mcp_server_ids_mapping[prefixed] = frozenset({server.server_id})
await manager.update_server(self._make_server("evict-openapi", MCPApprovalStatus.rejected))
assert "evict-openapi" not in manager.registry
assert prefixed not in global_mcp_tool_registry.tools
assert "demo_tool" not in manager.tool_name_to_mcp_server_name_mapping
assert prefixed not in manager.tool_name_to_mcp_server_name_mapping
assert "demo_tool" not in manager.tool_name_to_mcp_server_ids_mapping
assert prefixed not in manager.tool_name_to_mcp_server_ids_mapping
async def test_update_server_noop_for_unregistered_pending(self):
# update_server called with a pending row that was never registered