fix: mcp rest auth check

This commit is contained in:
Yuta Saito 2026-01-13 17:18:58 +09:00
parent a1bba8c99b
commit 658bbcc2d5
5 changed files with 89 additions and 10 deletions

View file

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

View file

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

View file

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

View file

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

View file

@ -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,
}),