diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 10e40b76efd..3c27a65f06a 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -195,6 +195,7 @@ class MCPServerManager: name=name_for_prefix, alias=alias, server_name=server_name, + spec_path=server_config.get("spec_path", None), url=server_config.get("url", None) or "", command=server_config.get("command", None) or "", args=server_config.get("args", None) or [], @@ -218,12 +219,163 @@ class MCPServerManager: access_groups=server_config.get("access_groups", None), ) self.config_mcp_servers[server_id] = new_server + + # Check if this is an OpenAPI-based server + spec_path = server_config.get("spec_path", None) + if spec_path: + verbose_logger.info( + f"Loading OpenAPI spec from {spec_path} for server {server_name}" + ) + self._register_openapi_tools( + spec_path=spec_path, + server=new_server, + base_url=server_config.get("url", ""), + ) + verbose_logger.debug( f"Loaded MCP Servers: {json.dumps(self.config_mcp_servers, indent=4, default=str)}" ) self.initialize_tool_name_to_mcp_server_name_mapping() + def _register_openapi_tools(self, spec_path: str, server: MCPServer, base_url: str): + """ + Register tools from an OpenAPI specification for a given server. + + This creates "virtual" MCP tools from OpenAPI endpoints that are: + 1. Registered in the global tool registry with server prefix + 2. Mapped to the server for routing + 3. Executed via the local tool handler + + Args: + spec_path: Path to the OpenAPI specification file + server: The MCPServer instance to register tools for + base_url: Base URL for API calls + """ + from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( + build_input_schema, + create_tool_function, + ) + from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( + get_base_url as get_openapi_base_url, + ) + from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( + load_openapi_spec, + ) + from litellm.proxy._experimental.mcp_server.tool_registry import ( + global_mcp_tool_registry, + ) + + try: + # Load OpenAPI spec + spec = load_openapi_spec(spec_path) + + # Use base_url from config if provided, otherwise extract from spec + if not base_url: + base_url = get_openapi_base_url(spec) + + verbose_logger.info( + f"Registering OpenAPI tools for server {server.name} with base URL: {base_url}" + ) + + # Get server prefix for tool naming + server_prefix = get_server_prefix(server) + + # Build headers from server configuration + headers = {} + + # Add authentication headers if configured + if server.authentication_token: + from litellm.types.mcp import MCPAuth + + if server.auth_type == MCPAuth.bearer_token: + headers["Authorization"] = f"Bearer {server.authentication_token}" + elif server.auth_type == MCPAuth.api_key: + headers["Authorization"] = f"ApiKey {server.authentication_token}" + elif server.auth_type == MCPAuth.basic: + headers["Authorization"] = f"Basic {server.authentication_token}" + + # Add any extra headers from server config + # Note: extra_headers is a List[str] of header names to forward, not a dict + # For OpenAPI tools, we'll just use the authentication headers + # If extra_headers were needed, they would be processed separately + + verbose_logger.debug( + f"Using headers for OpenAPI tools (excluding sensitive values): " + f"{list(headers.keys())}" + ) + + # Extract and register tools from OpenAPI paths + paths = spec.get("paths", {}) + registered_count = 0 + + verbose_logger.debug(f"Processing {len(paths)} paths from OpenAPI spec") + + for path, path_item in paths.items(): + for method in ["get", "post", "put", "delete", "patch"]: + if method not in path_item: + continue + + operation = path_item[method] + + # Generate tool name (without prefix initially) + operation_id = operation.get( + "operationId", f"{method}_{path.replace('/', '_')}" + ) + base_tool_name = operation_id.replace(" ", "_").lower() + + # Add server prefix to tool name + prefixed_tool_name = add_server_prefix_to_tool_name( + base_tool_name, server_prefix + ) + + # Get description + description = operation.get( + "summary", + operation.get("description", f"{method.upper()} {path}"), + ) + + # Build input schema using imported function + input_schema = build_input_schema(operation) + + # Create tool function with headers using imported function + tool_func = create_tool_function( + path, method, operation, base_url, headers=headers + ) + tool_func.__name__ = prefixed_tool_name + tool_func.__doc__ = description + + # Register tool with prefixed name in global registry + global_mcp_tool_registry.register_tool( + name=prefixed_tool_name, + description=description, + input_schema=input_schema, + handler=tool_func, + ) + + # Update tool name to server name mapping (for both prefixed and base names) + self.tool_name_to_mcp_server_name_mapping[base_tool_name] = ( + server_prefix + ) + self.tool_name_to_mcp_server_name_mapping[prefixed_tool_name] = ( + server_prefix + ) + + registered_count += 1 + verbose_logger.debug( + f"Registered OpenAPI tool: {prefixed_tool_name} for server {server.name}" + ) + + verbose_logger.info( + f"Successfully registered {registered_count} OpenAPI tools for server {server.name}" + ) + + except Exception as e: + verbose_logger.error( + f"Failed to register OpenAPI tools for server {server.name}: {str(e)}" + ) + raise e + def remove_server(self, mcp_server: LiteLLM_MCPServerTable): """ Remove a server from the registry @@ -469,6 +621,10 @@ class MCPServerManager: Returns: List[MCPTool]: List of tools available on the server with prefixed names """ + from litellm.proxy._experimental.mcp_server.tool_registry import ( + global_mcp_tool_registry, + ) + verbose_logger.debug(f"Connecting to url: {server.url}") verbose_logger.info(f"_get_tools_from_server for {server.name}...") @@ -481,7 +637,14 @@ class MCPServerManager: extra_headers=extra_headers, ) - tools = await self._fetch_tools_with_timeout(client, server.name) + ## HANDLE OPENAPI TOOLS + if server.spec_path: + _tools = global_mcp_tool_registry.list_tools(tool_prefix=server.name) + tools = global_mcp_tool_registry.convert_tools_to_mcp_sdk_tool_type( + _tools + ) + else: + tools = await self._fetch_tools_with_timeout(client, server.name) prefixed_or_original_tools = self._create_prefixed_tools( tools, server, add_prefix=add_prefix @@ -597,9 +760,15 @@ class MCPServerManager: Check if the tool is allowed or banned for the given server """ if server.allowed_tools: - return tool_name in server.allowed_tools + return ( + tool_name in server.allowed_tools + or f"{server.name}-{tool_name}" in server.allowed_tools + ) if server.disallowed_tools: - return tool_name not in server.disallowed_tools + return ( + tool_name not in server.disallowed_tools + and f"{server.name}-{tool_name}" not in server.disallowed_tools + ) return True async def check_tool_permission_for_key_team( @@ -621,18 +790,20 @@ class MCPServerManager: Raises: HTTPException: If tool is not allowed for this key/team """ - from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import MCPRequestHandler - + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( + MCPRequestHandler, + ) + if not user_api_key_auth: return - + # Check if tool is allowed is_allowed = await MCPRequestHandler.is_tool_allowed_for_server( tool_name=tool_name, server_id=server.server_id, user_api_key_auth=user_api_key_auth, ) - + if not is_allowed: raise HTTPException( status_code=403, @@ -641,6 +812,64 @@ class MCPServerManager: }, ) + async def _call_openapi_tool_handler( + self, + server: MCPServer, + tool_name: str, + arguments: Dict[str, Any], + ) -> CallToolResult: + """ + Call an OpenAPI tool handler directly. + + For OpenAPI servers, instead of using MCP protocol, we call the tool handler + that was registered during OpenAPI spec parsing. This handler makes direct + HTTP requests to the API. + + Args: + tool_name: The full tool name (with prefix) to call + arguments: Tool arguments to pass to the handler + + Returns: + CallToolResult with the response from the API + """ + from mcp.types import TextContent + + from litellm.proxy._experimental.mcp_server.tool_registry import ( + global_mcp_tool_registry, + ) + + # Get the tool from the registry + tool = global_mcp_tool_registry.get_tool(f"{server.name}-{tool_name}") + if tool is None: + # Tool not found in registry + error_msg = f"OpenAPI tool {tool_name} not found in registry" + verbose_logger.error(error_msg) + return CallToolResult( + content=[TextContent(type="text", text=error_msg)], + isError=True, + ) + + try: + # Call the tool handler with the arguments + # The handler is an async function that makes the HTTP request + handler_result = await tool.handler(**arguments) + + # Convert the handler result (string response) to CallToolResult format + result = CallToolResult( + content=[TextContent(type="text", text=str(handler_result))], + isError=False, + ) + + return result + + except Exception as e: + error_msg = f"Error calling OpenAPI tool {tool_name}: {str(e)}" + verbose_logger.error(error_msg) + return CallToolResult( + content=[TextContent(type="text", text=error_msg)], + isError=True, + ) + async def pre_call_tool_check( self, name: str, @@ -793,95 +1022,109 @@ class MCPServerManager: server=mcp_server, ) - # Get server-specific auth header if available - server_auth_header: Optional[Union[Dict[str, str], str]] = None - if mcp_server_auth_headers and mcp_server.alias: - server_auth_header = mcp_server_auth_headers.get(mcp_server.alias) - elif mcp_server_auth_headers and mcp_server.server_name: - server_auth_header = mcp_server_auth_headers.get(mcp_server.server_name) + # Prepare tasks for during hooks + tasks = [] + if proxy_logging_obj: + # Create synthetic LLM data for during hook processing + from litellm.types.llms.base import HiddenParams + from litellm.types.mcp import MCPDuringCallRequestObject - # Fall back to deprecated mcp_auth_header if no server-specific header found - if server_auth_header is None: - server_auth_header = mcp_auth_header - - # oauth2 headers - extra_headers: Optional[Dict[str, str]] = None - if mcp_server.auth_type == MCPAuth.oauth2: - extra_headers = oauth2_headers - - if mcp_server.extra_headers and raw_headers: - if extra_headers is None: - extra_headers = {} - for header in mcp_server.extra_headers: - if header in raw_headers: - extra_headers[header] = raw_headers[header] - - client = self._create_mcp_client( - server=mcp_server, - mcp_auth_header=server_auth_header, - extra_headers=extra_headers, - ) - - async with client: - # Use the original tool name (without prefix) for the actual call - call_tool_params = MCPCallToolRequestParams( - name=original_tool_name, + request_obj = MCPDuringCallRequestObject( + tool_name=name, arguments=arguments, + server_name=server_name_from_prefix, + start_time=start_time.timestamp() if start_time else None, + hidden_params=HiddenParams(), ) - tasks = [] - if proxy_logging_obj: - # Create synthetic LLM data for during hook processing - from litellm.types.llms.base import HiddenParams - from litellm.types.mcp import MCPDuringCallRequestObject - request_obj = MCPDuringCallRequestObject( - tool_name=name, + during_hook_kwargs = { + "name": name, + "arguments": arguments, + "server_name": server_name_from_prefix, + "user_api_key_auth": user_api_key_auth, + } + + synthetic_llm_data = proxy_logging_obj._convert_mcp_to_llm_format( + request_obj, during_hook_kwargs + ) + + during_hook_task = asyncio.create_task( + proxy_logging_obj.during_call_hook( + user_api_key_dict=user_api_key_auth, + data=synthetic_llm_data, + call_type="mcp_call", # type: ignore + ) + ) + tasks.append(during_hook_task) + + # For OpenAPI servers, call the tool handler directly instead of via MCP client + if mcp_server.spec_path: + verbose_logger.debug( + f"Calling OpenAPI tool {name} directly via HTTP handler" + ) + tasks.append( + asyncio.create_task( + self._call_openapi_tool_handler(mcp_server, name, arguments) + ) + ) + else: + # For regular MCP servers, use the MCP client + # Get server-specific auth header if available + server_auth_header: Optional[Union[Dict[str, str], str]] = None + if mcp_server_auth_headers and mcp_server.alias: + server_auth_header = mcp_server_auth_headers.get(mcp_server.alias) + elif mcp_server_auth_headers and mcp_server.server_name: + server_auth_header = mcp_server_auth_headers.get(mcp_server.server_name) + + # Fall back to deprecated mcp_auth_header if no server-specific header found + if server_auth_header is None: + server_auth_header = mcp_auth_header + + # oauth2 headers + extra_headers: Optional[Dict[str, str]] = None + if mcp_server.auth_type == MCPAuth.oauth2: + extra_headers = oauth2_headers + + if mcp_server.extra_headers and raw_headers: + if extra_headers is None: + extra_headers = {} + for header in mcp_server.extra_headers: + if header in raw_headers: + extra_headers[header] = raw_headers[header] + + client = self._create_mcp_client( + server=mcp_server, + mcp_auth_header=server_auth_header, + extra_headers=extra_headers, + ) + + async with client: + # Use the original tool name (without prefix) for the actual call + call_tool_params = MCPCallToolRequestParams( + name=original_tool_name, arguments=arguments, - server_name=server_name_from_prefix, - start_time=start_time.timestamp() if start_time else None, - hidden_params=HiddenParams(), ) + tasks.append(asyncio.create_task(client.call_tool(call_tool_params))) - during_hook_kwargs = { - "name": name, - "arguments": arguments, - "server_name": server_name_from_prefix, - "user_api_key_auth": user_api_key_auth, - } + try: + mcp_responses = await asyncio.gather(*tasks) - synthetic_llm_data = proxy_logging_obj._convert_mcp_to_llm_format( - request_obj, during_hook_kwargs - ) + # If proxy_logging_obj is None, the tool call result is at index 0 + # If proxy_logging_obj is not None, the tool call result is at index 1 (after the during hook task) + result_index = 1 if proxy_logging_obj else 0 + result = mcp_responses[result_index] - during_hook_task = asyncio.create_task( - proxy_logging_obj.during_call_hook( - user_api_key_dict=user_api_key_auth, - data=synthetic_llm_data, - call_type="mcp_call", # type: ignore - ) - ) - tasks.append(during_hook_task) - - tasks.append(asyncio.create_task(client.call_tool(call_tool_params))) - try: - mcp_responses = await asyncio.gather(*tasks) - - # If proxy_logging_obj is None, the tool call result is at index 0 - # If proxy_logging_obj is not None, the tool call result is at index 1 (after the during hook task) - result_index = 1 if proxy_logging_obj else 0 - result = mcp_responses[result_index] - - return cast(CallToolResult, result) - except ( - BlockedPiiEntityError, - GuardrailRaisedException, - HTTPException, - ) as e: - # Re-raise guardrail exceptions to properly fail the MCP call - verbose_logger.error( - f"Guardrail blocked MCP tool call during result check: {str(e)}" - ) - raise e + return cast(CallToolResult, result) + except ( + BlockedPiiEntityError, + GuardrailRaisedException, + HTTPException, + ) as e: + # Re-raise guardrail exceptions to properly fail the MCP call + verbose_logger.error( + f"Guardrail blocked MCP tool call during result check: {str(e)}" + ) + raise e ######################################################### # End of Methods that call the upstream MCP servers diff --git a/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py b/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py new file mode 100644 index 00000000000..72288f8e673 --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py @@ -0,0 +1,236 @@ +""" +This module is used to generate MCP tools from OpenAPI specs. +""" + +import json +from typing import Any, Dict, Optional + +import httpx + +from litellm._logging import verbose_logger +from litellm.proxy._experimental.mcp_server.tool_registry import ( + global_mcp_tool_registry, +) + +# Store the base URL and headers globally +BASE_URL = "" +HEADERS: Dict[str, str] = {} + + +def load_openapi_spec(filepath: str) -> Dict[str, Any]: + """Load OpenAPI specification from JSON file.""" + with open(filepath, "r") as f: + return json.load(f) + + +def get_base_url(spec: Dict[str, Any]) -> str: + """Extract base URL from OpenAPI spec.""" + # OpenAPI 3.x + if "servers" in spec and spec["servers"]: + return spec["servers"][0]["url"] + # OpenAPI 2.x (Swagger) + elif "host" in spec: + scheme = spec.get("schemes", ["https"])[0] + base_path = spec.get("basePath", "") + return f"{scheme}://{spec['host']}{base_path}" + return "" + + +def extract_parameters(operation: Dict[str, Any]) -> tuple: + """Extract parameter names from OpenAPI operation.""" + path_params = [] + query_params = [] + body_params = [] + + # OpenAPI 3.x and 2.x parameters + if "parameters" in operation: + for param in operation["parameters"]: + param_name = param["name"] + if param.get("in") == "path": + path_params.append(param_name) + elif param.get("in") == "query": + query_params.append(param_name) + elif param.get("in") == "body": + body_params.append(param_name) + + # OpenAPI 3.x requestBody + if "requestBody" in operation: + body_params.append("body") + + return path_params, query_params, body_params + + +def build_input_schema(operation: Dict[str, Any]) -> Dict[str, Any]: + """Build MCP input schema from OpenAPI operation.""" + properties = {} + required = [] + + # Process parameters + if "parameters" in operation: + for param in operation["parameters"]: + param_name = param["name"] + param_schema = param.get("schema", {}) + param_type = param_schema.get("type", "string") + + properties[param_name] = { + "type": param_type, + "description": param.get("description", ""), + } + + if param.get("required", False): + required.append(param_name) + + # Process requestBody (OpenAPI 3.x) + if "requestBody" in operation: + request_body = operation["requestBody"] + content = request_body.get("content", {}) + + # Try to get JSON schema + if "application/json" in content: + schema = content["application/json"].get("schema", {}) + properties["body"] = { + "type": "object", + "description": request_body.get("description", "Request body"), + "properties": schema.get("properties", {}), + } + if request_body.get("required", False): + required.append("body") + + return { + "type": "object", + "properties": properties, + "required": required if required else [], + } + + +def create_tool_function( + path: str, + method: str, + operation: Dict[str, Any], + base_url: str, + headers: Optional[Dict[str, str]] = None, +): + """Create a tool function for an OpenAPI operation. + + Args: + path: API endpoint path + method: HTTP method (get, post, put, delete, patch) + operation: OpenAPI operation object + base_url: Base URL for the API + headers: Optional headers to include in requests (e.g., authentication) + """ + if headers is None: + headers = {} + + path_params, query_params, body_params = extract_parameters(operation) + all_params = path_params + query_params + body_params + + # Build function signature dynamically + if all_params: + params_str = ", ".join(f"{p}: str = ''" for p in all_params) + else: + params_str = "" + + # Create the function code as a string + func_code = f''' +async def tool_function({params_str}) -> str: + """Dynamically generated tool function.""" + url = base_url + path + + # Replace path parameters + path_param_names = {path_params} + for param_name in path_param_names: + param_value = locals().get(param_name, "") + if param_value: + url = url.replace("{{" + param_name + "}}", str(param_value)) + + # Build query params + query_param_names = {query_params} + params = {{}} + for param_name in query_param_names: + param_value = locals().get(param_name, "") + if param_value: + params[param_name] = param_value + + # Build request body + body_param_names = {body_params} + json_body = None + if body_param_names: + body_value = locals().get("body", {{}}) + if isinstance(body_value, dict): + json_body = body_value + elif body_value: + # If it's a string, try to parse as JSON + import json as json_module + try: + json_body = json_module.loads(body_value) if isinstance(body_value, str) else {{"data": body_value}} + except: + json_body = {{"data": body_value}} + + # Make HTTP request + async with httpx.AsyncClient() as client: + if "{method.lower()}" == "get": + response = await client.get(url, params=params, headers=headers) + elif "{method.lower()}" == "post": + response = await client.post(url, params=params, json=json_body, headers=headers) + elif "{method.lower()}" == "put": + response = await client.put(url, params=params, json=json_body, headers=headers) + elif "{method.lower()}" == "delete": + response = await client.delete(url, params=params, headers=headers) + elif "{method.lower()}" == "patch": + response = await client.patch(url, params=params, json=json_body, headers=headers) + else: + return "Unsupported HTTP method: {method}" + + return response.text +''' + + # Execute the function code to create the actual function + local_vars = { + "httpx": httpx, + "headers": headers, + "base_url": base_url, + "path": path, + "method": method, + } + exec(func_code, local_vars) + + return local_vars["tool_function"] + + +def register_tools_from_openapi(spec: Dict[str, Any], base_url: str): + """Register MCP tools from OpenAPI specification.""" + paths = spec.get("paths", {}) + + for path, path_item in paths.items(): + for method in ["get", "post", "put", "delete", "patch"]: + if method in path_item: + operation = path_item[method] + + # Generate tool name + operation_id = operation.get( + "operationId", f"{method}_{path.replace('/', '_')}" + ) + tool_name = operation_id.replace(" ", "_").lower() + + # Get description + description = operation.get( + "summary", operation.get("description", f"{method.upper()} {path}") + ) + + # Build input schema + input_schema = build_input_schema(operation) + + # Create tool function + tool_func = create_tool_function(path, method, operation, base_url) + tool_func.__name__ = tool_name + tool_func.__doc__ = description + + # Register tool with local registry + global_mcp_tool_registry.register_tool( + name=tool_name, + description=description, + input_schema=input_schema, + handler=tool_func, + ) + verbose_logger.debug(f"Registered tool: {tool_name}") diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index d7ebfb805f3..37157b4e0ac 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -364,25 +364,25 @@ if MCP_AVAILABLE: def _tool_name_matches(tool_name: str, filter_list: List[str]) -> bool: """ Check if a tool name matches any name in the filter list. - + Checks both the full tool name and unprefixed version (without server prefix). This allows users to configure simple tool names regardless of prefixing. - + Args: tool_name: The tool name to check (may be prefixed like "server-tool_name") filter_list: List of tool names to match against - + Returns: True if the tool name (prefixed or unprefixed) is in the filter list """ from litellm.proxy._experimental.mcp_server.utils import ( get_server_name_prefix_tool_mcp, ) - + # Check if the full name is in the list if tool_name in filter_list: return True - + # Check if the unprefixed name is in the list unprefixed_name, _ = get_server_name_prefix_tool_mcp(tool_name) return unprefixed_name in filter_list @@ -393,34 +393,36 @@ if MCP_AVAILABLE: ) -> List[MCPTool]: """ Filter tools by allowed/disallowed tools configuration. - + If allowed_tools is set, only tools in that list are returned. If disallowed_tools is set, tools in that list are excluded. Tool names are matched with and without server prefixes for flexibility. - + Args: tools: List of tools to filter mcp_server: Server configuration with allowed_tools/disallowed_tools - + Returns: Filtered list of tools """ tools_to_return = tools - + # Filter by allowed_tools (whitelist) if mcp_server.allowed_tools: tools_to_return = [ - tool for tool in tools + tool + for tool in tools if _tool_name_matches(tool.name, mcp_server.allowed_tools) ] - + # Filter by disallowed_tools (blacklist) if mcp_server.disallowed_tools: tools_to_return = [ - tool for tool in tools_to_return + tool + for tool in tools_to_return if not _tool_name_matches(tool.name, mcp_server.disallowed_tools) ] - + return tools_to_return async def _get_tools_from_mcp_servers( @@ -497,17 +499,17 @@ if MCP_AVAILABLE: extra_headers=extra_headers, add_prefix=add_prefix, ) - + filtered_tools = filter_tools_by_allowed_tools(tools, server) - + filtered_tools = await filter_tools_by_key_team_permissions( tools=filtered_tools, server_id=server_id, user_api_key_auth=user_api_key_auth, ) - + all_tools.extend(filtered_tools) - + verbose_logger.debug( f"Successfully fetched {len(tools)} tools from server {server.name}, {len(filtered_tools)} after filtering" ) @@ -520,6 +522,7 @@ if MCP_AVAILABLE: verbose_logger.info( f"Successfully fetched {len(all_tools)} tools total from all MCP servers" ) + return all_tools async def filter_tools_by_key_team_permissions( @@ -529,7 +532,7 @@ if MCP_AVAILABLE: ) -> List[MCPTool]: """ Filter tools based on key/team mcp_tool_permissions. - + Note: Tool names in the DB are stored without server prefixes, but tool names from MCP servers are prefixed. We need to strip the prefix before comparing. @@ -551,7 +554,7 @@ if MCP_AVAILABLE: else: # No restrictions, return all tools filtered_tools = tools - + return filtered_tools async def _list_mcp_tools( @@ -596,30 +599,7 @@ if MCP_AVAILABLE: ) # Continue with empty managed tools list instead of failing completely - # Get tools from local registry - local_tools = [] - try: - local_tools_raw = global_mcp_tool_registry.list_tools() - - # Convert local tools to MCPTool format - for tool in local_tools_raw: - # Convert from litellm.types.mcp_server.tool_registry.MCPTool to mcp.types.Tool - mcp_tool = MCPTool( - name=tool.name, - description=tool.description, - inputSchema=tool.input_schema, - ) - local_tools.append(mcp_tool) - except Exception as e: - verbose_logger.exception( - f"Error getting tools from local registry: {str(e)}" - ) - # Continue with empty local tools list instead of failing completely - - # Combine all tools - all_tools = managed_tools + local_tools - - return all_tools + return managed_tools @client async def call_mcp_tool( @@ -680,33 +660,42 @@ if MCP_AVAILABLE: standard_logging_mcp_tool_call ) litellm_logging_obj.model = f"MCP: {name}" - # Try managed server tool first (pass the full prefixed name) - # Primary and recommended way to use MCP servers + # Check if tool exists in local registry first (for OpenAPI-based tools) + # These tools are registered with their prefixed names ######################################################### - 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) - arguments=arguments, - user_api_key_auth=user_api_key_auth, - mcp_auth_header=mcp_auth_header, - mcp_server_auth_headers=mcp_server_auth_headers, - oauth2_headers=oauth2_headers, - raw_headers=raw_headers, - litellm_logging_obj=litellm_logging_obj, - ) + local_tool = global_mcp_tool_registry.get_tool(name) + if local_tool: + verbose_logger.debug(f"Executing local registry tool: {name}") + response = await _handle_local_mcp_tool(name, arguments) - # Fall back to local tool registry (use original name) - ######################################################### - # Deprecated: Local MCP Server Tool + # Try managed MCP server tool (pass the full prefixed name) + # Primary and recommended way to use external MCP servers ######################################################### else: - response = await _handle_local_mcp_tool(original_tool_name, arguments) + 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) + arguments=arguments, + user_api_key_auth=user_api_key_auth, + mcp_auth_header=mcp_auth_header, + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + litellm_logging_obj=litellm_logging_obj, + ) + + # Fall back to local tool registry with original name (legacy support) + ######################################################### + # Deprecated: Local MCP Server Tool + ######################################################### + else: + response = await _handle_local_mcp_tool(original_tool_name, arguments) ######################################################### # Post MCP Tool Call Hook @@ -778,14 +767,21 @@ if MCP_AVAILABLE: Handle tool execution for local registry tools Note: Local tools don't use prefixes, so we use the original name """ + import inspect + tool = global_mcp_tool_registry.get_tool(name) if not tool: raise HTTPException(status_code=404, detail=f"Tool '{name}' not found") try: - result = tool.handler(**arguments) + # Check if handler is async or sync + if inspect.iscoroutinefunction(tool.handler): + result = await tool.handler(**arguments) + else: + result = tool.handler(**arguments) return [TextContent(text=str(result), type="text")] except Exception as e: + verbose_logger.exception(f"Error executing local tool {name}: {str(e)}") return [TextContent(text=f"Error: {str(e)}", type="text")] def _get_mcp_servers_in_path(path: str) -> Optional[List[str]]: diff --git a/litellm/proxy/_experimental/mcp_server/tool_registry.py b/litellm/proxy/_experimental/mcp_server/tool_registry.py index c08b7979683..bc69095fc43 100644 --- a/litellm/proxy/_experimental/mcp_server/tool_registry.py +++ b/litellm/proxy/_experimental/mcp_server/tool_registry.py @@ -1,6 +1,8 @@ import json from typing import Any, Callable, Dict, List, Optional +from mcp.types import Tool as MCPToolSDKTool + from litellm._logging import verbose_logger from litellm.proxy.types_utils.utils import get_instance_fn from litellm.types.mcp_server.tool_registry import MCPTool @@ -39,12 +41,30 @@ class MCPToolRegistry: """ return self.tools.get(name) - def list_tools(self) -> List[MCPTool]: + def list_tools(self, tool_prefix: Optional[str] = None) -> List[MCPTool]: """ List all registered tools """ + if tool_prefix: + return [ + tool + for tool in self.tools.values() + if tool.name.startswith(tool_prefix) + ] return list(self.tools.values()) + def convert_tools_to_mcp_sdk_tool_type( + self, tools: List[MCPTool] + ) -> List[MCPToolSDKTool]: + return [ + MCPToolSDKTool( + name=tool.name, + description=tool.description, + inputSchema=tool.input_schema, + ) + for tool in tools + ] + def load_tools_from_config( self, mcp_tools_config: Optional[Dict[str, Any]] = None ) -> None: diff --git a/litellm/proxy/_new_secret_config.yaml b/litellm/proxy/_new_secret_config.yaml index 52a6fc16ff1..95c6869b95b 100644 --- a/litellm/proxy/_new_secret_config.yaml +++ b/litellm/proxy/_new_secret_config.yaml @@ -1,6 +1,29 @@ model_list: - model_name: gpt-5-mini litellm_params: - model: azure/gpt-5-mini-2 - api_key: os.environ/AZURE_API_KEY_ALT - api_base: os.environ/AZURE_API_BASE_ALT + model: openai/gpt-4o-mini + api_base: "https://webhook.site/2f385e05-00aa-402b-86d1-efc9261471a5" + api_key: dummy + - model_name: "byok-wildcard/*" + litellm_params: + model: openai/* + - model_name: xai-grok-3 + litellm_params: + model: xai/grok-3 + - model_name: hosted_vllm/whisper-v3 + litellm_params: + model: hosted_vllm/whisper-v3 + api_base: "https://webhook.site/2f385e05-00aa-402b-86d1-efc9261471a5" + api_key: dummy + +mcp_servers: + my_api_mcp: + url: "http://0.0.0.0:8090" + spec_path: "/Users/krrishdholakia/Documents/temp_py_folder/example_openapi.json" + auth_type: none + allowed_tools: ["getpetbyid", "my_api_mcp-findpetsbystatus"] + + +litellm_settings: + callbacks: ["prometheus"] + custom_prometheus_metadata_labels: ["metadata.initiative", "metadata.business-unit"] diff --git a/litellm/types/mcp_server/mcp_server_manager.py b/litellm/types/mcp_server/mcp_server_manager.py index 3e0c2b20e39..a247dc3f614 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -17,6 +17,7 @@ class MCPServer(BaseModel): server_name: Optional[str] = None url: Optional[str] = None transport: MCPTransportType + spec_path: Optional[str] = None auth_type: Optional[MCPAuthType] = None authentication_token: Optional[str] = None mcp_info: Optional[MCPInfo] = None diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 66871aefc1c..769785b9214 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -1001,7 +1001,7 @@ async def test_list_tools_with_no_tool_permissions_shows_all(): async def test_list_tools_strips_prefix_when_matching_permissions(): """ Test that tool permission filtering correctly strips prefixes from tool names. - + Tools from MCP servers are prefixed (e.g., "GITMCP-fetch_litellm_documentation"), but allowed tools in DB are stored without prefix (e.g., "fetch_litellm_documentation"). The filtering should strip the prefix before comparing. @@ -1056,7 +1056,9 @@ async def test_list_tools_strips_prefix_when_matching_permissions(): tool1.inputSchema = {} tool2 = MagicMock() - tool2.name = "GITMCP-search_litellm_documentation" # Prefixed, not in allowed list + tool2.name = ( + "GITMCP-search_litellm_documentation" # Prefixed, not in allowed list + ) tool2.description = "Search docs" tool2.inputSchema = {} @@ -1093,3 +1095,76 @@ async def test_list_tools_strips_prefix_when_matching_permissions(): "GITMCP-fetch_litellm_documentation", "GITMCP-search_litellm_code", ] + + +def test_filter_tools_by_allowed_tools(): + """Test that filter_tools_by_allowed_tools filters tools correctly""" + from mcp.types import Tool + + from litellm.proxy._experimental.mcp_server.server import ( + filter_tools_by_allowed_tools, + ) + from litellm.types.mcp import MCPTransport + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + mcp_server = MCPServer( + server_id="my_api_mcp", + name="my_api_mcp", + alias="my_api_mcp", + transport=MCPTransport.http, + allowed_tools=["getpetbyid", "my_api_mcp-findpetsbystatus"], + disallowed_tools=None, + ) + tools_to_return = [ + Tool( + name="my_api_mcp-getpetbyid", + title=None, + description="Find pet by ID", + inputSchema={ + "type": "object", + "properties": {"petId": {"type": "integer", "description": ""}}, + "required": ["petId"], + }, + outputSchema=None, + annotations=None, + ), + Tool( + name="my_api_mcp-findpetsbystatus", + title=None, + description="Finds Pets by status", + inputSchema={ + "type": "object", + "properties": {"status": {"type": "string", "description": ""}}, + "required": ["status"], + }, + outputSchema=None, + annotations=None, + ), + Tool( + name="my_api_mcp-addpet", + title=None, + description="Add a new pet to the store", + inputSchema={ + "type": "object", + "properties": { + "body": { + "type": "object", + "description": "Request body", + "properties": { + "name": {"type": "string"}, + "status": {"type": "string"}, + }, + } + }, + "required": ["body"], + }, + outputSchema=None, + annotations=None, + ), + ] + + filtered_tools = filter_tools_by_allowed_tools(tools_to_return, mcp_server) + + assert len(filtered_tools) == 2 + assert filtered_tools[0].name == "my_api_mcp-getpetbyid" + assert filtered_tools[1].name == "my_api_mcp-findpetsbystatus" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 1bf5c0abb8a..0476f290a09 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -1150,6 +1150,78 @@ class TestMCPServerManager: user_api_key_auth=user_auth, ) + @pytest.mark.asyncio + async def test_allowed_tools_with_mixed_prefixed_and_unprefixed_names(self): + """ + Test that allowed_tools works with both unprefixed and prefixed tool names. + This tests the scenario where allowed_tools = ["getpetbyid", "my_api_mcp-findpetsbystatus"] + Both getpetbyid (unprefixed) and findpetsbystatus (called unprefixed but allowed via prefix) should work. + """ + manager = MCPServerManager() + + # Create server with mixed prefixed/unprefixed allowed_tools + server = MCPServer( + server_id="my_api_mcp", + name="my_api_mcp", + transport=MCPTransport.stdio, + allowed_tools=["getpetbyid", "my_api_mcp-findpetsbystatus"], + disallowed_tools=None, + ) + + # Mock dependencies - set object_permission and object_permission_id to None + # so permission checks return None (no restrictions) + user_api_key_auth = MagicMock() + user_api_key_auth.object_permission = None + user_api_key_auth.object_permission_id = None + proxy_logging_obj = MagicMock() + + # Mock the async methods that pre_call_tool_check calls + proxy_logging_obj._create_mcp_request_object_from_kwargs = MagicMock( + return_value={} + ) + proxy_logging_obj._convert_mcp_to_llm_format = MagicMock(return_value={}) + proxy_logging_obj.pre_call_hook = AsyncMock(return_value={}) + + # Test 1: Call getpetbyid (unprefixed in allowed_tools) - should succeed + await manager.pre_call_tool_check( + name="getpetbyid", + arguments={"petId": "1"}, + server_name_from_prefix="my_api_mcp", + user_api_key_auth=user_api_key_auth, + proxy_logging_obj=proxy_logging_obj, + server=server, + ) + + # Test 2: Call findpetsbystatus (prefixed in allowed_tools as "my_api_mcp-findpetsbystatus") - should succeed + await manager.pre_call_tool_check( + name="findpetsbystatus", + arguments={"status": "available"}, + server_name_from_prefix="my_api_mcp", + user_api_key_auth=user_api_key_auth, + proxy_logging_obj=proxy_logging_obj, + server=server, + ) + + # Test 3: Call a tool that's not in allowed_tools - should fail + with pytest.raises(HTTPException) as exc_info: + await manager.pre_call_tool_check( + name="deletepet", + arguments={"petId": "1"}, + server_name_from_prefix="my_api_mcp", + user_api_key_auth=user_api_key_auth, + proxy_logging_obj=proxy_logging_obj, + server=server, + ) + + assert exc_info.value.status_code == 403 + assert ( + "Tool deletepet is not allowed for server my_api_mcp" + in exc_info.value.detail["error"] + ) + assert ( + "Contact proxy admin to allow this tool" in exc_info.value.detail["error"] + ) + if __name__ == "__main__": pytest.main([__file__]) diff --git a/ui/litellm-dashboard/src/components/admins.tsx b/ui/litellm-dashboard/src/components/admins.tsx index 7c891c71dd0..88ec208cb4d 100644 --- a/ui/litellm-dashboard/src/components/admins.tsx +++ b/ui/litellm-dashboard/src/components/admins.tsx @@ -524,9 +524,7 @@ const AdminPanel: React.FC = ({