From 2843dab7fe496a61970bc1f193d4962d17ea4c41 Mon Sep 17 00:00:00 2001 From: YutaSaito <36355491+uc4w6c@users.noreply.github.com> Date: Thu, 13 Nov 2025 06:50:52 +0900 Subject: [PATCH] fix: allow tool call even when server name prefix is missing (#16425) * fix: allow tool call even when server name prefix is missing * fix: test * fix: test * fix: test --- .../mcp_server/mcp_server_manager.py | 60 ++++---- .../proxy/_experimental/mcp_server/server.py | 105 +++++++------ litellm/responses/main.py | 25 +++- .../mcp/litellm_proxy_mcp_handler.py | 141 +++++++++++------- .../responses/mcp/mcp_streaming_iterator.py | 14 +- .../mcp_tests/test_aresponses_api_with_mcp.py | 14 +- tests/mcp_tests/test_mcp_server.py | 21 +-- .../mcp_server/test_mcp_server.py | 83 +++++++++-- .../mcp_server/test_mcp_server_manager.py | 26 ++-- 9 files changed, 308 insertions(+), 181 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 5b1dc5933c3..7cd6546e9cf 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -952,7 +952,7 @@ class MCPServerManager: self, name: str, arguments: Dict[str, Any], - server_name_from_prefix: str, + server_name: str, user_api_key_auth: Optional[UserAPIKeyAuth], proxy_logging_obj: ProxyLogging, server: MCPServer, @@ -983,7 +983,7 @@ class MCPServerManager: pre_hook_kwargs = { "name": name, "arguments": arguments, - "server_name": server_name_from_prefix, + "server_name": server_name, "user_api_key_auth": user_api_key_auth, "user_api_key_user_id": ( getattr(user_api_key_auth, "user_id", None) @@ -1197,6 +1197,7 @@ class MCPServerManager: async def call_tool( self, + server_name: str, name: str, arguments: Dict[str, Any], user_api_key_auth: Optional[UserAPIKeyAuth] = None, @@ -1207,10 +1208,11 @@ class MCPServerManager: raw_headers: Optional[Dict[str, str]] = None, ) -> CallToolResult: """ - Call a tool with the given name and arguments (handles prefixed tool names) + Call a tool with the given name and arguments Args: - name: Tool name (can be prefixed with server name) + server_name: Server name + name: Tool name arguments: Tool arguments user_api_key_auth: User authentication mcp_auth_header: MCP auth header (deprecated) @@ -1223,26 +1225,12 @@ class MCPServerManager: """ start_time = datetime.datetime.now() - # Remove prefix if present to get the original tool name - original_tool_name, server_name_from_prefix = get_server_name_prefix_tool_mcp( - name - ) - # Get the MCP server - mcp_server = self._get_mcp_server_from_tool_name(name) + prefixed_tool_name = add_server_prefix_to_tool_name(name, server_name) + mcp_server = self._get_mcp_server_from_tool_name(prefixed_tool_name) if mcp_server is None: raise ValueError(f"Tool {name} not found") - # Validate that the server from prefix matches the actual server (if prefix was used) - if server_name_from_prefix: - expected_prefix = get_server_prefix(mcp_server) - if normalize_server_name(server_name_from_prefix) != normalize_server_name( - expected_prefix - ): - raise ValueError( - f"Tool {name} server prefix mismatch: expected {expected_prefix}, got {server_name_from_prefix}" - ) - ######################################################### # Pre MCP Tool Call Hook # Allow validation and modification of tool calls before execution @@ -1250,9 +1238,9 @@ class MCPServerManager: ######################################################### if proxy_logging_obj: await self.pre_call_tool_check( - name=original_tool_name, + name=name, arguments=arguments, - server_name_from_prefix=server_name_from_prefix, + server_name=server_name, user_api_key_auth=user_api_key_auth, proxy_logging_obj=proxy_logging_obj, server=mcp_server, @@ -1264,7 +1252,7 @@ class MCPServerManager: during_hook_task = self._create_during_hook_task( name=name, arguments=arguments, - server_name_from_prefix=server_name_from_prefix, + server_name_from_prefix=server_name, user_api_key_auth=user_api_key_auth, proxy_logging_obj=proxy_logging_obj, start_time=start_time, @@ -1285,7 +1273,7 @@ class MCPServerManager: # For regular MCP servers, use the MCP client return await self._call_regular_mcp_tool( mcp_server=mcp_server, - original_tool_name=original_tool_name, + original_tool_name=name, arguments=arguments, tasks=tasks, mcp_auth_header=mcp_auth_header, @@ -1369,12 +1357,16 @@ class MCPServerManager: # If not found and tool name is prefixed, try extracting server name from prefix if is_tool_name_prefixed(tool_name): - _, server_name_from_prefix = get_server_name_prefix_tool_mcp(tool_name) - for server in self.get_registry().values(): - if normalize_server_name(server.name) == normalize_server_name( - server_name_from_prefix - ): - return server + ( + original_tool_name, + server_name_from_prefix, + ) = get_server_name_prefix_tool_mcp(tool_name) + if original_tool_name in self.tool_name_to_mcp_server_name_mapping: + for server in self.get_registry().values(): + if normalize_server_name(server.name) == normalize_server_name( + server_name_from_prefix + ): + return server return None @@ -1414,13 +1406,13 @@ class MCPServerManager: return server return None - def get_mcp_server_names_from_ids(self, server_ids: List[str]) -> List[str]: - server_names = [] + def get_mcp_servers_from_ids(self, server_ids: List[str]) -> List[MCPServer]: + servers = [] registry = self.get_registry() for server in registry.values(): if server.server_id in server_ids: - server_names.append(server.name) - return server_names + servers.append(server) + return servers def get_mcp_server_by_name(self, server_name: str) -> Optional[MCPServer]: """ diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index e380d88ee70..6f4dafb5fe0 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -238,7 +238,7 @@ if MCP_AVAILABLE: ( user_api_key_auth, mcp_auth_header, - _, + mcp_servers, mcp_server_auth_headers, oauth2_headers, raw_headers, @@ -272,6 +272,7 @@ if MCP_AVAILABLE: response = await call_mcp_tool( user_api_key_auth=user_api_key_auth, mcp_auth_header=mcp_auth_header, + mcp_servers=mcp_servers, mcp_server_auth_headers=mcp_server_auth_headers, oauth2_headers=oauth2_headers, raw_headers=raw_headers, @@ -312,31 +313,32 @@ if MCP_AVAILABLE: async def _get_allowed_mcp_servers_from_mcp_server_names( mcp_servers: Optional[List[str]], - allowed_mcp_servers: List[str], - ) -> List[str]: + allowed_mcp_servers: List[MCPServer], + ) -> List[MCPServer]: """ Get the filtered MCP servers from the MCP server names """ - from typing import Set - filtered_server_ids: Set[str] = set() + filtered_server: dict[str, MCPServer] = {} # Filter servers based on mcp_servers parameter if provided if mcp_servers is not None: for server_or_group in mcp_servers: server_name_matched = False - for server_id in allowed_mcp_servers: - server = global_mcp_server_manager.get_mcp_server_by_id(server_id) - + for server in allowed_mcp_servers: if server: match_list = [ s.lower() - for s in [server.alias, server.server_name, server_id] + for s in [ + server.alias, + server.server_name, + server.server_id, + ] if s is not None ] if server_or_group.lower() in match_list: - filtered_server_ids.add(server_id) + filtered_server[server.server_id] = server server_name_matched = True break @@ -349,15 +351,16 @@ if MCP_AVAILABLE: ) # Only include servers that the user has access to for server_id in access_group_server_ids: - if server_id in allowed_mcp_servers: - filtered_server_ids.add(server_id) + for server in allowed_mcp_servers: + if server_id == server.server_id: + filtered_server[server.server_id] = server except Exception as e: verbose_logger.debug( f"Could not resolve '{server_or_group}' as access group: {e}" ) - if filtered_server_ids: - allowed_mcp_servers = list(filtered_server_ids) + if filtered_server: + return list(filtered_server.values()) return allowed_mcp_servers @@ -450,8 +453,11 @@ if MCP_AVAILABLE: return [] # Get allowed MCP servers based on user permissions - allowed_mcp_servers = await global_mcp_server_manager.get_allowed_mcp_servers( - user_api_key_auth + allowed_mcp_server_ids = ( + await global_mcp_server_manager.get_allowed_mcp_servers(user_api_key_auth) + ) + allowed_mcp_servers = global_mcp_server_manager.get_mcp_servers_from_ids( + allowed_mcp_server_ids ) if mcp_servers is not None: @@ -465,8 +471,7 @@ if MCP_AVAILABLE: # Get tools from each allowed server all_tools = [] - for server_id in allowed_mcp_servers: - server = global_mcp_server_manager.get_mcp_server_by_id(server_id) + for server in allowed_mcp_servers: if server is None: continue @@ -504,7 +509,7 @@ if MCP_AVAILABLE: filtered_tools = await filter_tools_by_key_team_permissions( tools=filtered_tools, - server_id=server_id, + server_id=server.server_id, user_api_key_auth=user_api_key_auth, ) @@ -607,6 +612,7 @@ if MCP_AVAILABLE: arguments: Optional[Dict[str, Any]] = None, user_api_key_auth: Optional[UserAPIKeyAuth] = None, mcp_auth_header: Optional[str] = None, + mcp_servers: Optional[List[str]] = None, mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]] = None, oauth2_headers: Optional[Dict[str, str]] = None, raw_headers: Optional[Dict[str, str]] = None, @@ -621,25 +627,31 @@ if MCP_AVAILABLE: status_code=400, detail="Request arguments are required" ) - # Remove prefix from tool name for logging and processing - original_tool_name, server_name_from_prefix = get_server_name_prefix_tool_mcp( - name - ) - ## CHECK IF USER IS ALLOWED TO CALL THIS TOOL allowed_mcp_server_ids = await MCPRequestHandler.get_allowed_mcp_servers( user_api_key_auth=user_api_key_auth, ) - allowed_mcp_servers = global_mcp_server_manager.get_mcp_server_names_from_ids( + allowed_mcp_servers = global_mcp_server_manager.get_mcp_servers_from_ids( allowed_mcp_server_ids ) - if not MCPRequestHandler.is_tool_allowed( + allowed_mcp_servers = await _get_allowed_mcp_servers_from_mcp_server_names( + mcp_servers=mcp_servers, allowed_mcp_servers=allowed_mcp_servers, - server_name=server_name_from_prefix, - ): + ) + server_name: Optional[str] + if len(allowed_mcp_servers) == 1: + original_tool_name, server_name = name, allowed_mcp_servers[0].server_name + else: + # Remove prefix from tool name for logging and processing + original_tool_name, server_name = get_server_name_prefix_tool_mcp(name) + + if not server_name or not MCPRequestHandler.is_tool_allowed( + allowed_mcp_servers=[server.name for server in allowed_mcp_servers], + server_name=server_name, + ): raise HTTPException( status_code=403, detail=f"User not allowed to call this tool. Allowed MCP servers: {allowed_mcp_servers}", @@ -649,16 +661,16 @@ if MCP_AVAILABLE: _get_standard_logging_mcp_tool_call( name=original_tool_name, # Use original name for logging arguments=arguments, - server_name=server_name_from_prefix, + server_name=server_name, ) ) litellm_logging_obj: Optional[LiteLLMLoggingObj] = kwargs.get( "litellm_logging_obj", None ) if litellm_logging_obj: - litellm_logging_obj.model_call_details["mcp_tool_call_metadata"] = ( - standard_logging_mcp_tool_call - ) + litellm_logging_obj.model_call_details[ + "mcp_tool_call_metadata" + ] = standard_logging_mcp_tool_call litellm_logging_obj.model = f"MCP: {name}" # Check if tool exists in local registry first (for OpenAPI-based tools) # These tools are registered with their prefixed names @@ -672,15 +684,16 @@ if MCP_AVAILABLE: # Primary and recommended way to use external MCP servers ######################################################### else: - mcp_server: Optional[MCPServer] = ( - global_mcp_server_manager._get_mcp_server_from_tool_name(name) - ) + mcp_server: Optional[ + MCPServer + ] = global_mcp_server_manager._get_mcp_server_from_tool_name(name) if mcp_server: standard_logging_mcp_tool_call["mcp_server_cost_info"] = ( mcp_server.mcp_info or {} ).get("mcp_server_cost_info") response = await _handle_managed_mcp_tool( - name=name, # Pass the full name (potentially prefixed) + server_name=server_name, + name=original_tool_name, # Pass the full name (potentially prefixed) arguments=arguments, user_api_key_auth=user_api_key_auth, mcp_auth_header=mcp_auth_header, @@ -734,6 +747,7 @@ if MCP_AVAILABLE: ) async def _handle_managed_mcp_tool( + server_name: str, name: str, arguments: Dict[str, Any], user_api_key_auth: Optional[UserAPIKeyAuth] = None, @@ -748,6 +762,7 @@ if MCP_AVAILABLE: from litellm.proxy.proxy_server import proxy_logging_obj call_tool_result = await global_mcp_server_manager.call_tool( + server_name=server_name, name=name, arguments=arguments, user_api_key_auth=user_api_key_auth, @@ -1050,14 +1065,16 @@ if MCP_AVAILABLE: ) auth_context_var.set(auth_user) - def get_auth_context() -> Tuple[ - Optional[UserAPIKeyAuth], - Optional[str], - Optional[List[str]], - Optional[Dict[str, Dict[str, str]]], - Optional[Dict[str, str]], - Optional[Dict[str, str]], - ]: + def get_auth_context() -> ( + Tuple[ + Optional[UserAPIKeyAuth], + Optional[str], + Optional[List[str]], + Optional[Dict[str, Dict[str, str]]], + Optional[Dict[str, str]], + Optional[Dict[str, str]], + ] + ): """ Get the UserAPIKeyAuth from the auth context variable. diff --git a/litellm/responses/main.py b/litellm/responses/main.py index 6450f550402..9cb1691c46f 100644 --- a/litellm/responses/main.py +++ b/litellm/responses/main.py @@ -167,11 +167,12 @@ async def aresponses_api_with_mcp( user_api_key_auth = kwargs.get("user_api_key_auth") # Get original MCP tools (for events) and OpenAI tools (for LLM) by reusing existing methods - original_mcp_tools = ( - await LiteLLM_Proxy_MCP_Handler._process_mcp_tools_without_openai_transform( - user_api_key_auth=user_api_key_auth, - mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy, - ) + ( + original_mcp_tools, + tool_server_map, + ) = await LiteLLM_Proxy_MCP_Handler._process_mcp_tools_without_openai_transform( + user_api_key_auth=user_api_key_auth, + mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy, ) openai_tools = LiteLLM_Proxy_MCP_Handler._transform_mcp_tools_to_openai( original_mcp_tools @@ -230,6 +231,7 @@ async def aresponses_api_with_mcp( mcp_discovery_events=mcp_discovery_events, call_params=call_params, previous_response_id=previous_response_id, + tool_server_map=tool_server_map, **kwargs, ) @@ -274,7 +276,9 @@ async def aresponses_api_with_mcp( "user_api_key_auth" ) tool_results = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( - tool_calls=tool_calls, user_api_key_auth=user_api_key_auth + tool_server_map=tool_server_map, + tool_calls=tool_calls, + user_api_key_auth=user_api_key_auth, ) if tool_results: @@ -320,13 +324,18 @@ async def aresponses_api_with_mcp( ) final_response = MCPEnhancedStreamingIterator( - base_iterator=final_response, mcp_events=tool_execution_events + tool_server_map=tool_server_map, + base_iterator=final_response, + mcp_events=tool_execution_events, ) # Add custom output elements to the final response (for non-streaming) elif isinstance(final_response, ResponsesAPIResponse): # Fetch MCP tools again for output elements (without OpenAI transformation) - mcp_tools_for_output = await LiteLLM_Proxy_MCP_Handler._process_mcp_tools_without_openai_transform( + ( + mcp_tools_for_output, + _, + ) = await LiteLLM_Proxy_MCP_Handler._process_mcp_tools_without_openai_transform( user_api_key_auth=user_api_key_auth, mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy, ) diff --git a/litellm/responses/mcp/litellm_proxy_mcp_handler.py b/litellm/responses/mcp/litellm_proxy_mcp_handler.py index 2600d53a171..7ad26e0a863 100644 --- a/litellm/responses/mcp/litellm_proxy_mcp_handler.py +++ b/litellm/responses/mcp/litellm_proxy_mcp_handler.py @@ -1,6 +1,7 @@ from typing import TYPE_CHECKING, Any, Dict, Iterable, List, Optional, Tuple, Union from litellm._logging import verbose_logger +from litellm.proxy._experimental.mcp_server.utils import get_server_name_prefix_tool_mcp from litellm.responses.main import aresponses from litellm.responses.streaming_iterator import BaseResponsesAPIStreamingIterator from litellm.types.llms.openai import ResponsesAPIResponse, ToolParam @@ -68,16 +69,24 @@ class LiteLLM_Proxy_MCP_Handler: async def _get_mcp_tools_from_manager( user_api_key_auth: Any, mcp_tools_with_litellm_proxy: Optional[Iterable[ToolParam]], - ) -> List[MCPTool]: + ) -> tuple[List[MCPTool], List[str]]: """ Get available tools from the MCP server manager. Args: user_api_key_auth: User authentication info for access control mcp_tools_with_litellm_proxy: ToolParam objects with server_url starting with "litellm_proxy" + + Returns: + List of MCP tools + List names of allowed MCP servers """ from litellm.proxy._experimental.mcp_server.server import ( _get_tools_from_mcp_servers, + _get_allowed_mcp_servers_from_mcp_server_names, + ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, ) mcp_servers: List[str] = [] @@ -92,15 +101,40 @@ class LiteLLM_Proxy_MCP_Handler: ): mcp_servers.append(server_url.split("/")[-1]) - return await _get_tools_from_mcp_servers( + tools = await _get_tools_from_mcp_servers( user_api_key_auth=user_api_key_auth, mcp_auth_header=None, mcp_servers=mcp_servers, mcp_server_auth_headers=None, ) + allowed_mcp_server_ids = ( + await global_mcp_server_manager.get_allowed_mcp_servers(user_api_key_auth) + ) + allowed_mcp_servers = global_mcp_server_manager.get_mcp_servers_from_ids( + allowed_mcp_server_ids + ) + + allowed_mcp_servers = await _get_allowed_mcp_servers_from_mcp_server_names( + mcp_servers=mcp_servers, + allowed_mcp_servers=allowed_mcp_servers, + ) + + server_names: List[str] = [] + for server in allowed_mcp_servers: + if server is None: + continue + server_name = getattr(server, "server_name", None) or getattr( + server, "alias", None + ) or getattr(server, "name", None) + if isinstance(server_name, str): + server_names.append(server_name) + + return tools, server_names @staticmethod - def _deduplicate_mcp_tools(mcp_tools: List[Any]) -> List[Any]: + def _deduplicate_mcp_tools( + mcp_tools: List[MCPTool], allowed_mcp_servers: List[str] + ) -> tuple[List[MCPTool], dict[str, str]]: """ Deduplicate MCP tools by name, keeping the first occurrence of each tool. @@ -109,28 +143,34 @@ class LiteLLM_Proxy_MCP_Handler: Returns: List of deduplicated MCP tools + The returned dictionary maps each tool_name to the server_name """ seen_names = set() deduplicated_tools = [] + tool_server_map: dict[str, str] = {} for tool in mcp_tools: - tool_name = ( - getattr(tool, "name", None) - if hasattr(tool, "name") - else tool.get("name") - if isinstance(tool, dict) - else None - ) + if isinstance(tool, dict): + tool_name = tool.get("name") + else: + tool_name = getattr(tool, "name", None) + if tool_name and tool_name not in seen_names: seen_names.add(tool_name) deduplicated_tools.append(tool) + if len(allowed_mcp_servers) == 1: + tool_server_map[tool_name] = allowed_mcp_servers[0] + else: + tool_server_map[tool_name], _ = get_server_name_prefix_tool_mcp( + tool_name + ) - return deduplicated_tools + return deduplicated_tools, tool_server_map @staticmethod def _filter_mcp_tools_by_allowed_tools( - mcp_tools: List[Any], mcp_tools_with_litellm_proxy: List[ToolParam] - ) -> List[Any]: + mcp_tools: List[MCPTool], mcp_tools_with_litellm_proxy: List[ToolParam] + ) -> List[MCPTool]: """Filter MCP tools based on allowed_tools parameter from the original tool configs.""" # Collect all allowed tool names from all MCP tool configs allowed_tool_names = set() @@ -147,13 +187,11 @@ class LiteLLM_Proxy_MCP_Handler: # Filter tools based on allowed names filtered_tools = [] for mcp_tool in mcp_tools: - tool_name = ( - getattr(mcp_tool, "name", None) - if hasattr(mcp_tool, "name") - else mcp_tool.get("name") - if isinstance(mcp_tool, dict) - else None - ) + if isinstance(mcp_tool, dict): + tool_name = mcp_tool.get("name") + else: + tool_name = getattr(mcp_tool, "name", None) + if tool_name and tool_name in allowed_tool_names: filtered_tools.append(mcp_tool) @@ -162,13 +200,9 @@ class LiteLLM_Proxy_MCP_Handler: @staticmethod async def _process_mcp_tools_to_openai_format( user_api_key_auth: Any, mcp_tools_with_litellm_proxy: List[ToolParam] - ) -> List[Any]: + ) -> tuple[List[Any], dict[str, str]]: """ - Centralized method to process MCP tools through the complete pipeline: - 1. Fetch tools from MCP manager - 2. Filter based on allowed_tools parameter - 3. Deduplicate tools by name - 4. Transform to OpenAI format + Centralized method to process MCP tools through the complete pipeline. Args: user_api_key_auth: User authentication info for access control @@ -176,40 +210,26 @@ class LiteLLM_Proxy_MCP_Handler: Returns: List of tools in OpenAI format ready to be sent to the LLM + The returned dictionary maps each tool_name to the server_name """ - if not mcp_tools_with_litellm_proxy: - return [] - - # Step 1: Fetch MCP tools from manager - mcp_tools_fetched = await LiteLLM_Proxy_MCP_Handler._get_mcp_tools_from_manager( - user_api_key_auth=user_api_key_auth, - mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy, + ( + deduplicated_mcp_tools, + tool_server_map, + ) = await LiteLLM_Proxy_MCP_Handler._process_mcp_tools_without_openai_transform( + user_api_key_auth, + mcp_tools_with_litellm_proxy, ) - # Step 2: Filter tools based on allowed_tools parameter - filtered_mcp_tools = ( - LiteLLM_Proxy_MCP_Handler._filter_mcp_tools_by_allowed_tools( - mcp_tools=mcp_tools_fetched, - mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy, - ) - ) - - # Step 3: Deduplicate tools after filtering - deduplicated_mcp_tools = LiteLLM_Proxy_MCP_Handler._deduplicate_mcp_tools( - filtered_mcp_tools - ) - - # Step 4: Transform to OpenAI format openai_tools = LiteLLM_Proxy_MCP_Handler._transform_mcp_tools_to_openai( deduplicated_mcp_tools ) - return openai_tools + return openai_tools, tool_server_map @staticmethod async def _process_mcp_tools_without_openai_transform( user_api_key_auth: Any, mcp_tools_with_litellm_proxy: List[ToolParam] - ) -> List[Any]: + ) -> tuple[List[Any], dict[str, str]]: """ Process MCP tools through filtering and deduplication pipeline without OpenAI transformation. This is useful for cases where we need the original MCP tool objects (e.g., for events). @@ -222,10 +242,13 @@ class LiteLLM_Proxy_MCP_Handler: List of filtered and deduplicated MCP tools in their original format """ if not mcp_tools_with_litellm_proxy: - return [] + return [], {} # Step 1: Fetch MCP tools from manager - mcp_tools_fetched = await LiteLLM_Proxy_MCP_Handler._get_mcp_tools_from_manager( + ( + mcp_tools_fetched, + allowed_mcp_servers, + ) = await LiteLLM_Proxy_MCP_Handler._get_mcp_tools_from_manager( user_api_key_auth=user_api_key_auth, mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy, ) @@ -239,11 +262,14 @@ class LiteLLM_Proxy_MCP_Handler: ) # Step 3: Deduplicate tools after filtering - deduplicated_mcp_tools = LiteLLM_Proxy_MCP_Handler._deduplicate_mcp_tools( - filtered_mcp_tools + ( + deduplicated_mcp_tools, + tool_server_map, + ) = LiteLLM_Proxy_MCP_Handler._deduplicate_mcp_tools( + filtered_mcp_tools, allowed_mcp_servers ) - return deduplicated_mcp_tools + return deduplicated_mcp_tools, tool_server_map @staticmethod def _transform_mcp_tools_to_openai(mcp_tools: List[Any]) -> List[Any]: @@ -371,7 +397,7 @@ class LiteLLM_Proxy_MCP_Handler: @staticmethod async def _execute_tool_calls( - tool_calls: List[Any], user_api_key_auth: Any + tool_server_map: dict[str, str], tool_calls: List[Any], user_api_key_auth: Any ) -> List[Dict[str, Any]]: """Execute tool calls and return results.""" from fastapi import HTTPException @@ -402,7 +428,10 @@ class LiteLLM_Proxy_MCP_Handler: # Import here to avoid circular import from litellm.proxy.proxy_server import proxy_logging_obj + server_name = tool_server_map[tool_name] + result = await global_mcp_server_manager.call_tool( + server_name=server_name, name=tool_name, arguments=parsed_arguments, user_api_key_auth=user_api_key_auth, @@ -549,6 +578,7 @@ class LiteLLM_Proxy_MCP_Handler: mcp_discovery_events: List[Any], call_params: Dict[str, Any], previous_response_id: Optional[str], + tool_server_map: dict[str, str], **kwargs, ) -> Any: """ @@ -577,6 +607,7 @@ class LiteLLM_Proxy_MCP_Handler: return MCPEnhancedStreamingIterator( base_iterator=None, # Will be created internally mcp_events=mcp_discovery_events, # Pre-generated MCP discovery events + tool_server_map=tool_server_map, mcp_tools_with_litellm_proxy=mcp_tools_with_litellm_proxy, user_api_key_auth=kwargs.get("user_api_key_auth"), original_request_params=request_params, diff --git a/litellm/responses/mcp/mcp_streaming_iterator.py b/litellm/responses/mcp/mcp_streaming_iterator.py index dcc660f0380..ea31f0f7f1d 100644 --- a/litellm/responses/mcp/mcp_streaming_iterator.py +++ b/litellm/responses/mcp/mcp_streaming_iterator.py @@ -257,6 +257,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): self, base_iterator: Any, # Can be None - will be created internally mcp_events: List[ResponsesAPIStreamingResponse], + tool_server_map: dict[str, str], mcp_tools_with_litellm_proxy: Optional[List[Any]] = None, user_api_key_auth: Any = None, original_request_params: Optional[Dict[str, Any]] = None, @@ -280,6 +281,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): self.mcp_events = ( mcp_events # Store the initial MCP events for backward compatibility ) + self.tool_server_map = tool_server_map # Iterator references self.base_iterator: Optional[Union[Any, ResponsesAPIResponse]] = ( @@ -506,7 +508,9 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): # Execute the tools tool_results = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( - tool_calls=tool_calls, user_api_key_auth=self.user_api_key_auth + tool_server_map=self.tool_server_map, + tool_calls=tool_calls, + user_api_key_auth=self.user_api_key_auth, ) # Create completion events and output_item.done events for tool execution @@ -518,9 +522,11 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): tool_name = "unknown" tool_arguments = "{}" for tool_call in tool_calls: - name, args, call_id = ( - LiteLLM_Proxy_MCP_Handler._extract_tool_call_details(tool_call) - ) + ( + name, + args, + call_id, + ) = LiteLLM_Proxy_MCP_Handler._extract_tool_call_details(tool_call) if call_id == tool_call_id: tool_name = name or "unknown" tool_arguments = args or "{}" diff --git a/tests/mcp_tests/test_aresponses_api_with_mcp.py b/tests/mcp_tests/test_aresponses_api_with_mcp.py index 64bc58eb40c..865a580f0ca 100644 --- a/tests/mcp_tests/test_aresponses_api_with_mcp.py +++ b/tests/mcp_tests/test_aresponses_api_with_mcp.py @@ -286,6 +286,8 @@ async def test_mcp_allowed_tools_filtering(): 'inputSchema': {'type': 'object', 'properties': {}} })() ] + + allowed_mcp_servers = ["gitmcp"] # Test Case 1: MCP tool config with allowed_tools specified mcp_tool_config_with_allowed_tools = [ @@ -381,8 +383,8 @@ async def test_mcp_allowed_tools_filtering(): ) # Then deduplicate the filtered tools - filtered_tools_deduplicated = LiteLLM_Proxy_MCP_Handler._deduplicate_mcp_tools( - filtered_tools_with_duplicates + filtered_tools_deduplicated, _ = LiteLLM_Proxy_MCP_Handler._deduplicate_mcp_tools( + filtered_tools_with_duplicates, [] ) # Should only return 1 tool (the duplicate should be removed) @@ -395,7 +397,7 @@ async def test_mcp_allowed_tools_filtering(): print("✓ Test Case 3: duplicate tools are properly deduplicated") # Test Case 3b: Test standalone deduplication method - standalone_deduplicated = LiteLLM_Proxy_MCP_Handler._deduplicate_mcp_tools(mock_mcp_tools_with_duplicates) + standalone_deduplicated, _ = LiteLLM_Proxy_MCP_Handler._deduplicate_mcp_tools(mock_mcp_tools_with_duplicates, allowed_mcp_servers) # Should return 2 unique tools (GitMCP-fetch_litellm_documentation and GitMCP-search_litellm_documentation) assert len(standalone_deduplicated) == 2, f"Expected 2 unique tools after standalone deduplication, got {len(standalone_deduplicated)}" @@ -510,7 +512,7 @@ async def test_streaming_mcp_events_validation(): patch.object(LiteLLM_Proxy_MCP_Handler, '_execute_tool_calls', new_callable=AsyncMock) as mock_execute_tools: # Setup MCP mocks - mock_get_tools.return_value = mock_mcp_tools + mock_get_tools.return_value = (mock_mcp_tools, ["test_server"]) def mock_execute_tool_calls_side_effect(tool_calls, user_api_key_auth): """Mock tool execution with realistic results""" @@ -695,7 +697,7 @@ async def test_streaming_responses_api_with_mcp_tools(): patch.object(LiteLLM_Proxy_MCP_Handler, '_execute_tool_calls', new_callable=AsyncMock) as mock_execute_tools: # Setup MCP mocks only - mock_get_tools.return_value = mock_mcp_tools + mock_get_tools.return_value = (mock_mcp_tools, ["litellm_proxy"]) # Create a dynamic mock that will match the actual tool call ID from the LLM response def mock_execute_tool_calls_side_effect(tool_calls, user_api_key_auth): @@ -1135,4 +1137,4 @@ async def test_no_duplicate_mcp_tools_in_streaming_e2e(): } - \ No newline at end of file + diff --git a/tests/mcp_tests/test_mcp_server.py b/tests/mcp_tests/test_mcp_server.py index 9ee5f6a9a60..fbdada93068 100644 --- a/tests/mcp_tests/test_mcp_server.py +++ b/tests/mcp_tests/test_mcp_server.py @@ -106,6 +106,7 @@ async def test_mcp_server_manager_https_server(): ] = expected_prefix result = await mcp_server_manager.call_tool( + server_name="zapier_mcp_server", name=f"{expected_prefix}-gmail_send_email", arguments={ "body": "Test", @@ -266,6 +267,7 @@ async def test_mcp_http_transport_call_tool_mock(): # Call the tool result = await test_manager.call_tool( + server_name="test_http_server", name="gmail_send_email", arguments={ "to": "test@example.com", @@ -332,6 +334,7 @@ async def test_mcp_http_transport_call_tool_error_mock(): # Call the tool with invalid data result = await test_manager.call_tool( + server_name="test_http_server", name="gmail_send_email", arguments={"to": "invalid-email", "subject": "Test", "body": "Test"}, proxy_logging_obj=None, @@ -370,6 +373,7 @@ async def test_mcp_http_transport_tool_not_found(): # Try to call a tool that doesn't exist in mapping with pytest.raises(ValueError, match="Tool nonexistent_tool not found"): await test_manager.call_tool( + server_name="test_http_server", name="nonexistent_tool", arguments={"param": "value"}, proxy_logging_obj=None, @@ -774,7 +778,7 @@ async def test_get_tools_from_mcp_servers(): mock_manager.get_allowed_mcp_servers = AsyncMock( return_value=["server1_id", "server2_id"] ) - mock_manager.get_mcp_server_by_id = mock_get_server_by_id + mock_manager.get_mcp_servers_from_ids = MagicMock(return_value=[mock_server_1, mock_server_2]) mock_manager._get_tools_from_server = AsyncMock(return_value=[mock_tool_1]) with patch( @@ -796,7 +800,7 @@ async def test_get_tools_from_mcp_servers(): mock_manager_2.get_allowed_mcp_servers = AsyncMock( return_value=["server1_id", "server2_id"] ) - mock_manager_2.get_mcp_server_by_id = mock_get_server_by_id + mock_manager_2.get_mcp_servers_from_ids = MagicMock(return_value=[mock_server_1, mock_server_2]) mock_manager_2._get_tools_from_server = AsyncMock( side_effect=lambda server, mcp_auth_header=None, extra_headers=None, add_prefix=False: ( [mock_tool_1] if server.server_id == "server1_id" else [mock_tool_2] @@ -824,7 +828,7 @@ async def test_get_tools_from_mcp_servers(): mock_manager.get_allowed_mcp_servers = AsyncMock( return_value=["server1_id", "server2_id", "server3_id"] ) - mock_manager.get_mcp_server_by_id = mock_get_server_by_id + mock_manager.get_mcp_servers_from_ids = MagicMock(return_value=[mock_server_1, mock_server_2, mock_server_3]) mock_manager._get_tools_from_server = AsyncMock(return_value=[mock_tool_1]) with patch( @@ -2071,7 +2075,7 @@ async def test_filter_tools_by_allowed_tools_integration(): mock_manager.get_allowed_mcp_servers = AsyncMock( return_value=["test-server-123"] ) - mock_manager.get_mcp_server_by_id = MagicMock(return_value=mock_server) + mock_manager.get_mcp_servers_from_ids = MagicMock(return_value=[mock_server]) # Mock the _get_tools_from_server method to return all tools mock_manager._get_tools_from_server = AsyncMock(return_value=mock_tools) @@ -2109,7 +2113,7 @@ async def test_filter_tools_by_allowed_tools_integration(): # Verify the manager methods were called correctly mock_manager.get_allowed_mcp_servers.assert_called_once_with(mock_user_auth) - mock_manager.get_mcp_server_by_id.assert_called_once_with("test-server-123") + mock_manager.get_mcp_servers_from_ids.assert_called_once_with(["test-server-123"]) mock_manager._get_tools_from_server.assert_called_once() @@ -2179,8 +2183,7 @@ async def test_filter_tools_by_disallowed_tools_integration(): mock_manager.get_allowed_mcp_servers = AsyncMock( return_value=["test-server-456"] ) - mock_manager.get_mcp_server_by_id = MagicMock(return_value=mock_server) - + mock_manager.get_mcp_servers_from_ids = MagicMock(return_value=[mock_server]) # Mock the _get_tools_from_server method to return all tools mock_manager._get_tools_from_server = AsyncMock(return_value=mock_tools) @@ -2217,7 +2220,7 @@ async def test_filter_tools_by_disallowed_tools_integration(): # Verify the manager methods were called correctly mock_manager.get_allowed_mcp_servers.assert_called_once_with(mock_user_auth) - mock_manager.get_mcp_server_by_id.assert_called_once_with("test-server-456") + mock_manager.get_mcp_servers_from_ids.assert_called_once_with(["test-server-456"]) mock_manager._get_tools_from_server.assert_called_once() @@ -2274,7 +2277,7 @@ async def test_filter_tools_no_restrictions_integration(): mock_manager.get_allowed_mcp_servers = AsyncMock( return_value=["test-server-000"] ) - mock_manager.get_mcp_server_by_id = MagicMock(return_value=mock_server) + mock_manager.get_mcp_servers_from_ids = MagicMock(return_value=[mock_server]) # Mock the _get_tools_from_server method to return all tools mock_manager._get_tools_from_server = AsyncMock(return_value=mock_tools) 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 769785b9214..2936b48c755 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 @@ -1,4 +1,5 @@ import asyncio +from typing import Optional from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -90,18 +91,27 @@ async def test_get_tools_from_mcp_servers_continues_when_one_server_fails(): working_server.alias = "working" working_server.allowed_tools = None working_server.disallowed_tools = None + working_server.server_id = "working_server" + working_server.server_name = "working_server" + working_server.auth_type = None + working_server.extra_headers = None failing_server = MagicMock() failing_server.name = "failing_server" failing_server.alias = "failing" failing_server.allowed_tools = None failing_server.disallowed_tools = None + failing_server.server_id = "failing_server" + failing_server.server_name = "failing_server" + failing_server.auth_type = None + failing_server.extra_headers = None # Mock global_mcp_server_manager mock_manager = MagicMock() mock_manager.get_allowed_mcp_servers = AsyncMock( return_value=["working_server", "failing_server"] ) + mock_manager.get_mcp_servers_from_ids = MagicMock(return_value=[working_server, failing_server]) mock_manager.get_mcp_server_by_id = lambda server_id: ( working_server if server_id == "working_server" else failing_server ) @@ -138,7 +148,7 @@ async def test_get_tools_from_mcp_servers_continues_when_one_server_fails(): result = await _get_tools_from_mcp_servers( user_api_key_auth=user_api_key_auth, mcp_auth_header=None, - mcp_servers=None, + mcp_servers=["working_server", "failing_server"], mcp_server_auth_headers=mcp_server_auth_headers, ) @@ -176,16 +186,29 @@ async def test_get_tools_from_mcp_servers_handles_all_servers_failing(): failing_server1 = MagicMock() failing_server1.name = "failing_server1" failing_server1.alias = "failing1" + failing_server1.allowed_tools = None + failing_server1.disallowed_tools = None + failing_server1.server_id = "failing_server1" + failing_server1.server_name = "failing_server1" + failing_server1.auth_type = None + failing_server1.extra_headers = None failing_server2 = MagicMock() failing_server2.name = "failing_server2" failing_server2.alias = "failing2" + failing_server2.allowed_tools = None + failing_server2.disallowed_tools = None + failing_server2.server_id = "failing_server2" + failing_server2.server_name = "failing_server2" + failing_server2.auth_type = None + failing_server2.extra_headers = None # Mock global_mcp_server_manager mock_manager = MagicMock() mock_manager.get_allowed_mcp_servers = AsyncMock( return_value=["failing_server1", "failing_server2"] ) + mock_manager.get_mcp_servers_from_ids = MagicMock(return_value=[failing_server1, failing_server2]) mock_manager.get_mcp_server_by_id = lambda server_id: ( failing_server1 if server_id == "failing_server1" else failing_server2 ) @@ -592,13 +615,14 @@ async def test_list_tools_single_server_unprefixed_names(): server.alias = "zapier" server.allowed_tools = None server.disallowed_tools = None + server.server_name = "server1" + server.auth_type = None + server.extra_headers = None # Mock manager: allow just one server and return a tool based on add_prefix flag mock_manager = MagicMock() mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["server1"]) - mock_manager.get_mcp_server_by_id = lambda server_id: ( - server if server_id == "server1" else None - ) + mock_manager.get_mcp_servers_from_ids = MagicMock(return_value=[server]) async def mock_get_tools_from_server( server, mcp_auth_header=None, extra_headers=None, add_prefix=False @@ -649,6 +673,9 @@ async def test_list_tools_multiple_servers_prefixed_names(): server1.alias = "zapier" server1.allowed_tools = None server1.disallowed_tools = None + server1.server_name = "server1" + server1.auth_type = None + server1.extra_headers = None server2 = MagicMock() server2.server_id = "server2" @@ -656,12 +683,16 @@ async def test_list_tools_multiple_servers_prefixed_names(): server2.alias = "jira" server2.allowed_tools = None server2.disallowed_tools = None + server2.server_name = "server2" + server2.auth_type = None + server2.extra_headers = None # Mock manager mock_manager = MagicMock() mock_manager.get_allowed_mcp_servers = AsyncMock( return_value=["server1", "server2"] ) + mock_manager.get_mcp_servers_from_ids = MagicMock(return_value=[server1, server2]) mock_manager.get_mcp_server_by_id = lambda server_id: ( server1 if server_id == "server1" else server2 ) @@ -710,13 +741,35 @@ async def test_call_mcp_tool_user_unauthorized_access(): object_permission_id="key-permission-123", ) - # Mock global_mcp_server_manager.get_mcp_server_names_from_ids to return + # Mock global_mcp_server_manager.get_mcp_servers_from_ids to return # a list that doesn't include "restricted_server" (the server the user is trying to access) with patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_names_from_ids" + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_allowed_mcp_servers", + AsyncMock(return_value=["allowed_server", "another_server"]), + ), patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_servers_from_ids" ) as mock_get_server_names: - # User has access to "allowed_server" but not "restricted_server" - mock_get_server_names.return_value = ["allowed_server", "another_server"] + allowed_server_obj = MagicMock() + allowed_server_obj.name = "allowed_server" + allowed_server_obj.server_name = "allowed_server" + allowed_server_obj.server_id = "allowed_server" + allowed_server_obj.alias = "allowed_server" + allowed_server_obj.allowed_tools = None + allowed_server_obj.disallowed_tools = None + allowed_server_obj.auth_type = None + allowed_server_obj.extra_headers = None + + another_server_obj = MagicMock() + another_server_obj.name = "another_server" + another_server_obj.server_name = "another_server" + another_server_obj.server_id = "another_server" + another_server_obj.alias = "another_server" + another_server_obj.allowed_tools = None + another_server_obj.disallowed_tools = None + another_server_obj.auth_type = None + another_server_obj.extra_headers = None + + mock_get_server_names.return_value = [allowed_server_obj, another_server_obj] # Try to call a tool from "restricted_server" - should raise HTTPException with 403 status with pytest.raises(HTTPException) as exc_info: @@ -770,10 +823,14 @@ async def test_list_tools_filters_by_key_team_permissions(): server.alias = "test" server.allowed_tools = None server.disallowed_tools = None + server.server_name = "server1" + server.auth_type = None + server.extra_headers = None # Mock manager mock_manager = MagicMock() mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["server1"]) + mock_manager.get_mcp_servers_from_ids = MagicMock(return_value=[server]) mock_manager.get_mcp_server_by_id = lambda server_id: server async def mock_get_tools_from_server( @@ -868,10 +925,14 @@ async def test_list_tools_with_team_tool_permissions_inheritance(): server.alias = "test" server.allowed_tools = None server.disallowed_tools = None + server.server_name = "server1" + server.auth_type = None + server.extra_headers = None # Mock manager mock_manager = MagicMock() mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["server1"]) + mock_manager.get_mcp_servers_from_ids = MagicMock(return_value=[server]) mock_manager.get_mcp_server_by_id = lambda server_id: server async def mock_get_tools_from_server( @@ -951,10 +1012,14 @@ async def test_list_tools_with_no_tool_permissions_shows_all(): server.alias = "test" server.allowed_tools = None server.disallowed_tools = None + server.server_name = "server1" + server.auth_type = None + server.extra_headers = None # Mock manager mock_manager = MagicMock() mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["server1"]) + mock_manager.get_mcp_servers_from_ids = MagicMock(return_value=[server]) mock_manager.get_mcp_server_by_id = lambda server_id: server async def mock_get_tools_from_server( @@ -1044,7 +1109,7 @@ async def test_list_tools_strips_prefix_when_matching_permissions(): # Mock manager mock_manager = MagicMock() mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["gitmcp_server"]) - mock_manager.get_mcp_server_by_id = lambda server_id: server + mock_manager.get_mcp_servers_from_ids = MagicMock(return_value=[server]) async def mock_get_tools_from_server( server, mcp_auth_header=None, extra_headers=None, add_prefix=True 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 cce318654cc..f36ae1aec0f 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 @@ -453,7 +453,7 @@ class TestMCPServerManager: await manager.pre_call_tool_check( name="allowed_tool", arguments={"param": "value"}, - server_name_from_prefix="test-server", + server_name="test-server", user_api_key_auth=user_api_key_auth, proxy_logging_obj=proxy_logging_obj, server=server, @@ -482,7 +482,7 @@ class TestMCPServerManager: await manager.pre_call_tool_check( name="blocked_tool", arguments={"param": "value"}, - server_name_from_prefix="test-server", + server_name="test-server", user_api_key_auth=user_api_key_auth, proxy_logging_obj=proxy_logging_obj, server=server, @@ -529,7 +529,7 @@ class TestMCPServerManager: await manager.pre_call_tool_check( name="allowed_tool", arguments={"param": "value"}, - server_name_from_prefix="test-server", + server_name="test-server", user_api_key_auth=user_api_key_auth, proxy_logging_obj=proxy_logging_obj, server=server, @@ -558,7 +558,7 @@ class TestMCPServerManager: await manager.pre_call_tool_check( name="banned_tool", arguments={"param": "value"}, - server_name_from_prefix="test-server", + server_name="test-server", user_api_key_auth=user_api_key_auth, proxy_logging_obj=proxy_logging_obj, server=server, @@ -605,7 +605,7 @@ class TestMCPServerManager: await manager.pre_call_tool_check( name="any_tool", arguments={"param": "value"}, - server_name_from_prefix="test-server", + server_name="test-server", user_api_key_auth=user_api_key_auth, proxy_logging_obj=proxy_logging_obj, server=server, @@ -644,7 +644,7 @@ class TestMCPServerManager: await manager.pre_call_tool_check( name="tool2", arguments={"param": "value"}, - server_name_from_prefix="test-server", + server_name="test-server", user_api_key_auth=user_api_key_auth, proxy_logging_obj=proxy_logging_obj, server=server, @@ -655,7 +655,7 @@ class TestMCPServerManager: await manager.pre_call_tool_check( name="tool3", arguments={"param": "value"}, - server_name_from_prefix="test-server", + server_name="test-server", user_api_key_auth=user_api_key_auth, proxy_logging_obj=proxy_logging_obj, server=server, @@ -992,9 +992,9 @@ class TestMCPServerManager: # Should succeed await manager.pre_call_tool_check( + server_name="Test Server", name="read_wiki_structure", arguments={"repoName": "facebook/react"}, - server_name_from_prefix="test", user_api_key_auth=user_auth, proxy_logging_obj=proxy_logging, server=server, @@ -1038,9 +1038,9 @@ class TestMCPServerManager: # Should fail with 403 with pytest.raises(HTTPException) as exc_info: await manager.pre_call_tool_check( + server_name="Test Server", name="ask_question", arguments={"question": "test"}, - server_name_from_prefix="test", user_api_key_auth=user_auth, proxy_logging_obj=proxy_logging, server=server, @@ -1186,7 +1186,7 @@ class TestMCPServerManager: await manager.pre_call_tool_check( name="getpetbyid", arguments={"petId": "1"}, - server_name_from_prefix="my_api_mcp", + server_name="my_api_mcp", user_api_key_auth=user_api_key_auth, proxy_logging_obj=proxy_logging_obj, server=server, @@ -1196,7 +1196,7 @@ class TestMCPServerManager: await manager.pre_call_tool_check( name="findpetsbystatus", arguments={"status": "available"}, - server_name_from_prefix="my_api_mcp", + server_name="my_api_mcp", user_api_key_auth=user_api_key_auth, proxy_logging_obj=proxy_logging_obj, server=server, @@ -1207,7 +1207,7 @@ class TestMCPServerManager: await manager.pre_call_tool_check( name="deletepet", arguments={"petId": "1"}, - server_name_from_prefix="my_api_mcp", + server_name="my_api_mcp", user_api_key_auth=user_api_key_auth, proxy_logging_obj=proxy_logging_obj, server=server, @@ -1245,6 +1245,7 @@ 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" # Create mock client that tracks context manager usage mock_client = MagicMock() @@ -1302,6 +1303,7 @@ class TestMCPServerManager: # Call the tool result = await manager.call_tool( + server_name="test-server", name="test_tool", arguments={"param": "value"}, user_api_key_auth=user_api_key_auth,