mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
test: add permission test
This commit is contained in:
parent
22aad95bb1
commit
df37770a70
3 changed files with 472 additions and 148 deletions
|
|
@ -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]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue