mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-06 08:16:43 +00:00
fix: rest_endpoints own allowed_mcp_servers for MCP calls
This commit is contained in:
parent
658bbcc2d5
commit
22aad95bb1
2 changed files with 126 additions and 48 deletions
|
|
@ -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)}")
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue