From 658bbcc2d56283d961a7f0a8b6cac7820e908eb2 Mon Sep 17 00:00:00 2001 From: Yuta Saito Date: Tue, 13 Jan 2026 17:18:58 +0900 Subject: [PATCH] fix: mcp rest auth check --- .../mcp_server/auth/user_api_key_auth_mcp.py | 2 +- .../mcp_server/rest_endpoints.py | 80 +++++++++++++++++-- .../proxy/_experimental/mcp_server/server.py | 5 ++ .../src/components/mcp_tools/mcp_tools.tsx | 2 +- .../src/components/networking.tsx | 10 ++- 5 files changed, 89 insertions(+), 10 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index b43f4217177..49d6ac7d898 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -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 diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 642cb0cec2d..cf391a3a74f 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -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: diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 2adfe97c611..6c4a6c37c90 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -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 diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_tools.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_tools.tsx index 5a971e8b471..37c61023f23 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_tools.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_tools.tsx @@ -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; diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index 0a378eba3b2..920448e8ac1 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -6128,12 +6128,17 @@ export const listMCPTools = async (accessToken: string, serverId: string) => { } }; -export const callMCPTool = async (accessToken: string, toolName: string, toolArguments: Record) => { +export const callMCPTool = async ( + accessToken: string, + serverId: string, + toolName: string, + toolArguments: Record, +) => { 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 = { [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, }),