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