mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix: mcp rest auth check
This commit is contained in:
parent
a1bba8c99b
commit
658bbcc2d5
5 changed files with 89 additions and 10 deletions
|
|
@ -516,7 +516,7 @@ class MCPRequestHandler:
|
|||
Check if the tool is allowed for the given user/key based on permissions
|
||||
"""
|
||||
if len(allowed_mcp_servers) == 0:
|
||||
return True
|
||||
return False
|
||||
elif server_name in allowed_mcp_servers:
|
||||
return True
|
||||
return False
|
||||
|
|
|
|||
|
|
@ -1,9 +1,10 @@
|
|||
import importlib
|
||||
from typing import Dict, List, Optional, Union
|
||||
|
||||
from fastapi import APIRouter, Depends, Query, Request
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, Request
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.proxy._experimental.mcp_server.ui_session_utils import build_effective_auth_contexts
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.types.mcp import MCPAuth
|
||||
|
|
@ -22,7 +23,7 @@ router = APIRouter(
|
|||
)
|
||||
|
||||
if MCP_AVAILABLE:
|
||||
from litellm.experimental_mcp_client.client import MCPTool
|
||||
from mcp.types import Tool as MCPTool
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
global_mcp_server_manager,
|
||||
)
|
||||
|
|
@ -134,11 +135,30 @@ if MCP_AVAILABLE:
|
|||
MCPRequestHandler._get_mcp_server_auth_headers_from_headers(headers)
|
||||
)
|
||||
|
||||
auth_contexts = await build_effective_auth_contexts(user_api_key_dict)
|
||||
|
||||
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 = 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_set:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={
|
||||
"error": "access_denied",
|
||||
"message": f"The key is not allowed to access server {server_id}",
|
||||
},
|
||||
)
|
||||
server = global_mcp_server_manager.get_mcp_server_by_id(server_id)
|
||||
if server is None:
|
||||
return {
|
||||
|
|
@ -165,9 +185,22 @@ if MCP_AVAILABLE:
|
|||
"message": f"Failed to get tools from server {server.name}: {str(e)}",
|
||||
}
|
||||
else:
|
||||
# Query all servers
|
||||
if not allowed_server_ids:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={
|
||||
"error": "access_denied",
|
||||
"message": "The key is not allowed to access any MCP servers.",
|
||||
},
|
||||
)
|
||||
|
||||
# Query all servers the user has access to
|
||||
errors = []
|
||||
for server in global_mcp_server_manager.get_registry().values():
|
||||
for allowed_server_id in allowed_server_ids:
|
||||
server = global_mcp_server_manager.get_mcp_server_by_id(allowed_server_id)
|
||||
if server is None:
|
||||
continue
|
||||
|
||||
server_auth_header = _get_server_auth_header(
|
||||
server, mcp_server_auth_headers, mcp_auth_header
|
||||
)
|
||||
|
|
@ -225,6 +258,17 @@ if MCP_AVAILABLE:
|
|||
|
||||
try:
|
||||
data = await request.json()
|
||||
# Server ID permission check (server_id is required)
|
||||
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",
|
||||
# },
|
||||
# )
|
||||
|
||||
data = await add_litellm_data_to_request(
|
||||
data=data,
|
||||
request=request,
|
||||
|
|
@ -252,12 +296,36 @@ if MCP_AVAILABLE:
|
|||
if mcp_server_auth_headers:
|
||||
data["mcp_server_auth_headers"] = mcp_server_auth_headers
|
||||
data["raw_headers"] = raw_headers_from_request
|
||||
|
||||
|
||||
# Extract user_api_key_auth from metadata and add to top level
|
||||
# call_mcp_tool expects user_api_key_auth as a top-level parameter
|
||||
if "metadata" in data and "user_api_key_auth" in data["metadata"]:
|
||||
data["user_api_key_auth"] = data["metadata"]["user_api_key_auth"]
|
||||
|
||||
|
||||
# Get all auth contexts
|
||||
auth_contexts = await build_effective_auth_contexts(user_api_key_dict)
|
||||
|
||||
# Collect allowed server IDs from all contexts
|
||||
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)
|
||||
|
||||
# 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}",
|
||||
},
|
||||
)
|
||||
|
||||
# Restrict to the specified server only
|
||||
data["mcp_servers"] = [server_id]
|
||||
|
||||
result = await call_mcp_tool(**data)
|
||||
return result
|
||||
except BlockedPiiEntityError as e:
|
||||
|
|
|
|||
|
|
@ -1240,6 +1240,11 @@ if MCP_AVAILABLE:
|
|||
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
|
||||
|
|
|
|||
|
|
@ -40,7 +40,7 @@ const MCPToolsViewer = ({
|
|||
if (!accessToken) throw new Error("Access Token required");
|
||||
|
||||
try {
|
||||
const result: CallMCPToolResponse = await callMCPTool(accessToken, args.tool.name, args.arguments);
|
||||
const result: CallMCPToolResponse = await callMCPTool(accessToken, serverId, args.tool.name, args.arguments);
|
||||
return result;
|
||||
} catch (error) {
|
||||
throw error;
|
||||
|
|
|
|||
|
|
@ -6128,12 +6128,17 @@ export const listMCPTools = async (accessToken: string, serverId: string) => {
|
|||
}
|
||||
};
|
||||
|
||||
export const callMCPTool = async (accessToken: string, toolName: string, toolArguments: Record<string, any>) => {
|
||||
export const callMCPTool = async (
|
||||
accessToken: string,
|
||||
serverId: string,
|
||||
toolName: string,
|
||||
toolArguments: Record<string, any>,
|
||||
) => {
|
||||
try {
|
||||
// Construct base URL
|
||||
let url = proxyBaseUrl ? `${proxyBaseUrl}/mcp-rest/tools/call` : `/mcp-rest/tools/call`;
|
||||
|
||||
console.log("Calling MCP tool:", toolName, "with arguments:", toolArguments);
|
||||
console.log("Calling MCP tool:", toolName, "with arguments:", toolArguments, "for server:", serverId);
|
||||
|
||||
const headers: Record<string, string> = {
|
||||
[globalLitellmHeaderName]: `Bearer ${accessToken}`,
|
||||
|
|
@ -6144,6 +6149,7 @@ export const callMCPTool = async (accessToken: string, toolName: string, toolArg
|
|||
method: "POST",
|
||||
headers,
|
||||
body: JSON.stringify({
|
||||
server_id: serverId,
|
||||
name: toolName,
|
||||
arguments: toolArguments,
|
||||
}),
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue