From df161f1e4522487e5b86d838a0f2e2da2e320c9a Mon Sep 17 00:00:00 2001 From: "Jugal D. Bhatt" <55304795+jugaldb@users.noreply.github.com> Date: Fri, 4 Jul 2025 00:09:37 +0530 Subject: [PATCH] Litellm mcp tool prefix (#12289) * prefix all server tools from the backend * change var name * fix tests --- .../mcp_server/mcp_server_manager.py | 94 ++++++++++++---- .../proxy/_experimental/mcp_server/server.py | 43 +++++--- .../proxy/_experimental/mcp_server/utils.py | 64 +++++++++++ tests/mcp_tests/test_mcp_server.py | 100 +++++++++++++----- 4 files changed, 240 insertions(+), 61 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index ed25889d3ae..6cafaeec3c4 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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): diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 98dfda8eb25..8dd4a6be6bf 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -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: diff --git a/litellm/proxy/_experimental/mcp_server/utils.py b/litellm/proxy/_experimental/mcp_server/utils.py index bad5f060fb8..ac969560e7c 100644 --- a/litellm/proxy/_experimental/mcp_server/utils.py +++ b/litellm/proxy/_experimental/mcp_server/utils.py @@ -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 diff --git a/tests/mcp_tests/test_mcp_server.py b/tests/mcp_tests/test_mcp_server.py index 031e26cf16b..7737bf73a20 100644 --- a/tests/mcp_tests/test_mcp_server.py +++ b/tests/mcp_tests/test_mcp_server.py @@ -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 = {}