Merge pull request #19051 from BerriAI/litellm_fix_mcp-rest-auth-checks

[fix] mcp rest auth checks
This commit is contained in:
YutaSaito 2026-01-14 10:39:28 +09:00 • committed by GitHub
commit a66e007574
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
8 changed files with 677 additions and 192 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,13 @@
import importlib
from datetime import datetime
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,13 +26,14 @@ 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,
)
from litellm.proxy._experimental.mcp_server.server import (
ListMCPToolsRestAPIResponseObject,
call_mcp_tool,
MCPServer,
execute_mcp_tool,
filter_tools_by_allowed_tools,
)
@ -134,11 +139,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 +189,24 @@ 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 +264,30 @@ if MCP_AVAILABLE:
try:
data = await request.json()
# 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",
},
)
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,
request=request,
@ -252,13 +315,55 @@ 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"]
result = await call_mcp_tool(**data)
# 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}",
},
)
# 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)
# 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)}")
@ -301,7 +406,6 @@ if MCP_AVAILABLE:
# /health/tools/list -> List tools from MCP server
# For these routes users will dynamically pass the MCP connection params, they don't need to be on the MCP registry
########################################################
from litellm.proxy._experimental.mcp_server.server import MCPServer
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
NewMCPServerRequest,
)

View file

@ -1200,47 +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,
)
# Track resolved MCP server for both permission checks and dispatch
mcp_server: Optional[MCPServer] = None
@ -1359,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,

View file

@ -13,11 +13,14 @@ import litellm
from litellm.types.utils import StandardLoggingPayload
from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy._experimental.mcp_server.server import (
mcp_server_tool_call,
mcp_server_tool_call,
set_auth_context,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
MCPServerManager,
MCPServerManager,
)
from litellm.proxy.proxy_server import LiteLLM_ObjectPermissionTable
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.mcp import MCPPostCallResponseObject
from litellm.types.utils import HiddenParams
from mcp.types import Tool as MCPTool, CallToolResult, TextContent
@ -34,6 +37,20 @@ class TestMCPLogger(CustomLogger):
print(f"Captured standard_logging_payload: {self.standard_logging_payload}")
def _set_authorized_user(server_ids):
"""Configure auth context with permission to call the specified servers."""
server_list = list(server_ids)
user_auth = UserAPIKeyAuth(
api_key="test",
user_id="test_user",
object_permission=LiteLLM_ObjectPermissionTable(
object_permission_id="mcp-test-permissions",
mcp_servers=server_list,
),
)
set_auth_context(user_api_key_auth=user_auth, mcp_servers=server_list)
@pytest.mark.asyncio
async def test_mcp_cost_tracking():
# Create a mock tool call result
@ -87,6 +104,8 @@ async def test_mcp_cost_tracking():
with patch('litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager', local_mcp_server_manager), \
patch('litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager', local_mcp_server_manager):
_set_authorized_user(local_mcp_server_manager.get_all_mcp_server_ids())
print("tool_name_to_mcp_server_name_mapping", local_mcp_server_manager.tool_name_to_mcp_server_name_mapping)
# Manually add the tool mapping to ensure it's available (since mocking might not capture it properly)
@ -197,6 +216,8 @@ async def test_mcp_cost_tracking_per_tool():
with patch('litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager', local_mcp_server_manager), \
patch('litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager', local_mcp_server_manager):
_set_authorized_user(local_mcp_server_manager.get_all_mcp_server_ids())
print("tool_name_to_mcp_server_name_mapping", local_mcp_server_manager.tool_name_to_mcp_server_name_mapping)
# Test 1: Call expensive_tool - should cost 5.0
@ -327,6 +348,8 @@ async def test_mcp_tool_call_hook():
with patch('litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager', local_mcp_server_manager), \
patch('litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager', local_mcp_server_manager):
_set_authorized_user(local_mcp_server_manager.get_all_mcp_server_ids())
print("tool_name_to_mcp_server_name_mapping", local_mcp_server_manager.tool_name_to_mcp_server_name_mapping)
# Call mcp tool using the correct separator format (- not /)

View file

@ -1,6 +1,7 @@
# Create server parameters for stdio connection
import os
import sys
from litellm.proxy.proxy_server import LiteLLM_ObjectPermissionTable
import pytest
from unittest.mock import AsyncMock, MagicMock, patch
from contextlib import asynccontextmanager
@ -630,8 +631,15 @@ async def test_list_tools_rest_api_server_not_found():
from fastapi import Query
from litellm.proxy._types import UserAPIKeyAuth
# Mock UserAPIKeyAuth
mock_user_auth = UserAPIKeyAuth(api_key="test", user_id="test")
# Mock UserAPIKeyAuth with explicit permission to access the requested server id
mock_user_auth = UserAPIKeyAuth(
api_key="test",
user_id="test",
object_permission=LiteLLM_ObjectPermissionTable(
object_permission_id="dummy",
mcp_servers=["non_existent_server_id"],
),
)
# Mock request
mock_request = MagicMock()
@ -704,7 +712,16 @@ async def test_list_tools_rest_api_success():
)
# Mock UserAPIKeyAuth
mock_user_auth = UserAPIKeyAuth(api_key="test", user_id="test")
mock_user_auth = UserAPIKeyAuth(
api_key="test",
user_id="test",
object_permission=LiteLLM_ObjectPermissionTable(
object_permission_id="dummy",
mcp_servers=list(
global_mcp_server_manager.get_all_mcp_server_ids()
),
),
)
# Get the server ID
server_id = list(global_mcp_server_manager.get_registry().keys())[0]
@ -1718,6 +1735,7 @@ async def test_list_tool_rest_api_with_server_specific_auth():
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
MCPRequestHandler,
)
from litellm.proxy._types import UserAPIKeyAuth
# Create mock request with server-specific auth headers
mock_request = MagicMock()
@ -1727,10 +1745,6 @@ async def test_list_tool_rest_api_with_server_specific_auth():
"x-mcp-slack-authorization": "Bearer slack_token",
}
# Create mock user_api_key_dict
mock_user_api_key_dict = MagicMock()
mock_user_api_key_dict.user_id = "test_user"
# Mock the MCPRequestHandler methods
with patch.object(
MCPRequestHandler, "_get_mcp_auth_header_from_headers"
@ -1748,6 +1762,9 @@ async def test_list_tool_rest_api_with_server_specific_auth():
with patch(
"litellm.proxy._experimental.mcp_server.rest_endpoints.global_mcp_server_manager"
) as mock_manager:
mock_manager.get_allowed_mcp_servers = AsyncMock(
return_value=["test-server-123"]
)
# Create a mock server
mock_server = MagicMock()
mock_server.server_id = "test-server-123"
@ -1757,6 +1774,15 @@ async def test_list_tool_rest_api_with_server_specific_auth():
mock_manager.get_mcp_server_by_id.return_value = mock_server
mock_user_api_key_dict = UserAPIKeyAuth(
api_key="test",
user_id="test_user",
object_permission=LiteLLM_ObjectPermissionTable(
object_permission_id="dummy",
mcp_servers=[mock_server.server_id],
),
)
# Mock the _get_tools_for_single_server function
with patch(
"litellm.proxy._experimental.mcp_server.rest_endpoints._get_tools_for_single_server"
@ -1803,6 +1829,7 @@ async def test_list_tool_rest_api_with_default_auth():
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
MCPRequestHandler,
)
from litellm.proxy._types import UserAPIKeyAuth
# Create mock request with default auth header only
mock_request = MagicMock()
@ -1811,10 +1838,6 @@ async def test_list_tool_rest_api_with_default_auth():
"x-mcp-authorization": "Bearer default_token",
}
# Create mock user_api_key_dict
mock_user_api_key_dict = MagicMock()
mock_user_api_key_dict.user_id = "test_user"
# Mock the MCPRequestHandler methods
with patch.object(
MCPRequestHandler, "_get_mcp_auth_header_from_headers"
@ -1829,6 +1852,9 @@ async def test_list_tool_rest_api_with_default_auth():
with patch(
"litellm.proxy._experimental.mcp_server.rest_endpoints.global_mcp_server_manager"
) as mock_manager:
mock_manager.get_allowed_mcp_servers = AsyncMock(
return_value=["test-server-123"]
)
# Create a mock server
mock_server = MagicMock()
mock_server.server_id = "test-server-123"
@ -1838,6 +1864,15 @@ async def test_list_tool_rest_api_with_default_auth():
mock_manager.get_mcp_server_by_id.return_value = mock_server
mock_user_api_key_dict = UserAPIKeyAuth(
api_key="test",
user_id="test_user",
object_permission=LiteLLM_ObjectPermissionTable(
object_permission_id="dummy",
mcp_servers=[mock_server.server_id],
),
)
# Mock the _get_tools_for_single_server function
with patch(
"litellm.proxy._experimental.mcp_server.rest_endpoints._get_tools_for_single_server"
@ -1884,6 +1919,7 @@ async def test_list_tool_rest_api_all_servers_with_auth():
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
MCPRequestHandler,
)
from litellm.proxy._types import UserAPIKeyAuth
# Create mock request with server-specific auth headers
mock_request = MagicMock()
@ -1893,10 +1929,6 @@ async def test_list_tool_rest_api_all_servers_with_auth():
"x-mcp-slack-authorization": "Bearer slack_token",
}
# Create mock user_api_key_dict
mock_user_api_key_dict = MagicMock()
mock_user_api_key_dict.user_id = "test_user"
# Mock the MCPRequestHandler methods
with patch.object(
MCPRequestHandler, "_get_mcp_auth_header_from_headers"
@ -1929,6 +1961,23 @@ async def test_list_tool_rest_api_all_servers_with_auth():
"zapier": mock_zapier_server,
"slack": mock_slack_server,
}
mock_manager.get_allowed_mcp_servers = AsyncMock(
return_value=["zapier", "slack"]
)
mock_manager.get_mcp_server_by_id.side_effect = (
lambda server_id: mock_manager.get_registry.return_value.get(
server_id
)
)
mock_user_api_key_dict = UserAPIKeyAuth(
api_key="test",
user_id="test_user",
object_permission=LiteLLM_ObjectPermissionTable(
object_permission_id="dummy",
mcp_servers=["zapier", "slack"],
),
)
# Mock the _get_tools_for_single_server function
with patch(
@ -1971,17 +2020,15 @@ async def test_list_tool_rest_api_all_servers_with_auth():
assert result["tools"][0].name == "send_email"
assert result["tools"][1].name == "send_message"
# Verify that _get_tools_for_single_server was called for both servers with correct auth headers
# Verify that _get_tools_for_single_server was called for both servers
assert mock_get_tools.call_count == 2
calls = mock_get_tools.call_args_list
server_auth_map = {
call_args[0][0]: call_args[0][1]
for call_args in mock_get_tools.call_args_list
}
# First call should be for zapier server with zapier auth
assert calls[0][0][0] == mock_zapier_server # server
assert calls[0][0][1] == "Bearer zapier_token" # server_auth_header
# Second call should be for slack server with slack auth
assert calls[1][0][0] == mock_slack_server # server
assert calls[1][0][1] == "Bearer slack_token" # server_auth_header
assert server_auth_map.get(mock_zapier_server) == "Bearer zapier_token"
assert server_auth_map.get(mock_slack_server) == "Bearer slack_token"
@pytest.mark.asyncio

View file

@ -1,6 +1,8 @@
from typing import Dict, Optional
import json
from typing import Any, Dict, Optional
import pytest
from fastapi import HTTPException
from starlette.requests import Request
from litellm.proxy._experimental.mcp_server import rest_endpoints
@ -12,8 +14,21 @@ from litellm.proxy._types import NewMCPServerRequest, UserAPIKeyAuth
from litellm.types.mcp import MCPAuth
def _build_request(headers: Optional[Dict[str, str]] = None) -> Request:
def _build_request(
headers: Optional[Dict[str, str]] = None,
*,
path: str = "/mcp-rest/test/tools/list",
method: str = "POST",
json_body: Optional[Any] = None,
body: Optional[bytes] = None,
) -> Request:
headers = headers or {}
if json_body is not None:
body_bytes = json.dumps(json_body).encode("utf-8")
elif body is not None:
body_bytes = body
else:
body_bytes = b""
raw_headers = [
(key.lower().encode("latin-1"), value.encode("latin-1"))
for key, value in headers.items()
@ -21,13 +36,18 @@ def _build_request(headers: Optional[Dict[str, str]] = None) -> Request:
scope = {
"type": "http",
"http_version": "1.1",
"method": "POST",
"path": "/mcp-rest/test/tools/list",
"method": method,
"path": path,
"headers": raw_headers,
}
state = {"sent": False}
async def receive():
return {"type": "http.request", "body": b"", "more_body": False}
if state["sent"]:
return {"type": "http.request", "body": b"", "more_body": False}
state["sent"] = True
return {"type": "http.request", "body": body_bytes, "more_body": False}
return Request(scope, receive=receive)
@ -53,146 +73,380 @@ def _route_has_dependency(route, dependency) -> bool:
return any(getattr(dep, "call", None) == dependency for dep in dependant.dependencies)
@pytest.mark.asyncio
async def test_execute_with_mcp_client_redacts_stack_trace(monkeypatch):
def fake_create_client(*args, **kwargs):
return object()
class TestExecuteWithMcpClient:
@pytest.mark.asyncio
async def test_redacts_stack_trace(self, monkeypatch):
def fake_create_client(*args, **kwargs):
return object()
monkeypatch.setattr(
rest_endpoints.global_mcp_server_manager,
"_create_mcp_client",
fake_create_client,
)
monkeypatch.setattr(
rest_endpoints.global_mcp_server_manager,
"_create_mcp_client",
fake_create_client,
)
async def failing_operation(client):
raise RuntimeError("boom")
async def failing_operation(client):
raise RuntimeError("boom")
payload = NewMCPServerRequest(
server_name="example",
url="https://example.com",
auth_type=MCPAuth.none,
)
payload = NewMCPServerRequest(
server_name="example",
url="https://example.com",
auth_type=MCPAuth.none,
)
result = await rest_endpoints._execute_with_mcp_client(
payload, failing_operation
)
result = await rest_endpoints._execute_with_mcp_client(
payload, failing_operation
)
assert result["status"] == "error"
assert "stack_trace" not in result
assert result["status"] == "error"
assert "stack_trace" not in result
def test_test_connection_requires_auth_dependency():
route = _get_route("/mcp-rest/test/connection", "POST")
assert _route_has_dependency(route, user_api_key_auth)
class TestTestConnection:
def test_requires_auth_dependency(self):
route = _get_route("/mcp-rest/test/connection", "POST")
assert _route_has_dependency(route, user_api_key_auth)
@pytest.mark.asyncio
async def test_test_tools_list_forwards_mcp_auth_header(monkeypatch):
"""Ensure credential-based auth forwards the auth_value to the MCP client."""
class TestTestToolsList:
pytestmark = pytest.mark.asyncio
captured: dict = {}
async def test_forwards_mcp_auth_header(self, monkeypatch):
"""Ensure credential-based auth forwards the auth_value to the MCP client."""
async def fake_execute(
request,
operation,
mcp_auth_header=None,
oauth2_headers=None,
raw_headers=None,
):
captured["mcp_auth_header"] = mcp_auth_header
captured["oauth2_headers"] = oauth2_headers
return {
"tools": [],
"error": None,
"message": "Successfully retrieved tools",
captured: dict = {}
async def fake_execute(
request,
operation,
mcp_auth_header=None,
oauth2_headers=None,
raw_headers=None,
):
captured["mcp_auth_header"] = mcp_auth_header
captured["oauth2_headers"] = oauth2_headers
return {
"tools": [],
"error": None,
"message": "Successfully retrieved tools",
}
monkeypatch.setattr(
rest_endpoints, "_execute_with_mcp_client", fake_execute, raising=False
)
oauth_call_counter = {"count": 0}
def fake_oauth(headers):
oauth_call_counter["count"] += 1
return {"Authorization": "Bearer oauth"}
monkeypatch.setattr(
auth_mcp.MCPRequestHandler,
"_get_oauth2_headers_from_headers",
staticmethod(fake_oauth),
raising=False,
)
request = _build_request()
payload = NewMCPServerRequest(
server_name="example",
url="https://example.com",
auth_type=MCPAuth.api_key,
credentials={"auth_value": "secret-key"},
)
result = await rest_endpoints.test_tools_list(
request, payload, user_api_key_dict=UserAPIKeyAuth()
)
assert result["message"] == "Successfully retrieved tools"
assert captured["mcp_auth_header"] == "secret-key"
assert captured["oauth2_headers"] is None
assert oauth_call_counter["count"] == 0
async def test_extracts_oauth2_headers(self, monkeypatch):
"""Ensure oauth2 auth type pulls oauth headers and omits MCP auth header."""
captured: dict = {}
async def fake_execute(
request,
operation,
mcp_auth_header=None,
oauth2_headers=None,
raw_headers=None,
):
captured["mcp_auth_header"] = mcp_auth_header
captured["oauth2_headers"] = oauth2_headers
return {
"tools": [],
"error": None,
"message": "Successfully retrieved tools",
}
monkeypatch.setattr(
rest_endpoints, "_execute_with_mcp_client", fake_execute, raising=False
)
oauth_headers = {"Authorization": "Bearer oauth"}
oauth_call_counter = {"count": 0}
def fake_oauth(headers):
oauth_call_counter["count"] += 1
return oauth_headers
monkeypatch.setattr(
auth_mcp.MCPRequestHandler,
"_get_oauth2_headers_from_headers",
staticmethod(fake_oauth),
raising=False,
)
request = _build_request({"authorization": "Bearer incoming"})
payload = NewMCPServerRequest(
server_name="example",
url="https://example.com",
auth_type=MCPAuth.oauth2,
)
result = await rest_endpoints.test_tools_list(
request, payload, user_api_key_dict=UserAPIKeyAuth()
)
assert result["message"] == "Successfully retrieved tools"
assert captured["mcp_auth_header"] is None
assert captured["oauth2_headers"] == oauth_headers
assert oauth_call_counter["count"] == 1
class TestListToolsRestAPI:
pytestmark = pytest.mark.asyncio
async def test_rejects_disallowed_server(self, monkeypatch):
async def fake_contexts(user_api_key_auth):
return [user_api_key_auth]
async def fake_get_allowed_mcp_servers(*args, **kwargs):
return []
monkeypatch.setattr(
rest_endpoints,
"build_effective_auth_contexts",
fake_contexts,
raising=False,
)
monkeypatch.setattr(
rest_endpoints.global_mcp_server_manager,
"get_allowed_mcp_servers",
fake_get_allowed_mcp_servers,
raising=False,
)
request = _build_request(path="/mcp-rest/tools/list", method="GET")
result = await rest_endpoints.list_tool_rest_api(
request,
server_id="server-1",
user_api_key_dict=UserAPIKeyAuth(),
)
assert result["tools"] == []
assert result["error"] == "unexpected_error"
assert "access_denied" in result["message"]
assert "server server-1" in result["message"]
async def test_lists_tools_for_allowed_server(self, monkeypatch):
async def fake_contexts(user_api_key_auth):
return [user_api_key_auth]
async def fake_get_allowed_mcp_servers(*args, **kwargs):
return ["server-1"]
class StubServer:
alias = "server-1"
server_name = "server-1"
name = "stub"
allowed_tools = None
mcp_info = {"server_name": "stub"}
stub_server = StubServer()
captured = {"called": False}
async def fake_get_tools(server, server_auth_header, raw_headers=None):
captured["called"] = True
captured["server"] = server
captured["auth_header"] = server_auth_header
return ["tool-1"]
monkeypatch.setattr(
rest_endpoints,
"build_effective_auth_contexts",
fake_contexts,
raising=False,
)
monkeypatch.setattr(
rest_endpoints.global_mcp_server_manager,
"get_allowed_mcp_servers",
fake_get_allowed_mcp_servers,
raising=False,
)
monkeypatch.setattr(
rest_endpoints.global_mcp_server_manager,
"get_mcp_server_by_id",
lambda server_id: stub_server if server_id == "server-1" else None,
raising=False,
)
monkeypatch.setattr(
rest_endpoints,
"_get_tools_for_single_server",
fake_get_tools,
raising=False,
)
request = _build_request(path="/mcp-rest/tools/list", method="GET")
result = await rest_endpoints.list_tool_rest_api(
request,
server_id="server-1",
user_api_key_dict=UserAPIKeyAuth(),
)
assert captured["called"] is True
assert captured["server"] is stub_server
assert result["tools"] == ["tool-1"]
assert result["error"] is None
assert result["message"] == "Successfully retrieved tools"
class TestCallToolRestAPI:
pytestmark = pytest.mark.asyncio
async def test_rejects_disallowed_server(self, monkeypatch):
async def fake_contexts(user_api_key_auth):
return [user_api_key_auth]
async def fake_get_allowed_mcp_servers(*args, **kwargs):
return []
async def fake_add_litellm_data_to_request(**kwargs):
return kwargs.get("data", {})
monkeypatch.setattr(
rest_endpoints,
"build_effective_auth_contexts",
fake_contexts,
raising=False,
)
monkeypatch.setattr(
rest_endpoints.global_mcp_server_manager,
"get_allowed_mcp_servers",
fake_get_allowed_mcp_servers,
raising=False,
)
monkeypatch.setattr(
"litellm.proxy.proxy_server.add_litellm_data_to_request",
fake_add_litellm_data_to_request,
raising=False,
)
request_payload = {
"server_id": "server-1",
"name": "demo-tool",
"arguments": {"foo": "bar"},
}
request = _build_request(
path="/mcp-rest/tools/call",
method="POST",
json_body=request_payload,
)
monkeypatch.setattr(
rest_endpoints, "_execute_with_mcp_client", fake_execute, raising=False
)
with pytest.raises(HTTPException) as exc_info:
await rest_endpoints.call_tool_rest_api(
request,
user_api_key_dict=UserAPIKeyAuth(),
)
oauth_call_counter = {"count": 0}
assert exc_info.value.status_code == 403
assert exc_info.value.detail["error"] == "access_denied"
assert "server server-1" in exc_info.value.detail["message"]
def fake_oauth(headers):
oauth_call_counter["count"] += 1
return {"Authorization": "Bearer oauth"}
async def test_executes_tool_when_allowed(self, monkeypatch):
async def fake_contexts(user_api_key_auth):
return [user_api_key_auth]
monkeypatch.setattr(
auth_mcp.MCPRequestHandler,
"_get_oauth2_headers_from_headers",
staticmethod(fake_oauth),
raising=False,
)
async def fake_get_allowed_mcp_servers(*args, **kwargs):
return ["server-1"]
request = _build_request()
payload = NewMCPServerRequest(
server_name="example",
url="https://example.com",
auth_type=MCPAuth.api_key,
credentials={"auth_value": "secret-key"},
)
class StubServer:
alias = "server-1"
server_name = "server-1"
name = "stub"
allowed_tools = None
mcp_info = {"server_name": "stub"}
result = await rest_endpoints.test_tools_list(
request, payload, user_api_key_dict=UserAPIKeyAuth()
)
stub_server = StubServer()
assert result["message"] == "Successfully retrieved tools"
assert captured["mcp_auth_header"] == "secret-key"
assert captured["oauth2_headers"] is None
assert oauth_call_counter["count"] == 0
async def fake_add_litellm_data_to_request(**kwargs):
return kwargs.get("data", {})
captured = {}
@pytest.mark.asyncio
async def test_test_tools_list_extracts_oauth2_headers(monkeypatch):
"""Ensure oauth2 auth type pulls oauth headers and omits MCP auth header."""
async def fake_execute_mcp_tool(**kwargs):
captured.update(kwargs)
return {"result": "ok"}
captured: dict = {}
monkeypatch.setattr(
rest_endpoints,
"build_effective_auth_contexts",
fake_contexts,
raising=False,
)
monkeypatch.setattr(
rest_endpoints.global_mcp_server_manager,
"get_allowed_mcp_servers",
fake_get_allowed_mcp_servers,
raising=False,
)
monkeypatch.setattr(
rest_endpoints.global_mcp_server_manager,
"get_mcp_server_by_id",
lambda server_id: stub_server if server_id == "server-1" else None,
raising=False,
)
monkeypatch.setattr(
"litellm.proxy.proxy_server.add_litellm_data_to_request",
fake_add_litellm_data_to_request,
raising=False,
)
monkeypatch.setattr(
"litellm.proxy.proxy_server.proxy_config",
{},
raising=False,
)
monkeypatch.setattr(
rest_endpoints,
"execute_mcp_tool",
fake_execute_mcp_tool,
raising=False,
)
async def fake_execute(
request,
operation,
mcp_auth_header=None,
oauth2_headers=None,
raw_headers=None,
):
captured["mcp_auth_header"] = mcp_auth_header
captured["oauth2_headers"] = oauth2_headers
return {
"tools": [],
"error": None,
"message": "Successfully retrieved tools",
request_payload = {
"server_id": "server-1",
"name": "demo-tool",
"arguments": {"foo": "bar"},
}
request = _build_request(
path="/mcp-rest/tools/call",
method="POST",
json_body=request_payload,
)
monkeypatch.setattr(
rest_endpoints, "_execute_with_mcp_client", fake_execute, raising=False
)
result = await rest_endpoints.call_tool_rest_api(
request,
user_api_key_dict=UserAPIKeyAuth(),
)
oauth_headers = {"Authorization": "Bearer oauth"}
oauth_call_counter = {"count": 0}
def fake_oauth(headers):
oauth_call_counter["count"] += 1
return oauth_headers
monkeypatch.setattr(
auth_mcp.MCPRequestHandler,
"_get_oauth2_headers_from_headers",
staticmethod(fake_oauth),
raising=False,
)
request = _build_request({"authorization": "Bearer incoming"})
payload = NewMCPServerRequest(
server_name="example",
url="https://example.com",
auth_type=MCPAuth.oauth2,
)
result = await rest_endpoints.test_tools_list(
request, payload, user_api_key_dict=UserAPIKeyAuth()
)
assert result["message"] == "Successfully retrieved tools"
assert captured["mcp_auth_header"] is None
assert captured["oauth2_headers"] == oauth_headers
assert oauth_call_counter["count"] == 1
assert result == {"result": "ok"}
assert captured["name"] == "demo-tool"
assert captured["arguments"] == {"foo": "bar"}
assert captured["allowed_mcp_servers"] == [stub_server]

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