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
This commit is contained in:
YutaSaito 2025-11-13 06:50:52 +09:00 • committed by GitHub
parent ae3178d5d4
commit 2843dab7fe
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
9 changed files with 308 additions and 181 deletions

View file

@ -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]:
"""

View file

@ -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.

View file

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

View file

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

View file

@ -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 "{}"

View file

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

View file

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

View file

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

View file

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