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..48f7a8b0b7b 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -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, ) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 2adfe97c611..f22040a7dd9 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -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, diff --git a/tests/mcp_tests/test_mcp_logging.py b/tests/mcp_tests/test_mcp_logging.py index d27626dd447..aeebb2e913e 100644 --- a/tests/mcp_tests/test_mcp_logging.py +++ b/tests/mcp_tests/test_mcp_logging.py @@ -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 /) diff --git a/tests/mcp_tests/test_mcp_server.py b/tests/mcp_tests/test_mcp_server.py index 8fb0e80cc39..66dd2bdc6b3 100644 --- a/tests/mcp_tests/test_mcp_server.py +++ b/tests/mcp_tests/test_mcp_server.py @@ -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 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py index 0c6d0921952..c026ea232b7 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py @@ -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] 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, }),