From ae1fc804450eff7d2793291b232e5e2c0a0ac955 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Fri, 6 Feb 2026 14:44:17 -0800 Subject: [PATCH] feat(mcp/): support custom code guardrails for mcp calls allows custom code guardrails to work on mcp input --- .../mcp_server/mcp_server_manager.py | 248 ++++-------------- .../mcp_server/rest_endpoints.py | 213 +++------------ litellm/proxy/common_request_processing.py | 1 + .../custom_code/custom_code_guardrail.py | 4 +- .../unified_guardrail/unified_guardrail.py | 1 - .../custom_code/CustomCodeModal.tsx | 42 ++- .../components/playground/chat_ui/ChatUI.tsx | 1 - 7 files changed, 129 insertions(+), 381 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 1d7c0eaa116..f5b0f152fd8 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -14,8 +14,6 @@ import re from typing import Any, Callable, Dict, List, Literal, Optional, Set, Tuple, Union, cast from urllib.parse import urlparse -import anyio - from fastapi import HTTPException from httpx import HTTPStatusError from mcp import ReadResourceResult, Resource @@ -38,7 +36,6 @@ from litellm.llms.custom_httpx.http_handler import get_async_httpx_client from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( MCPRequestHandler, ) -from litellm.proxy._experimental.mcp_server.oauth2_token_cache import resolve_mcp_auth from litellm.proxy._experimental.mcp_server.utils import ( MCP_TOOL_PREFIX_SEPARATOR, add_server_prefix_to_name, @@ -56,7 +53,6 @@ from litellm.proxy._types import ( MCPTransportType, UserAPIKeyAuth, ) -from litellm.proxy.auth.ip_address_utils import IPAddressUtils from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper from litellm.proxy.utils import ProxyLogging from litellm.types.llms.custom_http import httpxSpecialProvider @@ -70,7 +66,7 @@ from litellm.types.utils import CallTypes try: from mcp.shared.tool_name_validation import ( - validate_tool_name, # pyright: ignore[reportAssignmentType] + validate_tool_name, # type: ignore[reportAssignmentType] ) from mcp.shared.tool_name_validation import SEP_986_URL except ImportError: @@ -78,12 +74,12 @@ except ImportError: SEP_986_URL = "https://github.com/modelcontextprotocol/protocol/blob/main/proposals/0001-tool-name-validation.md" - class _ToolNameValidationResult(BaseModel): + class ToolNameValidationResult(BaseModel): is_valid: bool = True warnings: list = [] - def validate_tool_name(name: str) -> _ToolNameValidationResult: # type: ignore[misc] - return _ToolNameValidationResult() + def validate_tool_name(name: str) -> ToolNameValidationResult: # type: ignore[misc] + return ToolNameValidationResult() # Probe includes characters on both sides of the separator to mimic real prefixed tool names. @@ -329,9 +325,6 @@ class MCPServerManager: access_groups=server_config.get("access_groups", None), static_headers=server_config.get("static_headers", None), allow_all_keys=bool(server_config.get("allow_all_keys", False)), - available_on_public_internet=bool( - server_config.get("available_on_public_internet", False) - ), ) self.config_mcp_servers[server_id] = new_server @@ -341,7 +334,7 @@ class MCPServerManager: verbose_logger.info( f"Loading OpenAPI spec from {spec_path} for server {server_name}" ) - await self._register_openapi_tools( + self._register_openapi_tools( spec_path=spec_path, server=new_server, base_url=server_config.get("url", ""), @@ -353,9 +346,7 @@ class MCPServerManager: self.initialize_tool_name_to_mcp_server_name_mapping() - async def _register_openapi_tools( - self, spec_path: str, server: MCPServer, base_url: str - ): + def _register_openapi_tools(self, spec_path: str, server: MCPServer, base_url: str): """ Register tools from an OpenAPI specification for a given server. @@ -377,15 +368,15 @@ class MCPServerManager: get_base_url as get_openapi_base_url, ) from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( - load_openapi_spec_async, + load_openapi_spec, ) from litellm.proxy._experimental.mcp_server.tool_registry import ( global_mcp_tool_registry, ) try: - # Load OpenAPI spec (async to avoid "called from within a running event loop") - spec = await load_openapi_spec_async(spec_path) + # Load OpenAPI spec + spec = load_openapi_spec(spec_path) # Use base_url from config if provided, otherwise extract from spec if not base_url: @@ -632,9 +623,6 @@ class MCPServerManager: allowed_tools=getattr(mcp_server, "allowed_tools", None), disallowed_tools=getattr(mcp_server, "disallowed_tools", None), allow_all_keys=mcp_server.allow_all_keys, - available_on_public_internet=bool( - getattr(mcp_server, "available_on_public_internet", False) - ), updated_at=getattr(mcp_server, "updated_at", None), ) return new_server @@ -673,47 +661,24 @@ class MCPServerManager: return [ server.server_id for server in self.get_registry().values() - if server.allow_all_keys is True + if server.allow_all_keys ] async def get_allowed_mcp_servers( self, user_api_key_auth: Optional[UserAPIKeyAuth] = None ) -> List[str]: """ - Get the allowed MCP Servers for the user. - - Priority: - 1. If object_permission.mcp_servers is explicitly set, use it (even for admins) - 2. If admin and no object_permission, return all servers - 3. Otherwise, use standard permission checks + Get the allowed MCP Servers for the user """ from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view + # If admin, get all servers + if user_api_key_auth and _user_has_admin_view(user_api_key_auth): + return list(self.get_registry().keys()) + allow_all_server_ids = self.get_allow_all_keys_server_ids() try: - # Check if object_permission.mcp_servers is explicitly set - has_explicit_object_permission = False - if user_api_key_auth and user_api_key_auth.object_permission: - # Check if mcp_servers is explicitly set (not None, empty list is valid) - if user_api_key_auth.object_permission.mcp_servers is not None: - has_explicit_object_permission = True - verbose_logger.debug( - f"Object permission mcp_servers explicitly set: {user_api_key_auth.object_permission.mcp_servers}" - ) - - # If admin but NO explicit object permission, get all servers - if ( - user_api_key_auth - and _user_has_admin_view(user_api_key_auth) - and not has_explicit_object_permission - ): - verbose_logger.debug( - "Admin user without explicit object_permission - returning all servers" - ) - return list(self.get_registry().keys()) - - # Get allowed servers from object permissions (respects object_permission even for admins) allowed_mcp_servers = await MCPRequestHandler.get_allowed_mcp_servers( user_api_key_auth ) @@ -732,23 +697,6 @@ class MCPServerManager: verbose_logger.warning(f"Failed to get allowed MCP servers: {str(e)}.") return allow_all_server_ids - def filter_server_ids_by_ip( - self, server_ids: List[str], client_ip: Optional[str] - ) -> List[str]: - """ - Filter server IDs by client IP — external callers only see public servers. - - Returns server_ids unchanged when client_ip is None (no filtering). - """ - if client_ip is None: - return server_ids - return [ - sid - for sid in server_ids - if (s := self.get_mcp_server_by_id(sid)) is not None - and self._is_server_accessible_from_ip(s, client_ip) - ] - async def get_tools_for_server(self, server_id: str) -> List[MCPTool]: """ Get the tools for a given server @@ -859,7 +807,7 @@ class MCPServerManager: return resolved_env - async def _create_mcp_client( + def _create_mcp_client( self, server: MCPServer, mcp_auth_header: Optional[Union[str, Dict[str, str]]] = None, @@ -869,22 +817,13 @@ class MCPServerManager: """ Create an MCPClient instance for the given server. - Auth resolution (single place for all auth logic): - 1. ``mcp_auth_header`` — per-request/per-user override - 2. OAuth2 client_credentials token — auto-fetched and cached - 3. ``server.authentication_token`` — static token from config/DB - Args: - server: The server configuration. - mcp_auth_header: Optional per-request auth override. - extra_headers: Additional headers to forward. - stdio_env: Environment variables for stdio transport. + server (MCPServer): The server configuration + mcp_auth_header: MCP auth header to be passed to the MCP server. This is optional and will be used if provided. Returns: - Configured MCP client instance. + MCPClient: Configured MCP client instance """ - auth_value = await resolve_mcp_auth(server, mcp_auth_header) - transport = server.transport or MCPTransport.sse # Handle stdio transport @@ -903,7 +842,7 @@ class MCPServerManager: server_url="", # Not used for stdio transport_type=transport, auth_type=server.auth_type, - auth_value=auth_value, + auth_value=mcp_auth_header or server.authentication_token, timeout=60.0, stdio_config=stdio_config, extra_headers=extra_headers, @@ -915,7 +854,7 @@ class MCPServerManager: server_url=server_url, transport_type=transport, auth_type=server.auth_type, - auth_value=auth_value, + auth_value=mcp_auth_header or server.authentication_token, timeout=60.0, extra_headers=extra_headers, ) @@ -955,7 +894,7 @@ class MCPServerManager: stdio_env = self._build_stdio_env(server, raw_headers) - client = await self._create_mcp_client( + client = self._create_mcp_client( server=server, mcp_auth_header=mcp_auth_header, extra_headers=extra_headers, @@ -1015,7 +954,7 @@ class MCPServerManager: stdio_env = self._build_stdio_env(server, raw_headers) - client = await self._create_mcp_client( + client = self._create_mcp_client( server=server, mcp_auth_header=mcp_auth_header, extra_headers=extra_headers, @@ -1059,7 +998,7 @@ class MCPServerManager: stdio_env = self._build_stdio_env(server, raw_headers) - client = await self._create_mcp_client( + client = self._create_mcp_client( server=server, mcp_auth_header=mcp_auth_header, extra_headers=extra_headers, @@ -1103,7 +1042,7 @@ class MCPServerManager: stdio_env = self._build_stdio_env(server, raw_headers) - client = await self._create_mcp_client( + client = self._create_mcp_client( server=server, mcp_auth_header=mcp_auth_header, extra_headers=extra_headers, @@ -1144,7 +1083,7 @@ class MCPServerManager: stdio_env = self._build_stdio_env(server, raw_headers) - client = await self._create_mcp_client( + client = self._create_mcp_client( server=server, mcp_auth_header=mcp_auth_header, extra_headers=extra_headers, @@ -1174,7 +1113,7 @@ class MCPServerManager: stdio_env = self._build_stdio_env(server, raw_headers) - client = await self._create_mcp_client( + client = self._create_mcp_client( server=server, mcp_auth_header=mcp_auth_header, extra_headers=extra_headers, @@ -1439,9 +1378,6 @@ class MCPServerManager: """ Fetch tools from MCP client with timeout and error handling. - Uses anyio.fail_after() instead of asyncio.wait_for() to avoid conflicts - with the MCP SDK's anyio TaskGroup. See GitHub issue #20715 for details. - Args: client: MCP client instance server_name: Name of the server for logging @@ -1449,12 +1385,24 @@ class MCPServerManager: Returns: List of tools from the server """ - try: - with anyio.fail_after(30.0): + + async def _list_tools_task(): + try: tools = await client.list_tools() verbose_logger.debug(f"Tools from {server_name}: {tools}") return tools - except TimeoutError: + except asyncio.CancelledError: + verbose_logger.warning(f"Client operation cancelled for {server_name}") + return [] + except Exception as e: + verbose_logger.warning( + f"Client operation failed for {server_name}: {str(e)}" + ) + return [] + + try: + return await asyncio.wait_for(_list_tools_task(), timeout=30.0) + except asyncio.TimeoutError: verbose_logger.warning(f"Timeout while listing tools from {server_name}") return [] except asyncio.CancelledError: @@ -1969,7 +1917,7 @@ class MCPServerManager: stdio_env = self._build_stdio_env(mcp_server, raw_headers) - client = await self._create_mcp_client( + client = self._create_mcp_client( server=mcp_server, mcp_auth_header=server_auth_header, extra_headers=extra_headers, @@ -2145,8 +2093,8 @@ class MCPServerManager: Note: This now handles prefixed tool names """ for server in self.get_registry().values(): - if server.needs_user_oauth_token: - # Skip OAuth2 servers that rely on user-provided tokens + if server.auth_type == MCPAuth.oauth2: + # Skip OAuth2 servers for now as they may require user-specific tokens continue tools = await self._get_tools_from_server(server) for tool in tools: @@ -2254,43 +2202,6 @@ class MCPServerManager: servers.append(server) return servers - def _get_general_settings(self) -> Dict[str, Any]: - """Get general_settings, importing lazily to avoid circular imports.""" - try: - from litellm.proxy.proxy_server import ( - general_settings as proxy_general_settings, - ) - - return proxy_general_settings - except ImportError: - # Fallback if proxy_server not available - return {} - - def _is_server_accessible_from_ip( - self, server: MCPServer, client_ip: Optional[str] - ) -> bool: - """ - Check if a server is accessible from the given client IP. - - - If client_ip is None, no IP filtering is applied (internal callers). - - If the server has available_on_public_internet=True, it's always accessible. - - Otherwise, only internal/private IPs can access it. - """ - if client_ip is None: - return True - if server.available_on_public_internet: - return True - # Check backwards compat: litellm.public_mcp_servers - public_ids = set(litellm.public_mcp_servers or []) - if server.server_id in public_ids: - return True - # Non-public server: only accessible from internal IPs - general_settings = self._get_general_settings() - internal_networks = IPAddressUtils.parse_internal_networks( - general_settings.get("mcp_internal_ip_ranges") - ) - return IPAddressUtils.is_internal_ip(client_ip, internal_networks) - def get_mcp_server_by_id(self, server_id: str) -> Optional[MCPServer]: """ Get the MCP Server from the server id @@ -2303,72 +2214,27 @@ class MCPServerManager: def get_public_mcp_servers(self) -> List[MCPServer]: """ - Get the public MCP servers (available_on_public_internet=True flag on server). - Also includes servers from litellm.public_mcp_servers for backwards compat. + Get the public MCP servers """ servers: List[MCPServer] = [] - public_ids = set(litellm.public_mcp_servers or []) - for server in self.get_registry().values(): - if server.available_on_public_internet or server.server_id in public_ids: + if litellm.public_mcp_servers is None: + return servers + for server_id in litellm.public_mcp_servers: + server = self.get_mcp_server_by_id(server_id) + if server: servers.append(server) return servers - def get_mcp_server_by_name( - self, server_name: str, client_ip: Optional[str] = None - ) -> Optional[MCPServer]: + def get_mcp_server_by_name(self, server_name: str) -> Optional[MCPServer]: """ - Get the MCP Server from the server name. - - Uses priority-based matching to avoid collisions: - 1. First pass: exact alias match (highest priority) - 2. Second pass: exact server_name match - 3. Third pass: exact name match (lowest priority) - - Args: - server_name: The server name to look up. - client_ip: Optional client IP for access control. When provided, - non-public servers are hidden from external IPs. + Get the MCP Server from the server name """ registry = self.get_registry() - # Pass 1: Match by alias (highest priority) - for server in registry.values(): - if server.alias == server_name: - if not self._is_server_accessible_from_ip(server, client_ip): - return None - return server - # Pass 2: Match by server_name for server in registry.values(): if server.server_name == server_name: - if not self._is_server_accessible_from_ip(server, client_ip): - return None - return server - # Pass 3: Match by name (lowest priority) - for server in registry.values(): - if server.name == server_name: - if not self._is_server_accessible_from_ip(server, client_ip): - return None return server return None - def get_filtered_registry( - self, client_ip: Optional[str] = None - ) -> Dict[str, MCPServer]: - """ - Get registry filtered by client IP access control. - - Args: - client_ip: Optional client IP. When provided, non-public servers - are hidden from external IPs. When None, returns all servers. - """ - registry = self.get_registry() - if client_ip is None: - return registry - return { - k: v - for k, v in registry.items() - if self._is_server_accessible_from_ip(v, client_ip) - } - def _generate_stable_server_id( self, server_name: str, @@ -2441,7 +2307,7 @@ class MCPServerManager: should_skip_health_check = False # Skip if auth_type is oauth2 - if server.needs_user_oauth_token: + if server.auth_type == MCPAuth.oauth2: should_skip_health_check = True # Skip if auth_type is not none and authentication_token is missing elif ( @@ -2456,7 +2322,7 @@ class MCPServerManager: if server.static_headers: extra_headers.update(server.static_headers) - client = await self._create_mcp_client( + client = self._create_mcp_client( server=server, mcp_auth_header=None, extra_headers=extra_headers, @@ -2474,9 +2340,6 @@ class MCPServerManager: except asyncio.TimeoutError: health_check_error = "Health check timed out after 10 seconds" status = "unhealthy" - except asyncio.CancelledError: - health_check_error = "Health check was cancelled" - status = "unknown" except Exception as e: health_check_error = str(e) status = "unhealthy" @@ -2601,7 +2464,6 @@ class MCPServerManager: token_url=server.token_url, registration_url=server.registration_url, allow_all_keys=server.allow_all_keys, - available_on_public_internet=server.available_on_public_internet, ) async def get_all_mcp_servers_unfiltered(self) -> List[LiteLLM_MCPServerTable]: diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index aed81afd254..eb47fb4f608 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -1,6 +1,6 @@ import importlib from datetime import datetime -from typing import Any, Awaitable, Callable, Dict, List, Optional, Union +from typing import Dict, List, Optional, Union from fastapi import APIRouter, Depends, HTTPException, Query, Request @@ -10,7 +10,6 @@ from litellm.proxy._experimental.mcp_server.ui_session_utils import ( ) from litellm.proxy._experimental.mcp_server.utils import merge_mcp_headers from litellm.proxy._types import UserAPIKeyAuth -from litellm.proxy.auth.ip_address_utils import IPAddressUtils from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.types.mcp import MCPAuth from litellm.types.utils import CallTypes @@ -37,7 +36,6 @@ if MCP_AVAILABLE: from litellm.proxy._experimental.mcp_server.server import ( ListMCPToolsRestAPIResponseObject, MCPServer, - _tool_name_matches, execute_mcp_tool, filter_tools_by_allowed_tools, ) @@ -80,87 +78,10 @@ if MCP_AVAILABLE: for tool in tools ] - def _extract_mcp_headers_from_request( - request: Request, - mcp_request_handler_cls, - ) -> tuple: - """ - Extract MCP auth headers from HTTP request. - - Returns: - Tuple of (mcp_auth_header, mcp_server_auth_headers, raw_headers) - """ - headers = request.headers - raw_headers = dict(headers) - mcp_auth_header = mcp_request_handler_cls._get_mcp_auth_header_from_headers( - headers - ) - mcp_server_auth_headers = ( - mcp_request_handler_cls._get_mcp_server_auth_headers_from_headers(headers) - ) - return mcp_auth_header, mcp_server_auth_headers, raw_headers - - async def _resolve_allowed_mcp_servers_with_ip_filter( - request: Request, - user_api_key_dict: UserAPIKeyAuth, - server_id: str, - ) -> List[MCPServer]: - """ - Resolve allowed MCP servers for a tool call with IP filtering. - - Args: - request: The HTTP request object - user_api_key_dict: The user's API key auth object - server_id: The server ID to validate access for - - Returns: - List of allowed MCPServer objects - - Raises: - HTTPException: If the server_id is not allowed - """ - # Get all auth contexts - auth_contexts = await build_effective_auth_contexts(user_api_key_dict) - - # Collect allowed server IDs from all contexts, then apply IP filtering - _rest_client_ip = IPAddressUtils.get_mcp_client_ip(request) - allowed_server_ids_set = set() - for auth_context in auth_contexts: - servers = await global_mcp_server_manager.get_allowed_mcp_servers( - user_api_key_auth=auth_context, - ) - allowed_server_ids_set.update(servers) - - allowed_server_ids_set = set( - global_mcp_server_manager.filter_server_ids_by_ip( - list(allowed_server_ids_set), _rest_client_ip - ) - ) - - # Check if the specified server_id is allowed - if server_id not in allowed_server_ids_set: - raise HTTPException( - status_code=403, - detail={ - "error": "access_denied", - "message": f"The key is not allowed to access server {server_id}", - }, - ) - - # Build allowed_mcp_servers list (only include allowed servers) - allowed_mcp_servers: List[MCPServer] = [] - for allowed_server_id in allowed_server_ids_set: - server = global_mcp_server_manager.get_mcp_server_by_id(allowed_server_id) - if server is not None: - allowed_mcp_servers.append(server) - - return allowed_mcp_servers - async def _get_tools_for_single_server( server, server_auth_header, raw_headers: Optional[Dict[str, str]] = None, - user_api_key_auth: Optional[UserAPIKeyAuth] = None, ): """Helper function to get tools for a single server.""" tools = await global_mcp_server_manager._get_tools_from_server( @@ -175,29 +96,6 @@ if MCP_AVAILABLE: if server.allowed_tools is not None and len(server.allowed_tools) > 0: tools = filter_tools_by_allowed_tools(tools, server) - # Filter tools based on user_api_key_auth.object_permission.mcp_tool_permissions - # This provides per-key/team/org control over which tools can be accessed - if ( - user_api_key_auth - and user_api_key_auth.object_permission - and user_api_key_auth.object_permission.mcp_tool_permissions - ): - allowed_tools_for_server = ( - user_api_key_auth.object_permission.mcp_tool_permissions.get( - server.server_id - ) - ) - if ( - allowed_tools_for_server is not None - and len(allowed_tools_for_server) > 0 - ): - # Filter tools to only include those in the allowed list - tools = [ - tool - for tool in tools - if _tool_name_matches(tool.name, allowed_tools_for_server) - ] - return _create_tool_response_objects(tools, server.mcp_info) async def _resolve_allowed_mcp_servers_for_tool_call( @@ -222,7 +120,9 @@ if MCP_AVAILABLE: ) allowed_mcp_servers: List[MCPServer] = [] for allowed_server_id in allowed_server_ids_set: - server = global_mcp_server_manager.get_mcp_server_by_id(allowed_server_id) + server = global_mcp_server_manager.get_mcp_server_by_id( + allowed_server_id + ) if server is not None: allowed_mcp_servers.append(server) return allowed_mcp_servers @@ -273,25 +173,21 @@ if MCP_AVAILABLE: auth_contexts = await build_effective_auth_contexts(user_api_key_dict) - _rest_client_ip = IPAddressUtils.get_mcp_client_ip(request) - allowed_server_ids_set = set() for auth_context in auth_contexts: servers = await global_mcp_server_manager.get_allowed_mcp_servers( - user_api_key_auth=auth_context, + user_api_key_auth=auth_context ) allowed_server_ids_set.update(servers) - allowed_server_ids = global_mcp_server_manager.filter_server_ids_by_ip( - list(allowed_server_ids_set), _rest_client_ip - ) + allowed_server_ids = list(allowed_server_ids_set) list_tools_result = [] error_message = None # If server_id is specified, only query that specific server if server_id: - if server_id not in allowed_server_ids: + if server_id not in allowed_server_ids_set: raise HTTPException( status_code=403, detail={ @@ -313,10 +209,7 @@ if MCP_AVAILABLE: try: list_tools_result = await _get_tools_for_single_server( - server, - server_auth_header, - raw_headers_from_request, - user_api_key_dict, + server, server_auth_header, raw_headers_from_request ) except Exception as e: verbose_logger.exception( @@ -352,10 +245,7 @@ if MCP_AVAILABLE: try: tools_result = await _get_tools_for_single_server( - server, - server_auth_header, - raw_headers_from_request, - user_api_key_dict, + server, server_auth_header, raw_headers_from_request ) list_tools_result.extend(tools_result) except Exception as e: @@ -449,10 +339,21 @@ if MCP_AVAILABLE: ) ) - # Extract MCP auth headers from request and add to data dict - mcp_auth_header, mcp_server_auth_headers, raw_headers_from_request = ( - _extract_mcp_headers_from_request(request, MCPRequestHandler) + # FIX: Extract MCP auth headers from request + # The UI sends bearer token in x-mcp-auth header and server-specific headers, + # but they weren't being extracted and passed to call_mcp_tool. + # This fix ensures auth headers are properly extracted from the HTTP request + # and passed through to the MCP server for authentication. + headers = request.headers + raw_headers_from_request = dict(headers) + mcp_auth_header = MCPRequestHandler._get_mcp_auth_header_from_headers( + headers ) + mcp_server_auth_headers = ( + MCPRequestHandler._get_mcp_server_auth_headers_from_headers(headers) + ) + + # Add extracted headers to data dict to pass to call_mcp_tool if mcp_auth_header: data["mcp_auth_header"] = mcp_auth_header if mcp_server_auth_headers: @@ -464,9 +365,8 @@ if MCP_AVAILABLE: if "metadata" in data and "user_api_key_auth" in data["metadata"]: data["user_api_key_auth"] = data["metadata"]["user_api_key_auth"] - # Resolve allowed MCP servers with IP filtering - allowed_mcp_servers = await _resolve_allowed_mcp_servers_with_ip_filter( - request, user_api_key_dict, server_id + allowed_mcp_servers = await _resolve_allowed_mcp_servers_for_tool_call( + user_api_key_dict, server_id ) # Call execute_mcp_tool directly (permission checks already done) @@ -528,50 +428,24 @@ if MCP_AVAILABLE: NewMCPServerRequest, ) - def _extract_credentials( - request: NewMCPServerRequest, - ) -> tuple: - """ - Extract OAuth credentials from the nested ``request.credentials`` dict. - - Returns: - (client_id, client_secret, scopes) — any value may be ``None``. - """ - creds = request.credentials if isinstance(request.credentials, dict) else {} - client_id: Optional[str] = creds.get("client_id") - client_secret: Optional[str] = creds.get("client_secret") - scopes_raw = creds.get("scopes") - scopes: Optional[List[str]] = scopes_raw if isinstance(scopes_raw, list) else None - return client_id, client_secret, scopes - async def _execute_with_mcp_client( request: NewMCPServerRequest, - operation: Callable[..., Awaitable[Any]], + operation, mcp_auth_header: Optional[Union[str, Dict[str, str]]] = None, oauth2_headers: Optional[Dict[str, str]] = None, raw_headers: Optional[Dict[str, str]] = None, - ) -> dict: + ): """ - Create a temporary MCP client from *request*, run *operation*, and return the result. - - For M2M OAuth servers (those with ``client_id``, ``client_secret``, and - ``token_url``), the incoming ``oauth2_headers`` are dropped so that - ``resolve_mcp_auth`` can auto-fetch a token via ``client_credentials``. + Common helper to create MCP client, execute operation, and ensure proper cleanup. Args: - request: MCP server configuration submitted by the UI. - operation: Async callable that receives the created client and returns a result dict. - mcp_auth_header: Pre-resolved credential header (API-key / bearer token). - oauth2_headers: Headers extracted from the incoming request (may contain the - litellm API key — must NOT be forwarded for M2M servers). - raw_headers: Raw request headers forwarded for stdio env construction. + request: MCP server configuration + operation: Async function that takes a client and returns the operation result Returns: - The dict returned by *operation*, or an error dict on failure. + Operation result or error response """ try: - client_id, client_secret, scopes = _extract_credentials(request) - server_model = MCPServer( server_id=request.server_id or "", name=request.alias or request.server_name or "", @@ -583,30 +457,18 @@ if MCP_AVAILABLE: args=request.args, env=request.env, static_headers=request.static_headers, - client_id=client_id, - client_secret=client_secret, - token_url=request.token_url, - scopes=scopes, - authorization_url=request.authorization_url, - registration_url=request.registration_url, ) stdio_env = global_mcp_server_manager._build_stdio_env( server_model, raw_headers ) - # For M2M OAuth servers, drop the incoming Authorization header so that - # resolve_mcp_auth can auto-fetch a token via client_credentials. - effective_oauth2_headers = ( - None if server_model.has_client_credentials else oauth2_headers - ) - merged_headers = merge_mcp_headers( - extra_headers=effective_oauth2_headers, + extra_headers=oauth2_headers, static_headers=request.static_headers, ) - client = await global_mcp_server_manager._create_mcp_client( + client = global_mcp_server_manager._create_mcp_client( server=server_model, mcp_auth_header=mcp_auth_header, extra_headers=merged_headers, @@ -615,14 +477,11 @@ if MCP_AVAILABLE: return await operation(client) - except (KeyboardInterrupt, SystemExit): - raise - except BaseException as e: - verbose_logger.error("Error in MCP operation: %s", e, exc_info=True) + except Exception as e: + verbose_logger.error(f"Error in MCP operation: {e}", exc_info=True) return { "status": "error", - "error": True, - "message": "Failed to connect to MCP server. Check proxy logs for details.", + "message": "An internal error has occurred while testing the MCP server.", } @router.post("/test/connection", dependencies=[Depends(user_api_key_auth)]) diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 480c6855343..e63bbed78c4 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -520,6 +520,7 @@ class ProxyBaseLLMRequestProcessing: "adelete_interaction", "acancel_interaction", "asend_message", + "call_mcp_tool", ], version: Optional[str] = None, user_model: Optional[str] = None, diff --git a/litellm/proxy/guardrails/guardrail_hooks/custom_code/custom_code_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/custom_code/custom_code_guardrail.py index c4698b1242d..68f9dfd7abc 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/custom_code/custom_code_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/custom_code/custom_code_guardrail.py @@ -29,7 +29,6 @@ Example custom code (async with HTTP): """ import asyncio -import inspect import threading from typing import TYPE_CHECKING, Any, Dict, Literal, Optional, Type, cast @@ -116,6 +115,9 @@ class CustomCodeGuardrail(CustomGuardrail): GuardrailEventHooks.pre_call, GuardrailEventHooks.during_call, GuardrailEventHooks.post_call, + GuardrailEventHooks.pre_mcp_call, + GuardrailEventHooks.during_mcp_call, + GuardrailEventHooks.logging_only, ] super().__init__( diff --git a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py index 13bdd6fe58d..cc05358baf7 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py @@ -87,7 +87,6 @@ class UnifiedLLMGuardrails(CustomLogger): if CallTypes(call_type) not in endpoint_guardrail_translation_mappings: return data except ValueError: - return data # handle unmapped call types endpoint_translation = endpoint_guardrail_translation_mappings[ diff --git a/ui/litellm-dashboard/src/components/guardrails/custom_code/CustomCodeModal.tsx b/ui/litellm-dashboard/src/components/guardrails/custom_code/CustomCodeModal.tsx index 5535688a9fd..1b8a8c034bd 100644 --- a/ui/litellm-dashboard/src/components/guardrails/custom_code/CustomCodeModal.tsx +++ b/ui/litellm-dashboard/src/components/guardrails/custom_code/CustomCodeModal.tsx @@ -137,6 +137,9 @@ const MODE_OPTIONS = [ { value: "post_call", label: "post_call (Response)" }, { value: "during_call", label: "during_call (Parallel)" }, { value: "logging_only", label: "logging_only" }, + { value: "pre_mcp_call", label: "pre_mcp_call (Before MCP Tool Call)" }, + { value: "post_mcp_call", label: "post_mcp_call (After MCP Tool Call)" }, + { value: "during_mcp_call", label: "during_mcp_call (During MCP Tool Call)" }, ]; // Data for editing an existing guardrail @@ -144,7 +147,7 @@ export interface EditGuardrailData { guardrail_id: string; guardrail_name: string; litellm_params: { - mode?: string; + mode?: string | string[]; default_on?: boolean; custom_code?: string; [key: string]: any; @@ -169,7 +172,7 @@ const CustomCodeModal: React.FC = ({ }) => { const isEditMode = !!editData; const [guardrailName, setGuardrailName] = useState(""); - const [mode, setMode] = useState("pre_call"); + const [mode, setMode] = useState(["pre_call"]); const [defaultOn, setDefaultOn] = useState(false); const [selectedTemplate, setSelectedTemplate] = useState("empty"); const [code, setCode] = useState(CODE_TEMPLATES.empty.code); @@ -241,20 +244,27 @@ const CustomCodeModal: React.FC = ({ setCode(CODE_TEMPLATES[templateKey as keyof typeof CODE_TEMPLATES].code); }; + // Normalize mode from API (string or string[]) to string[] + const normalizeMode = (m: string | string[] | undefined): string[] => { + if (m === undefined || m === null) return ["pre_call"]; + if (Array.isArray(m)) return m.length ? m : ["pre_call"]; + return [m]; + }; + // Reset form when modal opens or editData changes useEffect(() => { if (visible) { if (editData) { // Edit mode: populate with existing data setGuardrailName(editData.guardrail_name || ""); - setMode(editData.litellm_params?.mode || "pre_call"); + setMode(normalizeMode(editData.litellm_params?.mode)); setDefaultOn(editData.litellm_params?.default_on || false); setCode(editData.litellm_params?.custom_code || CODE_TEMPLATES.empty.code); setSelectedTemplate(""); // No template selected in edit mode } else { // Create mode: reset to defaults setGuardrailName(""); - setMode("pre_call"); + setMode(["pre_call"]); setDefaultOn(false); setSelectedTemplate("empty"); setCode(CODE_TEMPLATES.empty.code); @@ -319,7 +329,11 @@ const CustomCodeModal: React.FC = ({ if (guardrailName !== editData.guardrail_name) { updateData.guardrail_name = guardrailName; } - if (mode !== editData.litellm_params?.mode) { + const existingMode = normalizeMode(editData.litellm_params?.mode); + const modeChanged = + mode.length !== existingMode.length || + mode.some((m, i) => m !== existingMode[i]); + if (modeChanged) { updateData.litellm_params.mode = mode; } if (defaultOn !== editData.litellm_params?.default_on) { @@ -382,10 +396,20 @@ const CustomCodeModal: React.FC = ({ parsedInput.texts = []; } + // Use first request-like or response-like mode for test input_type + const requestModes = ["pre_call", "pre_mcp_call"]; + const responseModes = ["post_call", "post_mcp_call"]; + const testInputType: "request" | "response" = + mode.some((m) => requestModes.includes(m)) + ? "request" + : mode.some((m) => responseModes.includes(m)) + ? "response" + : "request"; + const response = await testCustomCodeGuardrail(accessToken, { custom_code: code, test_input: parsedInput, - input_type: mode as "request" | "response", + input_type: testInputType, request_data: { model: "test-model", metadata: {}, @@ -443,14 +467,16 @@ const CustomCodeModal: React.FC = ({ placeholder="e.g., block-pii-custom" /> -
- +
+