feat(mcp/): support custom code guardrails for mcp calls

allows custom code guardrails to work on mcp input
This commit is contained in:
Krrish Dholakia 2026-02-06 14:44:17 -08:00 • committed by Shin
parent 15881e1d70
commit ae1fc80445
7 changed files with 129 additions and 381 deletions

View file

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

View file

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

View file

@ -520,6 +520,7 @@ class ProxyBaseLLMRequestProcessing:
"adelete_interaction",
"acancel_interaction",
"asend_message",
"call_mcp_tool",
],
version: Optional[str] = None,
user_model: Optional[str] = None,

View file

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

View file

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

View file

@ -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<CustomCodeModalProps> = ({
}) => {
const isEditMode = !!editData;
const [guardrailName, setGuardrailName] = useState("");
const [mode, setMode] = useState<string>("pre_call");
const [mode, setMode] = useState<string[]>(["pre_call"]);
const [defaultOn, setDefaultOn] = useState(false);
const [selectedTemplate, setSelectedTemplate] = useState<string>("empty");
const [code, setCode] = useState(CODE_TEMPLATES.empty.code);
@ -241,20 +244,27 @@ const CustomCodeModal: React.FC<CustomCodeModalProps> = ({
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<CustomCodeModalProps> = ({
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<CustomCodeModalProps> = ({
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<CustomCodeModalProps> = ({
placeholder="e.g., block-pii-custom"
/>
</div>
<div className="w-[180px]">
<label className="block text-xs font-medium text-gray-600 mb-1">Mode</label>
<div className="w-[280px]">
<label className="block text-xs font-medium text-gray-600 mb-1">Mode (can select multiple)</label>
<Select
mode="multiple"
value={mode}
onChange={setMode}
options={MODE_OPTIONS}
className="w-full"
size="middle"
placeholder="Select modes"
/>
</div>
<div className="w-[180px]">

View file

@ -1129,7 +1129,6 @@ const ChatUI: React.FC<ChatUIProps> = ({
mcpServerId,
selectedMCPDirectTool,
mcpToolArguments,
selectedGuardrails.length > 0 ? { guardrails: selectedGuardrails } : undefined,
);
const resultText =
result?.content?.length > 0