mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
feat(mcp/): support custom code guardrails for mcp calls
allows custom code guardrails to work on mcp input
This commit is contained in:
parent
15881e1d70
commit
ae1fc80445
7 changed files with 129 additions and 381 deletions
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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)])
|
||||
|
|
|
|||
|
|
@ -520,6 +520,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
"adelete_interaction",
|
||||
"acancel_interaction",
|
||||
"asend_message",
|
||||
"call_mcp_tool",
|
||||
],
|
||||
version: Optional[str] = None,
|
||||
user_model: Optional[str] = None,
|
||||
|
|
|
|||
|
|
@ -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__(
|
||||
|
|
|
|||
|
|
@ -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[
|
||||
|
|
|
|||
|
|
@ -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]">
|
||||
|
|
|
|||
|
|
@ -1129,7 +1129,6 @@ const ChatUI: React.FC<ChatUIProps> = ({
|
|||
mcpServerId,
|
||||
selectedMCPDirectTool,
|
||||
mcpToolArguments,
|
||||
selectedGuardrails.length > 0 ? { guardrails: selectedGuardrails } : undefined,
|
||||
);
|
||||
const resultText =
|
||||
result?.content?.length > 0
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue