test: add permission test

This commit is contained in:
Yuta Saito 2026-01-14 07:58:21 +09:00
parent 22aad95bb1
commit df37770a70
3 changed files with 472 additions and 148 deletions

View file

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

View file

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

View file

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