From 7fe0d9501e3a1a78e4d74825be54589d9921cf68 Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Thu, 16 Jul 2026 23:19:06 -0700 Subject: [PATCH] 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. --- .../mcp_server/mcp_server_manager.py | 161 ++++++--- .../proxy/_experimental/mcp_server/server.py | 33 +- .../panw_prisma_airs/panw_prisma_airs.py | 25 +- .../mcp/litellm_proxy_mcp_handler.py | 1 + tests/mcp_tests/test_mcp_logging.py | 53 +-- tests/mcp_tests/test_mcp_server.py | 58 ++-- .../mcp_server/test_mcp_server.py | 315 +++++++++++++++++- .../mcp_server/test_mcp_server_manager.py | 144 ++++++-- 8 files changed, 651 insertions(+), 139 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 115ff2e492c..b8bf986776e 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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)) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 68a61b85175..0705c392f7f 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -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 diff --git a/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py b/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py index 192c8c9bc77..e3317f23e1e 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py +++ b/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py @@ -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: diff --git a/litellm/responses/mcp/litellm_proxy_mcp_handler.py b/litellm/responses/mcp/litellm_proxy_mcp_handler.py index e03f0296109..460be2b0427 100644 --- a/litellm/responses/mcp/litellm_proxy_mcp_handler.py +++ b/litellm/responses/mcp/litellm_proxy_mcp_handler.py @@ -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: diff --git a/tests/mcp_tests/test_mcp_logging.py b/tests/mcp_tests/test_mcp_logging.py index 55b49aa0d29..d8193c3074f 100644 --- a/tests/mcp_tests/test_mcp_logging.py +++ b/tests/mcp_tests/test_mcp_logging.py @@ -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 /) diff --git a/tests/mcp_tests/test_mcp_server.py b/tests/mcp_tests/test_mcp_server.py index 515bf1233aa..e8307f31181 100644 --- a/tests/mcp_tests/test_mcp_server.py +++ b/tests/mcp_tests/test_mcp_server.py @@ -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"}) ) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index a983ac3ff48..6b878b11a74 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -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.""" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 80bf08a5eba..646c7a835da 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -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