mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge pull request #19051 from BerriAI/litellm_fix_mcp-rest-auth-checks
[fix] mcp rest auth checks
This commit is contained in:
commit
a66e007574
8 changed files with 677 additions and 192 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,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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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 /)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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