mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
ae3178d5d4
commit
2843dab7fe
9 changed files with 308 additions and 181 deletions
|
|
@ -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]:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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 "{}"
|
||||
|
|
|
|||
|
|
@ -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():
|
|||
}
|
||||
|
||||
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue