fix: rest_endpoints own allowed_mcp_servers for MCP calls

This commit is contained in:
Yuta Saito 2026-01-14 06:39:46 +09:00
parent 658bbcc2d5
commit 22aad95bb1
2 changed files with 126 additions and 48 deletions

View file

@ -1,4 +1,5 @@
import importlib
from datetime import datetime
from typing import Dict, List, Optional, Union
from fastapi import APIRouter, Depends, HTTPException, Query, Request
@ -29,7 +30,9 @@ if MCP_AVAILABLE:
)
from litellm.proxy._experimental.mcp_server.server import (
ListMCPToolsRestAPIResponseObject,
MCPServer,
call_mcp_tool,
execute_mcp_tool,
filter_tools_by_allowed_tools,
)
@ -258,16 +261,29 @@ if MCP_AVAILABLE:
try:
data = await request.json()
# Server ID permission check (server_id is required)
# Validate required parameters early
server_id = data.get("server_id")
# if not server_id:
# raise HTTPException(
# status_code=400,
# detail={
# "error": "missing_parameter",
# "message": "server_id is required in request body",
# },
# )
if not server_id:
raise HTTPException(
status_code=400,
detail={
"error": "missing_parameter",
"message": "server_id is required in request body",
},
)
tool_name = data.get("name")
if not tool_name:
raise HTTPException(
status_code=400,
detail={
"error": "missing_parameter",
"message": "name is required in request body",
},
)
tool_arguments = data.get("arguments")
data = await add_litellm_data_to_request(
data=data,
@ -323,10 +339,26 @@ if MCP_AVAILABLE:
},
)
# Restrict to the specified server only
data["mcp_servers"] = [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)
result = await call_mcp_tool(**data)
# Call execute_mcp_tool directly (permission checks already done)
result = await execute_mcp_tool(
name=tool_name,
arguments=tool_arguments,
allowed_mcp_servers=allowed_mcp_servers,
start_time=datetime.now(),
user_api_key_auth=data.get("user_api_key_auth"),
mcp_auth_header=data.get("mcp_auth_header"),
mcp_server_auth_headers=data.get("mcp_server_auth_headers"),
oauth2_headers=data.get("oauth2_headers"),
raw_headers=data.get("raw_headers"),
litellm_logging_obj=data.get("litellm_logging_obj"),
)
return result
except BlockedPiiEntityError as e:
verbose_logger.error(f"BlockedPiiEntityError in MCP tool call: {str(e)}")

View file

@ -1200,52 +1200,38 @@ if MCP_AVAILABLE:
return managed_resource_templates
@client
async def call_mcp_tool(
async def execute_mcp_tool(
name: str,
arguments: Optional[Dict[str, Any]] = None,
arguments: Dict[str, Any],
allowed_mcp_servers: List[MCPServer],
start_time: datetime,
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
mcp_auth_header: Optional[str] = None,
mcp_servers: Optional[List[str]] = None,
mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]] = None,
oauth2_headers: Optional[Dict[str, str]] = None,
raw_headers: Optional[Dict[str, str]] = None,
**kwargs: Any,
) -> CallToolResult:
"""
Call a specific tool with the provided arguments (handles prefixed tool names)
Execute MCP tool.
This function assumes permission checks have already been performed.
Args:
name: Tool name (may include server prefix)
arguments: Tool arguments
allowed_mcp_servers: Pre-validated list of servers the user can access
start_time: Start time for logging
user_api_key_auth: Optional user API key auth for logging
mcp_auth_header: Optional MCP auth header
mcp_server_auth_headers: Optional server-specific auth headers
oauth2_headers: Optional OAuth2 headers
raw_headers: Optional raw HTTP headers
**kwargs: Additional arguments (e.g., litellm_logging_obj)
Returns:
CallToolResult: Tool execution result
"""
start_time = datetime.now()
if arguments is None:
raise HTTPException(
status_code=400, detail="Request arguments are required"
)
## CHECK IF USER IS ALLOWED TO CALL THIS TOOL
allowed_mcp_server_ids = (
await global_mcp_server_manager.get_allowed_mcp_servers(
user_api_key_auth=user_api_key_auth,
)
)
allowed_mcp_servers: List[MCPServer] = []
for allowed_mcp_server_id in allowed_mcp_server_ids:
allowed_server = global_mcp_server_manager.get_mcp_server_by_id(
allowed_mcp_server_id
)
if allowed_server is not None:
allowed_mcp_servers.append(allowed_server)
allowed_mcp_servers = await _get_allowed_mcp_servers_from_mcp_server_names(
mcp_servers=mcp_servers,
allowed_mcp_servers=allowed_mcp_servers,
)
if not allowed_mcp_servers:
raise HTTPException(
status_code=403,
detail="User not allowed to call this tool.",
)
# Track resolved MCP server for both permission checks and dispatch
mcp_server: Optional[MCPServer] = None
@ -1364,6 +1350,66 @@ if MCP_AVAILABLE:
)
return response
@client
async def call_mcp_tool(
name: str,
arguments: Optional[Dict[str, Any]] = None,
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
mcp_auth_header: Optional[str] = None,
mcp_servers: Optional[List[str]] = None,
mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]] = None,
oauth2_headers: Optional[Dict[str, str]] = None,
raw_headers: Optional[Dict[str, str]] = None,
**kwargs: Any,
) -> CallToolResult:
"""
Call a specific tool with the provided arguments (handles prefixed tool names).
"""
start_time = datetime.now()
if arguments is None:
raise HTTPException(
status_code=400, detail="Request arguments are required"
)
## CHECK IF USER IS ALLOWED TO CALL THIS TOOL
allowed_mcp_server_ids = (
await global_mcp_server_manager.get_allowed_mcp_servers(
user_api_key_auth=user_api_key_auth,
)
)
allowed_mcp_servers: List[MCPServer] = []
for allowed_mcp_server_id in allowed_mcp_server_ids:
allowed_server = global_mcp_server_manager.get_mcp_server_by_id(
allowed_mcp_server_id
)
if allowed_server is not None:
allowed_mcp_servers.append(allowed_server)
allowed_mcp_servers = await _get_allowed_mcp_servers_from_mcp_server_names(
mcp_servers=mcp_servers,
allowed_mcp_servers=allowed_mcp_servers,
)
if not allowed_mcp_servers:
raise HTTPException(
status_code=403,
detail="User not allowed to call this tool.",
)
# Delegate to execute_mcp_tool for execution
return await execute_mcp_tool(
name=name,
arguments=arguments,
allowed_mcp_servers=allowed_mcp_servers,
start_time=start_time,
user_api_key_auth=user_api_key_auth,
mcp_auth_header=mcp_auth_header,
mcp_server_auth_headers=mcp_server_auth_headers,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
**kwargs,
)
async def mcp_get_prompt(
name: str,
arguments: Optional[Dict[str, Any]] = None,