Litellm mcp tool prefix (#12289)

* prefix all server tools from the backend

* change var name

* fix tests
This commit is contained in:
Jugal D. Bhatt 2025-07-04 00:09:37 +05:30 • committed by GitHub
parent 1dee70efba
commit df161f1e45
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 240 additions and 61 deletions

View file

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

View file

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

View file

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

View file

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