mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Litellm mcp tool prefix (#12289)
* prefix all server tools from the backend * change var name * fix tests
This commit is contained in:
parent
1dee70efba
commit
df161f1e45
4 changed files with 240 additions and 61 deletions
|
|
@ -30,6 +30,7 @@ from litellm.proxy._types import (
|
|||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPInfo, MCPServer
|
||||
from litellm.proxy._experimental.mcp_server.utils import add_server_prefix_to_tool_name, normalize_server_name, get_server_name_prefix_tool_mcp, is_tool_name_prefixed
|
||||
|
||||
|
||||
class MCPServerManager:
|
||||
|
|
@ -216,13 +217,14 @@ class MCPServerManager:
|
|||
|
||||
async def _get_tools_from_server(self, server: MCPServer, mcp_auth_header: Optional[str] = None) -> List[MCPTool]:
|
||||
"""
|
||||
Helper method to get tools from a single MCP server.
|
||||
Helper method to get tools from a single MCP server with prefixed names.
|
||||
|
||||
Args:
|
||||
server (MCPServer): The server to query tools from
|
||||
mcp_auth_header: Optional auth header for MCP server
|
||||
|
||||
Returns:
|
||||
List[MCPTool]: List of tools available on the server
|
||||
List[MCPTool]: List of tools available on the server with prefixed names
|
||||
"""
|
||||
verbose_logger.debug(f"Connecting to url: {server.url}")
|
||||
verbose_logger.info("_get_tools_from_server...")
|
||||
|
|
@ -233,7 +235,7 @@ class MCPServerManager:
|
|||
server=server,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
)
|
||||
|
||||
|
||||
# Create a task for the client operations to ensure proper cancellation handling
|
||||
async def _list_tools_task():
|
||||
async with client:
|
||||
|
|
@ -243,12 +245,26 @@ class MCPServerManager:
|
|||
|
||||
try:
|
||||
tools = await _list_tools_task()
|
||||
|
||||
# Update tool to server mapping
|
||||
for tool in tools:
|
||||
self.tool_name_to_mcp_server_name_mapping[tool.name] = server.name
|
||||
|
||||
return tools
|
||||
# Create new tools with prefixed names
|
||||
prefixed_tools = []
|
||||
for tool in tools:
|
||||
# Create prefixed tool name
|
||||
prefixed_name = add_server_prefix_to_tool_name(tool.name, server.name)
|
||||
|
||||
# Create new tool with prefixed name
|
||||
prefixed_tool = MCPTool(
|
||||
name=prefixed_name,
|
||||
description=tool.description,
|
||||
inputSchema=tool.inputSchema
|
||||
)
|
||||
prefixed_tools.append(prefixed_tool)
|
||||
|
||||
# Update tool to server mapping with both original and prefixed names
|
||||
self.tool_name_to_mcp_server_name_mapping[tool.name] = server.name
|
||||
self.tool_name_to_mcp_server_name_mapping[prefixed_name] = server.name
|
||||
|
||||
return prefixed_tools
|
||||
except asyncio.CancelledError:
|
||||
verbose_logger.warning(f"Task cancelled while listing tools from {server.name}")
|
||||
raise # Re-raise the cancellation
|
||||
|
|
@ -264,32 +280,51 @@ class MCPServerManager:
|
|||
await client.disconnect()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
async def call_tool(
|
||||
self,
|
||||
name: str,
|
||||
arguments: Dict[str, Any],
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
|
||||
mcp_auth_header: Optional[str] = None,
|
||||
self,
|
||||
name: str,
|
||||
arguments: Dict[str, Any],
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
|
||||
mcp_auth_header: Optional[str] = None,
|
||||
) -> CallToolResult:
|
||||
"""
|
||||
Call a tool with the given name and arguments
|
||||
Call a tool with the given name and arguments (handles prefixed tool names)
|
||||
|
||||
Args:
|
||||
name: Tool name (can be prefixed with server name)
|
||||
arguments: Tool arguments
|
||||
user_api_key_auth: User authentication
|
||||
mcp_auth_header: MCP auth header
|
||||
|
||||
Returns:
|
||||
CallToolResult from the MCP server
|
||||
"""
|
||||
# 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)
|
||||
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 and normalize_server_name(server_name_from_prefix) != normalize_server_name(mcp_server.name):
|
||||
raise ValueError(
|
||||
f"Tool {name} server prefix mismatch: expected {mcp_server.name}, got {server_name_from_prefix}")
|
||||
|
||||
client = self._create_mcp_client(
|
||||
server=mcp_server,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
)
|
||||
async with client:
|
||||
# Use the original tool name (without prefix) for the actual call
|
||||
call_tool_params = MCPCallToolRequestParams(
|
||||
name=name,
|
||||
name=original_tool_name,
|
||||
arguments=arguments,
|
||||
)
|
||||
return await client.call_tool(call_tool_params)
|
||||
|
||||
|
||||
#########################################################
|
||||
# End of Methods that call the upstream MCP servers
|
||||
#########################################################
|
||||
|
|
@ -312,20 +347,41 @@ class MCPServerManager:
|
|||
async def _initialize_tool_name_to_mcp_server_name_mapping(self):
|
||||
"""
|
||||
Call list_tools for each server and update the tool name to MCP server name mapping
|
||||
Note: This now handles prefixed tool names
|
||||
"""
|
||||
for server in self.get_registry().values():
|
||||
tools = await self._get_tools_from_server(server)
|
||||
for tool in tools:
|
||||
# The tool.name here is already prefixed from _get_tools_from_server
|
||||
# Extract original name for mapping
|
||||
original_name, _ = get_server_name_prefix_tool_mcp(tool.name)
|
||||
self.tool_name_to_mcp_server_name_mapping[original_name] = server.name
|
||||
self.tool_name_to_mcp_server_name_mapping[tool.name] = server.name
|
||||
|
||||
def _get_mcp_server_from_tool_name(self, tool_name: str) -> Optional[MCPServer]:
|
||||
"""
|
||||
Get the MCP Server from the tool name
|
||||
Get the MCP Server from the tool name (handles both prefixed and non-prefixed names)
|
||||
|
||||
Args:
|
||||
tool_name: Tool name (can be prefixed or non-prefixed)
|
||||
|
||||
Returns:
|
||||
MCPServer if found, None otherwise
|
||||
"""
|
||||
# First try with the original tool name
|
||||
if tool_name in self.tool_name_to_mcp_server_name_mapping:
|
||||
server_name = self.tool_name_to_mcp_server_name_mapping[tool_name]
|
||||
for server in self.get_registry().values():
|
||||
if server.name == self.tool_name_to_mcp_server_name_mapping[tool_name]:
|
||||
if normalize_server_name(server.name) == normalize_server_name(server_name):
|
||||
return server
|
||||
|
||||
# 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
|
||||
|
||||
return None
|
||||
|
||||
async def _add_mcp_servers_from_db_to_in_memory_registry(self):
|
||||
|
|
|
|||
|
|
@ -16,14 +16,16 @@ from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
|
|||
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
|
||||
MCPRequestHandler,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.utils import (
|
||||
LITELLM_MCP_SERVER_NAME,
|
||||
LITELLM_MCP_SERVER_VERSION,
|
||||
LITELLM_MCP_SERVER_DESCRIPTION,
|
||||
)
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPInfo
|
||||
from litellm.types.utils import StandardLoggingMCPToolCall
|
||||
from litellm.utils import client
|
||||
|
||||
LITELLM_MCP_SERVER_NAME = "litellm-mcp-server"
|
||||
LITELLM_MCP_SERVER_VERSION = "1.0.0"
|
||||
LITELLM_MCP_SERVER_DESCRIPTION = "MCP Server for LiteLLM"
|
||||
|
||||
# Check if MCP is available
|
||||
# "mcp" requires python 3.10 or higher, but several litellm users use python 3.8
|
||||
|
|
@ -65,6 +67,7 @@ if MCP_AVAILABLE:
|
|||
from litellm.proxy._experimental.mcp_server.tool_registry import (
|
||||
global_mcp_tool_registry,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.utils import get_server_name_prefix_tool_mcp
|
||||
|
||||
######################################################
|
||||
############ MCP Tools List REST API Response Object #
|
||||
|
|
@ -249,23 +252,27 @@ if MCP_AVAILABLE:
|
|||
|
||||
@client
|
||||
async def call_mcp_tool(
|
||||
name: str,
|
||||
arguments: Optional[Dict[str, Any]] = None,
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
|
||||
mcp_auth_header: Optional[str] = None,
|
||||
**kwargs: Any
|
||||
name: str,
|
||||
arguments: Optional[Dict[str, Any]] = None,
|
||||
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
|
||||
mcp_auth_header: Optional[str] = None,
|
||||
**kwargs: Any
|
||||
) -> List[Union[MCPTextContent, MCPImageContent, MCPEmbeddedResource]]:
|
||||
"""
|
||||
Call a specific tool with the provided arguments
|
||||
Call a specific tool with the provided arguments (handles prefixed tool names)
|
||||
"""
|
||||
if arguments is None:
|
||||
raise HTTPException(
|
||||
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)
|
||||
|
||||
standard_logging_mcp_tool_call: StandardLoggingMCPToolCall = (
|
||||
_get_standard_logging_mcp_tool_call(
|
||||
name=name,
|
||||
name=original_tool_name, # Use original name for logging
|
||||
arguments=arguments,
|
||||
)
|
||||
)
|
||||
|
|
@ -283,17 +290,17 @@ if MCP_AVAILABLE:
|
|||
standard_logging_mcp_tool_call.get("mcp_server_name")
|
||||
)
|
||||
|
||||
# Try managed server tool first
|
||||
# Try managed server tool first (pass the full prefixed name)
|
||||
if name in global_mcp_server_manager.tool_name_to_mcp_server_name_mapping:
|
||||
return await _handle_managed_mcp_tool(
|
||||
name=name,
|
||||
name=name, # Pass the full name (potentially prefixed)
|
||||
arguments=arguments,
|
||||
user_api_key_auth=user_api_key_auth,
|
||||
mcp_auth_header=mcp_auth_header,
|
||||
)
|
||||
|
||||
# Fall back to local tool registry
|
||||
return await _handle_local_mcp_tool(name, arguments)
|
||||
# Fall back to local tool registry (use original name)
|
||||
return await _handle_local_mcp_tool(original_tool_name, arguments)
|
||||
|
||||
def _get_standard_logging_mcp_tool_call(
|
||||
name: str,
|
||||
|
|
@ -331,9 +338,12 @@ if MCP_AVAILABLE:
|
|||
return call_tool_result.content # type: ignore[return-value]
|
||||
|
||||
async def _handle_local_mcp_tool(
|
||||
name: str, arguments: Dict[str, Any]
|
||||
name: str, arguments: Dict[str, Any]
|
||||
) -> List[Union[MCPTextContent, MCPImageContent, MCPEmbeddedResource]]:
|
||||
"""Handle tool execution for local registry tools"""
|
||||
"""
|
||||
Handle tool execution for local registry tools
|
||||
Note: Local tools don't use prefixes, so we use the original name
|
||||
"""
|
||||
tool = global_mcp_tool_registry.get_tool(name)
|
||||
if not tool:
|
||||
raise HTTPException(status_code=404, detail=f"Tool '{name}' not found")
|
||||
|
|
@ -344,6 +354,7 @@ if MCP_AVAILABLE:
|
|||
except Exception as e:
|
||||
return [MCPTextContent(text=f"Error: {str(e)}", type="text")]
|
||||
|
||||
|
||||
async def handle_streamable_http_mcp(
|
||||
scope: Scope, receive: Receive, send: Send
|
||||
) -> None:
|
||||
|
|
|
|||
|
|
@ -1,5 +1,16 @@
|
|||
"""
|
||||
MCP Server Utilities
|
||||
"""
|
||||
from typing import Tuple
|
||||
|
||||
import importlib
|
||||
|
||||
# Constants
|
||||
LITELLM_MCP_SERVER_NAME = "litellm-mcp-server"
|
||||
LITELLM_MCP_SERVER_VERSION = "1.0.0"
|
||||
LITELLM_MCP_SERVER_DESCRIPTION = "MCP Server for LiteLLM"
|
||||
MCP_TOOL_PREFIX_SEPARATOR = "/"
|
||||
MCP_TOOL_PREFIX_FORMAT = "{server_name}{separator}{tool_name}"
|
||||
|
||||
def is_mcp_available() -> bool:
|
||||
"""
|
||||
|
|
@ -10,3 +21,56 @@ def is_mcp_available() -> bool:
|
|||
return True
|
||||
except ImportError:
|
||||
return False
|
||||
|
||||
def normalize_server_name(server_name: str) -> str:
|
||||
"""
|
||||
Normalize server name by replacing spaces with underscores
|
||||
"""
|
||||
return server_name.replace(" ", "_")
|
||||
|
||||
def add_server_prefix_to_tool_name(tool_name: str, server_name: str) -> str:
|
||||
"""
|
||||
Add server name prefix to tool name
|
||||
|
||||
Args:
|
||||
tool_name: Original tool name
|
||||
server_name: MCP server name
|
||||
|
||||
Returns:
|
||||
Prefixed tool name in format: server_name::tool_name
|
||||
"""
|
||||
formatted_server_name = normalize_server_name(server_name)
|
||||
|
||||
return MCP_TOOL_PREFIX_FORMAT.format(
|
||||
server_name=formatted_server_name,
|
||||
separator=MCP_TOOL_PREFIX_SEPARATOR,
|
||||
tool_name=tool_name
|
||||
)
|
||||
|
||||
def get_server_name_prefix_tool_mcp(prefixed_tool_name: str) -> Tuple[str, str]:
|
||||
"""
|
||||
Remove server name prefix from tool name
|
||||
|
||||
Args:
|
||||
prefixed_tool_name: Tool name with server prefix
|
||||
|
||||
Returns:
|
||||
Tuple of (original_tool_name, server_name)
|
||||
"""
|
||||
if MCP_TOOL_PREFIX_SEPARATOR in prefixed_tool_name:
|
||||
parts = prefixed_tool_name.split(MCP_TOOL_PREFIX_SEPARATOR, 1)
|
||||
if len(parts) == 2:
|
||||
return parts[1], parts[0] # tool_name, server_name
|
||||
return prefixed_tool_name, "" # No prefix found, return original name
|
||||
|
||||
def is_tool_name_prefixed(tool_name: str) -> bool:
|
||||
"""
|
||||
Check if tool name has server prefix
|
||||
|
||||
Args:
|
||||
tool_name: Tool name to check
|
||||
|
||||
Returns:
|
||||
True if tool name is prefixed, False otherwise
|
||||
"""
|
||||
return MCP_TOOL_PREFIX_SEPARATOR in tool_name
|
||||
|
|
|
|||
|
|
@ -42,26 +42,76 @@ async def test_mcp_server_manager():
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_server_manager_https_server():
|
||||
mcp_server_manager.load_servers_from_config(
|
||||
{
|
||||
"zapier_mcp_server": {
|
||||
"url": os.environ.get("ZAPIER_MCP_HTTPS_SERVER_URL"),
|
||||
"transport": MCPTransport.http,
|
||||
# Create mock tools and results
|
||||
mock_tools = [
|
||||
MCPTool(
|
||||
name="gmail_send_email",
|
||||
description="Send an email via Gmail",
|
||||
inputSchema={
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"body": {"type": "string"},
|
||||
"message": {"type": "string"},
|
||||
"instructions": {"type": "string"}
|
||||
},
|
||||
"required": ["body"]
|
||||
}
|
||||
}
|
||||
)
|
||||
]
|
||||
|
||||
mock_result = CallToolResult(
|
||||
content=[TextContent(type="text", text="Email sent successfully")],
|
||||
isError=False
|
||||
)
|
||||
tools = await mcp_server_manager.list_tools()
|
||||
print("TOOLS FROM MCP SERVER MANAGER== ", tools)
|
||||
|
||||
result = await mcp_server_manager.call_tool(
|
||||
name="gmail_send_email",
|
||||
arguments={
|
||||
"body": "Test",
|
||||
"message": "Test",
|
||||
"instructions": "Test",
|
||||
},
|
||||
)
|
||||
print("RESULT FROM CALLING TOOL FROM MCP SERVER MANAGER== ", result)
|
||||
|
||||
# Create a mock MCPClient
|
||||
mock_client = AsyncMock()
|
||||
mock_client.list_tools = AsyncMock(return_value=mock_tools)
|
||||
mock_client.call_tool = AsyncMock(return_value=mock_result)
|
||||
mock_client.__aenter__ = AsyncMock(return_value=mock_client)
|
||||
mock_client.__aexit__ = AsyncMock(return_value=None)
|
||||
|
||||
# Mock the MCPClient constructor
|
||||
def mock_client_constructor(*args, **kwargs):
|
||||
return mock_client
|
||||
|
||||
with patch('litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient', mock_client_constructor):
|
||||
mcp_server_manager.load_servers_from_config(
|
||||
{
|
||||
"zapier_mcp_server": {
|
||||
"url": "https://test-mcp-server.com/mcp",
|
||||
"transport": MCPTransport.http,
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
tools = await mcp_server_manager.list_tools()
|
||||
print("TOOLS FROM MCP SERVER MANAGER== ", tools)
|
||||
|
||||
# Verify tools were returned and properly prefixed
|
||||
assert len(tools) == 1
|
||||
assert tools[0].name == "zapier_mcp_server/gmail_send_email"
|
||||
|
||||
result = await mcp_server_manager.call_tool(
|
||||
name="zapier_mcp_server/gmail_send_email",
|
||||
arguments={
|
||||
"body": "Test",
|
||||
"message": "Test",
|
||||
"instructions": "Test",
|
||||
},
|
||||
)
|
||||
print("RESULT FROM CALLING TOOL FROM MCP SERVER MANAGER== ", result)
|
||||
|
||||
# Verify result
|
||||
assert result.isError is False
|
||||
assert len(result.content) == 1
|
||||
assert isinstance(result.content[0], TextContent)
|
||||
assert result.content[0].text == "Email sent successfully"
|
||||
|
||||
# Verify client methods were called
|
||||
mock_client.__aenter__.assert_called()
|
||||
mock_client.list_tools.assert_called_once()
|
||||
mock_client.call_tool.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -127,16 +177,16 @@ async def test_mcp_http_transport_list_tools_mock():
|
|||
|
||||
# Assertions
|
||||
assert len(tools) == 2
|
||||
assert tools[0].name == "gmail_send_email"
|
||||
assert tools[1].name == "calendar_create_event"
|
||||
assert tools[0].name == "test_http_server/gmail_send_email"
|
||||
assert tools[1].name == "test_http_server/calendar_create_event"
|
||||
|
||||
# Verify client methods were called
|
||||
mock_client.__aenter__.assert_called()
|
||||
mock_client.list_tools.assert_called_once()
|
||||
|
||||
# Verify tool mapping was updated
|
||||
assert test_manager.tool_name_to_mcp_server_name_mapping["gmail_send_email"] == "test_http_server"
|
||||
assert test_manager.tool_name_to_mcp_server_name_mapping["calendar_create_event"] == "test_http_server"
|
||||
assert test_manager.tool_name_to_mcp_server_name_mapping["test_http_server/gmail_send_email"] == "test_http_server"
|
||||
assert test_manager.tool_name_to_mcp_server_name_mapping["test_http_server/calendar_create_event"] == "test_http_server"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -587,12 +637,10 @@ async def test_list_tools_rest_api_success():
|
|||
server_id=server_id,
|
||||
user_api_key_dict=mock_user_auth
|
||||
)
|
||||
|
||||
|
||||
assert isinstance(response, dict)
|
||||
assert len(response["tools"]) == 1
|
||||
assert response["tools"][0].name == "test_tool"
|
||||
assert response["error"] is None
|
||||
assert response["message"] == "Successfully retrieved tools"
|
||||
assert response["tools"][0].name == "test_server/test_tool"
|
||||
finally:
|
||||
# Restore original state
|
||||
global_mcp_server_manager.registry = {}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue