mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
add mcp tests
This commit is contained in:
parent
3448bf79f8
commit
3d22e0197f
5 changed files with 913 additions and 12 deletions
|
|
@ -1,11 +1,112 @@
|
|||
# math_server.py
|
||||
import argparse
|
||||
import os
|
||||
from typing import Optional
|
||||
|
||||
from mcp.server.fastmcp import FastMCP
|
||||
from starlette.requests import Request
|
||||
from starlette.responses import JSONResponse
|
||||
from starlette.types import ASGIApp, Receive, Scope, Send
|
||||
|
||||
DEFAULT_API_KEY = "test-api-key"
|
||||
DEFAULT_BEARER_TOKEN = "test-bearer-token"
|
||||
DEFAULT_AUTHORIZATION_VALUE = "Custom raw-auth-value"
|
||||
DEFAULT_CUSTOM_HEADER = "x-custom-token"
|
||||
DEFAULT_CUSTOM_HEADER_VALUE = "custom-header-value"
|
||||
DEFAULT_CLIENT_ID = "test-client"
|
||||
DEFAULT_CLIENT_SECRET = "test-secret"
|
||||
DEFAULT_OAUTH_ACCESS_TOKEN = "test-oauth-access-token"
|
||||
DEFAULT_OBO_ACCESS_TOKEN = "test-obo-exchanged-token"
|
||||
|
||||
TOKEN_EXCHANGE_GRANT_TYPE = "urn:ietf:params:oauth:grant-type:token-exchange"
|
||||
VALID_OAUTH_BEARER_TOKENS = {DEFAULT_OAUTH_ACCESS_TOKEN, DEFAULT_OBO_ACCESS_TOKEN}
|
||||
|
||||
mcp = FastMCP("Math")
|
||||
|
||||
_auth_mode = "none"
|
||||
_auth_secret: Optional[str] = None
|
||||
_client_id = DEFAULT_CLIENT_ID
|
||||
_client_secret = DEFAULT_CLIENT_SECRET
|
||||
|
||||
|
||||
def _request_is_authorized(headers) -> bool:
|
||||
if _auth_mode == "none":
|
||||
return True
|
||||
if _auth_mode == "api_key":
|
||||
return headers.get("x-api-key") == _auth_secret
|
||||
if _auth_mode == "bearer_token":
|
||||
return headers.get("authorization") == f"Bearer {_auth_secret}"
|
||||
if _auth_mode == "authorization":
|
||||
return headers.get("authorization") == _auth_secret
|
||||
if _auth_mode == "custom_header":
|
||||
return headers.get(DEFAULT_CUSTOM_HEADER) == _auth_secret
|
||||
if _auth_mode == "oauth2":
|
||||
auth = headers.get("authorization") or ""
|
||||
return auth.startswith("Bearer ") and auth[len("Bearer ") :] in VALID_OAUTH_BEARER_TOKENS
|
||||
return False
|
||||
|
||||
|
||||
class _AuthMiddleware:
|
||||
def __init__(self, app: ASGIApp) -> None:
|
||||
self.app = app
|
||||
|
||||
async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
|
||||
if scope["type"] != "http":
|
||||
await self.app(scope, receive, send)
|
||||
return
|
||||
request = Request(scope)
|
||||
if request.url.path.startswith("/oauth/"):
|
||||
await self.app(scope, receive, send)
|
||||
return
|
||||
if not _request_is_authorized(request.headers):
|
||||
response = JSONResponse({"error": "unauthorized", "auth_mode": _auth_mode}, status_code=401)
|
||||
await response(scope, receive, send)
|
||||
return
|
||||
await self.app(scope, receive, send)
|
||||
|
||||
|
||||
@mcp.tool()
|
||||
def add(a: int, b: int) -> int:
|
||||
"""Add two numbers"""
|
||||
return a + b
|
||||
|
||||
|
||||
@mcp.tool()
|
||||
def multiply(a: int, b: int) -> int:
|
||||
"""Multiply two numbers"""
|
||||
return a * b
|
||||
|
||||
|
||||
@mcp.custom_route("/oauth/token", methods=["POST"])
|
||||
async def oauth_token(request: Request) -> JSONResponse:
|
||||
form = await request.form()
|
||||
grant_type = form.get("grant_type")
|
||||
|
||||
if form.get("client_id") != _client_id or form.get("client_secret") != _client_secret:
|
||||
return JSONResponse({"error": "invalid_client"}, status_code=401)
|
||||
|
||||
if grant_type == "client_credentials":
|
||||
return JSONResponse(
|
||||
{
|
||||
"access_token": DEFAULT_OAUTH_ACCESS_TOKEN,
|
||||
"token_type": "Bearer",
|
||||
"expires_in": 3600,
|
||||
}
|
||||
)
|
||||
|
||||
if grant_type == TOKEN_EXCHANGE_GRANT_TYPE:
|
||||
if not form.get("subject_token"):
|
||||
return JSONResponse({"error": "invalid_request"}, status_code=400)
|
||||
return JSONResponse(
|
||||
{
|
||||
"access_token": DEFAULT_OBO_ACCESS_TOKEN,
|
||||
"token_type": "Bearer",
|
||||
"expires_in": 3600,
|
||||
}
|
||||
)
|
||||
|
||||
return JSONResponse({"error": "unsupported_grant_type"}, status_code=400)
|
||||
|
||||
|
||||
def _parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(description="MCP math test server")
|
||||
|
|
@ -25,25 +126,23 @@ def _parse_args() -> argparse.Namespace:
|
|||
default=int(os.getenv("MCP_PORT", "0")),
|
||||
help="Port to bind when serving over HTTP",
|
||||
)
|
||||
parser.add_argument("--auth-mode", default="none")
|
||||
parser.add_argument("--auth-secret", default=None)
|
||||
parser.add_argument("--client-id", default=DEFAULT_CLIENT_ID)
|
||||
parser.add_argument("--client-secret", default=DEFAULT_CLIENT_SECRET)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
@mcp.tool()
|
||||
def add(a: int, b: int) -> int:
|
||||
"""Add two numbers"""
|
||||
return a + b
|
||||
|
||||
|
||||
@mcp.tool()
|
||||
def multiply(a: int, b: int) -> int:
|
||||
"""Multiply two numbers"""
|
||||
return a * b
|
||||
|
||||
|
||||
def main() -> None:
|
||||
global _auth_mode, _auth_secret, _client_id, _client_secret
|
||||
args = _parse_args()
|
||||
transport = (args.transport or "stdio").lower()
|
||||
|
||||
_auth_mode = args.auth_mode
|
||||
_auth_secret = args.auth_secret
|
||||
_client_id = args.client_id
|
||||
_client_secret = args.client_secret
|
||||
|
||||
if transport == "stdio":
|
||||
mcp.run(transport="stdio")
|
||||
return
|
||||
|
|
@ -53,6 +152,13 @@ def main() -> None:
|
|||
raise ValueError("HTTP transport requires a valid --port value")
|
||||
mcp.settings.host = args.host
|
||||
mcp.settings.port = args.port
|
||||
|
||||
original_app = mcp.streamable_http_app
|
||||
|
||||
def _app_with_auth():
|
||||
return _AuthMiddleware(original_app())
|
||||
|
||||
mcp.streamable_http_app = _app_with_auth # type: ignore[method-assign]
|
||||
mcp.run(transport="streamable-http")
|
||||
return
|
||||
|
||||
|
|
|
|||
74
tests/mcp_tests/test_configs/test_oauth2_mcp_config.yaml
Normal file
74
tests/mcp_tests/test_configs/test_oauth2_mcp_config.yaml
Normal file
|
|
@ -0,0 +1,74 @@
|
|||
model_list:
|
||||
- model_name: fake-model
|
||||
litellm_params:
|
||||
model: openai/fake
|
||||
api_key: fake-key
|
||||
|
||||
general_settings:
|
||||
mcp_internal_ip_ranges:
|
||||
- "10.0.0.0/8"
|
||||
- "192.168.0.0/16"
|
||||
- "0.0.0.0/0"
|
||||
|
||||
mcp_servers:
|
||||
math_no_auth:
|
||||
url: "http://localhost:0/mcp"
|
||||
transport: "http"
|
||||
|
||||
math_api_key:
|
||||
url: "http://localhost:0/mcp"
|
||||
transport: "http"
|
||||
auth_type: "api_key"
|
||||
auth_value: "test-api-key"
|
||||
|
||||
math_bearer_token:
|
||||
url: "http://localhost:0/mcp"
|
||||
transport: "http"
|
||||
auth_type: "bearer_token"
|
||||
auth_value: "test-bearer-token"
|
||||
|
||||
math_authorization:
|
||||
url: "http://localhost:0/mcp"
|
||||
transport: "http"
|
||||
auth_type: "authorization"
|
||||
auth_value: "Custom raw-auth-value"
|
||||
|
||||
math_custom_header:
|
||||
url: "http://localhost:0/mcp"
|
||||
transport: "http"
|
||||
extra_headers:
|
||||
- "x-custom-token"
|
||||
|
||||
# Machine to Machine (M2M)
|
||||
test_oauth2_server:
|
||||
url: "http://localhost:0/mcp"
|
||||
transport: "http"
|
||||
auth_type: "oauth2"
|
||||
client_id: "test-client"
|
||||
client_secret: "test-secret"
|
||||
token_url: "http://localhost:0/oauth/token"
|
||||
|
||||
# Interactive PKCE
|
||||
test_pkce_server:
|
||||
url: "http://localhost:0/mcp"
|
||||
transport: "http"
|
||||
auth_type: "oauth2"
|
||||
client_id: "test-client"
|
||||
client_secret: "test-secret"
|
||||
token_url: "http://localhost:0/oauth/token"
|
||||
authorization_url: "http://localhost:0/oauth/authorize"
|
||||
scopes: ["create_ticket", "update_ticket", "delete_ticket"]
|
||||
|
||||
# MCP OBO (OAuth2 Token Exchange, RFC 8693)
|
||||
internal_mcp_server:
|
||||
url: "http://localhost:0/mcp"
|
||||
transport: "http"
|
||||
auth_type: "oauth2_token_exchange"
|
||||
token_exchange_endpoint: "http://localhost:0/oauth/token"
|
||||
client_id: "test-client"
|
||||
client_secret: "test-secret"
|
||||
audience: "api://internal-tools-mcp"
|
||||
scopes:
|
||||
- "mcp.tools.read"
|
||||
- "mcp.tools.execute"
|
||||
subject_token_type: "urn:ietf:params:oauth:token-type:access_token"
|
||||
|
|
@ -446,3 +446,124 @@ async def test_mcp_tool_call_hook():
|
|||
logged_standard_logging_payload is not None
|
||||
), "Standard logging payload should not be None"
|
||||
assert logged_standard_logging_payload["response_cost"] == 1.42
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_auto_executed_mcp_tool_call_attributes_cost_to_key_and_team():
|
||||
"""
|
||||
Regression test: when an MCP tool is auto-executed via the responses handler
|
||||
(`_execute_tool_calls`), the resulting spend row must charge the cost to the
|
||||
calling key/team/user/end_user. Previously this path hand-rolled the logging
|
||||
metadata and dropped the attribution fields, so SpendLogs recorded the cost
|
||||
with NULL user/team_id/end_user (cost charged to nobody).
|
||||
|
||||
Asserting on the StandardLoggingPayload couples the two facts that matter for
|
||||
the spend table: `response_cost` (the spend column) and the
|
||||
`user_api_key_*` metadata (the user/team_id/end_user columns).
|
||||
"""
|
||||
from litellm.responses.mcp.litellm_proxy_mcp_handler import (
|
||||
LiteLLM_Proxy_MCP_Handler,
|
||||
)
|
||||
|
||||
litellm.logging_callback_manager._reset_all_callbacks()
|
||||
mock_result = CallToolResult(
|
||||
content=[TextContent(type="text", text="Test response")], isError=False
|
||||
)
|
||||
|
||||
mock_client = AsyncMock()
|
||||
mock_client.call_tool = AsyncMock(return_value=mock_result)
|
||||
mock_client.list_tools = AsyncMock(
|
||||
return_value=[
|
||||
MCPTool(
|
||||
name="add_tools",
|
||||
description="Test tool",
|
||||
inputSchema={
|
||||
"type": "object",
|
||||
"properties": {"test": {"type": "string"}},
|
||||
},
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
def mock_client_constructor(*args, **kwargs):
|
||||
return mock_client
|
||||
|
||||
local_mcp_server_manager = MCPServerManager()
|
||||
|
||||
expected_cost = 1.2
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient",
|
||||
mock_client_constructor,
|
||||
):
|
||||
await local_mcp_server_manager.load_servers_from_config(
|
||||
mcp_servers_config={
|
||||
"zapier_gmail_server": {
|
||||
"url": os.getenv("ZAPIER_MCP_HTTPS_SERVER_URL"),
|
||||
"mcp_info": {
|
||||
"mcp_server_cost_info": {
|
||||
"default_cost_per_query": expected_cost,
|
||||
}
|
||||
},
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
test_logger = TestMCPLogger()
|
||||
litellm.callbacks = [test_logger]
|
||||
|
||||
await local_mcp_server_manager._initialize_tool_name_to_mcp_server_name_mapping()
|
||||
local_mcp_server_manager.tool_name_to_mcp_server_name_mapping[
|
||||
"add_tools"
|
||||
] = "zapier_gmail_server"
|
||||
local_mcp_server_manager.tool_name_to_mcp_server_name_mapping[
|
||||
"zapier_gmail_server-add_tools"
|
||||
] = "zapier_gmail_server"
|
||||
|
||||
server_ids = local_mcp_server_manager.get_all_mcp_server_ids()
|
||||
user_auth = UserAPIKeyAuth(
|
||||
api_key="sk-attribution-test",
|
||||
user_id="user-123",
|
||||
team_id="team-456",
|
||||
end_user_id="end-user-789",
|
||||
object_permission=LiteLLM_ObjectPermissionTable(
|
||||
object_permission_id="mcp-test-permissions",
|
||||
mcp_servers=list(server_ids),
|
||||
),
|
||||
)
|
||||
|
||||
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,
|
||||
),
|
||||
):
|
||||
tool_name = "zapier_gmail_server-add_tools"
|
||||
await LiteLLM_Proxy_MCP_Handler._execute_tool_calls(
|
||||
tool_server_map={tool_name: "zapier_gmail_server"},
|
||||
tool_calls=[
|
||||
{
|
||||
"id": "call-1",
|
||||
"function": {"name": tool_name, "arguments": "{}"},
|
||||
}
|
||||
],
|
||||
user_api_key_auth=user_auth,
|
||||
litellm_call_id="cid",
|
||||
litellm_trace_id="tid",
|
||||
)
|
||||
|
||||
await asyncio.sleep(2)
|
||||
|
||||
payload = test_logger.standard_logging_payload
|
||||
assert payload is not None, "Standard logging payload should not be None"
|
||||
|
||||
assert payload["response_cost"] == expected_cost
|
||||
|
||||
metadata = payload["metadata"]
|
||||
assert metadata["user_api_key_user_id"] == "user-123"
|
||||
assert metadata["user_api_key_team_id"] == "team-456"
|
||||
assert metadata["user_api_key_end_user_id"] == "end-user-789"
|
||||
|
|
|
|||
343
tests/mcp_tests/test_proxy_mcp_auth_e2e.py
Normal file
343
tests/mcp_tests/test_proxy_mcp_auth_e2e.py
Normal file
|
|
@ -0,0 +1,343 @@
|
|||
import asyncio
|
||||
import socket
|
||||
import subprocess
|
||||
import sys
|
||||
import threading
|
||||
import time
|
||||
import typing
|
||||
from pathlib import Path
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
import uvicorn
|
||||
import yaml
|
||||
from mcp import ClientSession
|
||||
from mcp.client.streamable_http import streamablehttp_client
|
||||
|
||||
from litellm.proxy.proxy_server import (
|
||||
app as proxy_app,
|
||||
cleanup_router_config_variables,
|
||||
initialize,
|
||||
)
|
||||
from tests.mcp_tests.mcp_server import (
|
||||
DEFAULT_API_KEY,
|
||||
DEFAULT_AUTHORIZATION_VALUE,
|
||||
DEFAULT_BEARER_TOKEN,
|
||||
DEFAULT_CLIENT_ID,
|
||||
DEFAULT_CLIENT_SECRET,
|
||||
DEFAULT_CUSTOM_HEADER,
|
||||
DEFAULT_CUSTOM_HEADER_VALUE,
|
||||
)
|
||||
|
||||
CONFIG_TEMPLATE_PATH = Path("tests/mcp_tests/test_configs/test_oauth2_mcp_config.yaml")
|
||||
MCP_SERVER_SCRIPT = Path("tests/mcp_tests/mcp_server.py")
|
||||
PROJECT_ROOT = Path(__file__).resolve().parents[2]
|
||||
PROXY_START_TIMEOUT = 30
|
||||
PROXY_AUTHORIZATION_HEADER = "Bearer sk-1234"
|
||||
|
||||
# Each entry: proxy server name -> how to launch the upstream test MCP server.
|
||||
SERVER_SPECS: dict[str, dict[str, typing.Optional[str]]] = {
|
||||
"math_no_auth": {"auth_mode": "none", "auth_secret": None},
|
||||
"math_api_key": {"auth_mode": "api_key", "auth_secret": DEFAULT_API_KEY},
|
||||
"math_bearer_token": {
|
||||
"auth_mode": "bearer_token",
|
||||
"auth_secret": DEFAULT_BEARER_TOKEN,
|
||||
},
|
||||
"math_authorization": {
|
||||
"auth_mode": "authorization",
|
||||
"auth_secret": DEFAULT_AUTHORIZATION_VALUE,
|
||||
},
|
||||
"math_custom_header": {
|
||||
"auth_mode": "custom_header",
|
||||
"auth_secret": DEFAULT_CUSTOM_HEADER_VALUE,
|
||||
},
|
||||
"test_oauth2_server": {"auth_mode": "oauth2", "auth_secret": None},
|
||||
"test_pkce_server": {"auth_mode": "oauth2", "auth_secret": None},
|
||||
"internal_mcp_server": {"auth_mode": "oauth2", "auth_secret": None},
|
||||
}
|
||||
|
||||
|
||||
def _initialize_proxy(config_path: str) -> None:
|
||||
cleanup_router_config_variables()
|
||||
asyncio.run(initialize(config=config_path, debug=True))
|
||||
|
||||
|
||||
def _start_proxy_server(
|
||||
config_path: str,
|
||||
) -> tuple[str, uvicorn.Server, threading.Thread, socket.socket]:
|
||||
_initialize_proxy(config_path)
|
||||
|
||||
sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
|
||||
sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
|
||||
sock.bind(("127.0.0.1", 0))
|
||||
host, port = sock.getsockname()
|
||||
|
||||
config = uvicorn.Config(proxy_app, host=host, port=port, log_level="warning")
|
||||
server = uvicorn.Server(config)
|
||||
|
||||
def _run() -> None:
|
||||
loop = asyncio.new_event_loop()
|
||||
asyncio.set_event_loop(loop)
|
||||
loop.run_until_complete(server.serve(sockets=[sock]))
|
||||
|
||||
thread = threading.Thread(target=_run, daemon=True)
|
||||
thread.start()
|
||||
|
||||
start_time = time.time()
|
||||
while not server.started:
|
||||
if not thread.is_alive():
|
||||
raise RuntimeError("Proxy server failed to start")
|
||||
if time.time() - start_time > PROXY_START_TIMEOUT:
|
||||
raise TimeoutError("Proxy server did not start in time")
|
||||
time.sleep(0.05)
|
||||
|
||||
return f"http://{host}:{port}", server, thread, sock
|
||||
|
||||
|
||||
def _reserve_port(host: str = "127.0.0.1") -> int:
|
||||
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as sock:
|
||||
sock.bind((host, 0))
|
||||
return sock.getsockname()[1]
|
||||
|
||||
|
||||
def _start_mcp_server_process(*, auth_mode: str, port: int, auth_secret: typing.Optional[str]) -> subprocess.Popen:
|
||||
cmd = [
|
||||
sys.executable,
|
||||
str(MCP_SERVER_SCRIPT),
|
||||
"--transport",
|
||||
"http",
|
||||
"--host",
|
||||
"127.0.0.1",
|
||||
"--port",
|
||||
str(port),
|
||||
"--auth-mode",
|
||||
auth_mode,
|
||||
"--client-id",
|
||||
DEFAULT_CLIENT_ID,
|
||||
"--client-secret",
|
||||
DEFAULT_CLIENT_SECRET,
|
||||
]
|
||||
if auth_secret is not None:
|
||||
cmd.extend(["--auth-secret", auth_secret])
|
||||
|
||||
process = subprocess.Popen(cmd, cwd=str(PROJECT_ROOT), stdout=subprocess.PIPE, stderr=subprocess.PIPE)
|
||||
|
||||
start_time = time.time()
|
||||
while True:
|
||||
if process.poll() is not None:
|
||||
stdout, stderr = process.communicate()
|
||||
raise RuntimeError(
|
||||
f"MCP server exited early (auth_mode={auth_mode}).\n"
|
||||
f"STDOUT: {stdout.decode()}\nSTDERR: {stderr.decode()}"
|
||||
)
|
||||
try:
|
||||
with socket.create_connection(("127.0.0.1", port), timeout=0.1):
|
||||
break
|
||||
except OSError:
|
||||
if time.time() - start_time > PROXY_START_TIMEOUT:
|
||||
process.terminate()
|
||||
raise TimeoutError(f"MCP server did not start in time (auth_mode={auth_mode})")
|
||||
time.sleep(0.05)
|
||||
|
||||
return process
|
||||
|
||||
|
||||
@pytest.fixture(scope="session", autouse=True)
|
||||
def _clear_proxy_database_env() -> typing.Iterator[None]:
|
||||
mp = pytest.MonkeyPatch()
|
||||
mp.delenv("DATABASE_URL", raising=False)
|
||||
mp.setenv("LITELLM_MASTER_KEY", "sk-1234")
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
mp.undo()
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def mcp_auth_servers() -> typing.Iterator[dict[str, typing.Any]]:
|
||||
servers = {name: {**spec, "port": _reserve_port()} for name, spec in SERVER_SPECS.items()}
|
||||
|
||||
processes: list[subprocess.Popen] = []
|
||||
try:
|
||||
for spec in servers.values():
|
||||
process = _start_mcp_server_process(
|
||||
auth_mode=spec["auth_mode"],
|
||||
port=spec["port"],
|
||||
auth_secret=spec["auth_secret"],
|
||||
)
|
||||
spec["process"] = process
|
||||
spec["base_url"] = f"http://127.0.0.1:{spec['port']}"
|
||||
processes.append(process)
|
||||
yield servers
|
||||
finally:
|
||||
for process in processes:
|
||||
process.terminate()
|
||||
try:
|
||||
process.wait(timeout=5)
|
||||
except subprocess.TimeoutExpired:
|
||||
process.kill()
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def proxy_server_url(
|
||||
tmp_path_factory: pytest.TempPathFactory, mcp_auth_servers: dict[str, typing.Any]
|
||||
) -> typing.Iterator[str]:
|
||||
config = yaml.safe_load(CONFIG_TEMPLATE_PATH.read_text())
|
||||
|
||||
for server_name, spec in mcp_auth_servers.items():
|
||||
server_config = config["mcp_servers"][server_name]
|
||||
base_url = spec["base_url"]
|
||||
server_config["url"] = f"{base_url}/mcp"
|
||||
for endpoint_key in ("token_url", "authorization_url", "token_exchange_endpoint"):
|
||||
if endpoint_key in server_config:
|
||||
suffix = "authorize" if "authorization" in endpoint_key else "token"
|
||||
server_config[endpoint_key] = f"{base_url}/oauth/{suffix}"
|
||||
|
||||
config_path = tmp_path_factory.mktemp("mcp_auth_e2e") / "config.yaml"
|
||||
config_path.write_text(yaml.safe_dump(config))
|
||||
|
||||
server_url, server, thread, sock = _start_proxy_server(str(config_path))
|
||||
yield server_url
|
||||
|
||||
server.should_exit = True
|
||||
thread.join(timeout=10)
|
||||
sock.close()
|
||||
|
||||
|
||||
async def _call_add_tool(
|
||||
*,
|
||||
proxy_server_url: str,
|
||||
server_name: str,
|
||||
a: int,
|
||||
b: int,
|
||||
headers: typing.Optional[dict[str, str]] = None,
|
||||
) -> typing.Optional[str]:
|
||||
request_headers = {
|
||||
"Authorization": PROXY_AUTHORIZATION_HEADER,
|
||||
"x-mcp-servers": server_name,
|
||||
}
|
||||
if headers:
|
||||
request_headers.update(headers)
|
||||
|
||||
async with asyncio.timeout(20):
|
||||
async with streamablehttp_client(url=f"{proxy_server_url}/mcp", headers=request_headers) as (
|
||||
read,
|
||||
write,
|
||||
_get_session_id,
|
||||
):
|
||||
async with ClientSession(read, write) as session:
|
||||
await session.initialize()
|
||||
tools_result = await session.list_tools()
|
||||
assert any(tool.name.endswith("add") for tool in tools_result.tools)
|
||||
|
||||
result = await session.call_tool("add", arguments={"a": a, "b": b})
|
||||
assert result.content
|
||||
return getattr(result.content[0], "text", None)
|
||||
|
||||
|
||||
async def _list_tool_names(
|
||||
*,
|
||||
proxy_server_url: str,
|
||||
server_name: str,
|
||||
headers: typing.Optional[dict[str, str]] = None,
|
||||
) -> list[str]:
|
||||
request_headers = {
|
||||
"Authorization": PROXY_AUTHORIZATION_HEADER,
|
||||
"x-mcp-servers": server_name,
|
||||
}
|
||||
if headers:
|
||||
request_headers.update(headers)
|
||||
|
||||
async with asyncio.timeout(20):
|
||||
async with streamablehttp_client(url=f"{proxy_server_url}/mcp", headers=request_headers) as (
|
||||
read,
|
||||
write,
|
||||
_get_session_id,
|
||||
):
|
||||
async with ClientSession(read, write) as session:
|
||||
await session.initialize()
|
||||
tools_result = await session.list_tools()
|
||||
return [tool.name for tool in tools_result.tools]
|
||||
|
||||
|
||||
class TestProxyMcpAuthE2E:
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("server_name", "a", "b", "expected"),
|
||||
[
|
||||
("math_no_auth", 3, 4, "7"),
|
||||
("math_api_key", 5, 6, "11"),
|
||||
("math_bearer_token", 7, 8, "15"),
|
||||
("math_authorization", 1, 2, "3"),
|
||||
("test_oauth2_server", 9, 10, "19"),
|
||||
],
|
||||
)
|
||||
async def test_proxy_forwards_configured_credential(self, proxy_server_url, server_name, a, b, expected) -> None:
|
||||
result = await _call_add_tool(proxy_server_url=proxy_server_url, server_name=server_name, a=a, b=b)
|
||||
assert result == expected
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"server_name",
|
||||
["math_api_key", "math_bearer_token", "math_authorization", "test_oauth2_server"],
|
||||
)
|
||||
async def test_upstream_rejects_unauthenticated_request(self, mcp_auth_servers, server_name) -> None:
|
||||
base_url = mcp_auth_servers[server_name]["base_url"]
|
||||
async with httpx.AsyncClient() as client:
|
||||
response = await client.post(
|
||||
f"{base_url}/mcp",
|
||||
json={"jsonrpc": "2.0", "id": 1, "method": "initialize", "params": {}},
|
||||
)
|
||||
assert response.status_code == 401
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_oauth_m2m_ignores_caller_authorization(self, proxy_server_url) -> None:
|
||||
"""M2M servers must never forward the caller's Authorization; the proxy
|
||||
fetches its own client_credentials token. A bogus caller token must not
|
||||
break the call (proves the proxy substitutes its own upstream token)."""
|
||||
result = await _call_add_tool(
|
||||
proxy_server_url=proxy_server_url,
|
||||
server_name="test_oauth2_server",
|
||||
a=2,
|
||||
b=3,
|
||||
headers={
|
||||
"x-litellm-api-key": "Bearer sk-1234",
|
||||
"Authorization": "Bearer caller-supplied-bogus-token",
|
||||
},
|
||||
)
|
||||
assert result == "5"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_custom_header_passthrough(self, proxy_server_url) -> None:
|
||||
result = await _call_add_tool(
|
||||
proxy_server_url=proxy_server_url,
|
||||
server_name="math_custom_header",
|
||||
a=4,
|
||||
b=5,
|
||||
headers={f"x-mcp-math_custom_header-{DEFAULT_CUSTOM_HEADER}": DEFAULT_CUSTOM_HEADER_VALUE},
|
||||
)
|
||||
assert result == "9"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_custom_header_required_for_discovery(self, proxy_server_url) -> None:
|
||||
"""Without the per-server custom header the upstream rejects the request,
|
||||
so its tools must not be discoverable through the proxy."""
|
||||
tool_names = await _list_tool_names(proxy_server_url=proxy_server_url, server_name="math_custom_header")
|
||||
assert not any(name.endswith("add") for name in tool_names)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_obo_token_exchange(self, proxy_server_url) -> None:
|
||||
"""OBO: the proxy exchanges the caller's bearer (subject_token) for a
|
||||
scoped token and uses it upstream. Per the MCP OBO docs, tools/list and
|
||||
tools/call must both work with the user token in Authorization."""
|
||||
result = await _call_add_tool(
|
||||
proxy_server_url=proxy_server_url,
|
||||
server_name="internal_mcp_server",
|
||||
a=6,
|
||||
b=7,
|
||||
headers={
|
||||
"x-litellm-api-key": "Bearer sk-1234",
|
||||
"Authorization": "Bearer user-subject-jwt",
|
||||
},
|
||||
)
|
||||
assert result == "13"
|
||||
257
tests/mcp_tests/test_proxy_mcp_byok_pkce_e2e.py
Normal file
257
tests/mcp_tests/test_proxy_mcp_byok_pkce_e2e.py
Normal file
|
|
@ -0,0 +1,257 @@
|
|||
"""DB-backed end-to-end test for the interactive PKCE (BYOK) MCP OAuth flow.
|
||||
|
||||
This drives LiteLLM acting as the OAuth 2.1 authorization server:
|
||||
1. mint a UI session cookie
|
||||
2. POST /v1/mcp/oauth/authorize with a PKCE code_challenge -> authorization code
|
||||
3. POST /v1/mcp/oauth/token with the code_verifier -> access token (and the
|
||||
per-user credential is persisted to the DB)
|
||||
4. call the MCP server through the proxy using that access token; the proxy
|
||||
loads the stored credential and forwards it upstream
|
||||
|
||||
Requires a database (DATABASE_URL). Skipped otherwise. In CI this runs against
|
||||
the Postgres service wired into the MCP workflow.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import base64
|
||||
import hashlib
|
||||
import os
|
||||
import time
|
||||
import typing
|
||||
import uuid
|
||||
from urllib.parse import parse_qs, urlparse
|
||||
|
||||
import httpx
|
||||
import jwt
|
||||
import pytest
|
||||
import yaml
|
||||
from mcp import ClientSession
|
||||
from mcp.client.streamable_http import streamablehttp_client
|
||||
|
||||
from tests.mcp_tests.mcp_server import DEFAULT_OAUTH_ACCESS_TOKEN
|
||||
from tests.mcp_tests.test_proxy_mcp_auth_e2e import (
|
||||
_reserve_port,
|
||||
_start_mcp_server_process,
|
||||
_start_proxy_server,
|
||||
)
|
||||
|
||||
DATABASE_URL = os.getenv("DATABASE_URL")
|
||||
pytestmark = pytest.mark.skipif(not DATABASE_URL, reason="BYOK PKCE e2e requires a database (DATABASE_URL)")
|
||||
|
||||
MASTER_KEY = "sk-1234"
|
||||
BYOK_USER_ID = "byok-pkce-user"
|
||||
REDIRECT_URI = "http://127.0.0.1:8765/callback"
|
||||
|
||||
|
||||
def _pkce_pair() -> tuple[str, str]:
|
||||
verifier = base64.urlsafe_b64encode(os.urandom(40)).rstrip(b"=").decode()
|
||||
digest = hashlib.sha256(verifier.encode()).digest()
|
||||
challenge = base64.urlsafe_b64encode(digest).rstrip(b"=").decode()
|
||||
return verifier, challenge
|
||||
|
||||
|
||||
def _session_cookie() -> str:
|
||||
return jwt.encode(
|
||||
{
|
||||
"user_id": BYOK_USER_ID,
|
||||
"login_method": "username_password",
|
||||
"exp": int(time.time()) + 3600,
|
||||
},
|
||||
MASTER_KEY,
|
||||
algorithm="HS256",
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def upstream_server() -> typing.Iterator[dict[str, typing.Any]]:
|
||||
port = _reserve_port()
|
||||
process = _start_mcp_server_process(auth_mode="oauth2", port=port, auth_secret=None)
|
||||
try:
|
||||
yield {"base_url": f"http://127.0.0.1:{port}"}
|
||||
finally:
|
||||
process.terminate()
|
||||
try:
|
||||
process.wait(timeout=5)
|
||||
except Exception:
|
||||
process.kill()
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def byok_server_id(upstream_server: dict[str, typing.Any]) -> str:
|
||||
"""Create a BYOK MCP server row directly in the DB and return its id."""
|
||||
from litellm.proxy._types import NewMCPServerRequest
|
||||
from litellm.proxy._experimental.mcp_server.db import create_mcp_server
|
||||
from litellm.proxy.utils import PrismaClient, ProxyLogging
|
||||
from litellm.caching import DualCache
|
||||
|
||||
server_id = str(uuid.uuid4())
|
||||
|
||||
async def _create() -> None:
|
||||
prisma_client = PrismaClient(
|
||||
database_url=DATABASE_URL,
|
||||
proxy_logging_obj=ProxyLogging(user_api_key_cache=DualCache()),
|
||||
)
|
||||
await prisma_client.connect()
|
||||
try:
|
||||
await create_mcp_server(
|
||||
prisma_client,
|
||||
NewMCPServerRequest(
|
||||
server_id=server_id,
|
||||
alias="byok_pkce_server",
|
||||
url=f"{upstream_server['base_url']}/mcp",
|
||||
transport="http",
|
||||
auth_type="oauth2",
|
||||
is_byok=True,
|
||||
allow_all_keys=True,
|
||||
),
|
||||
touched_by="byok-pkce-e2e",
|
||||
)
|
||||
finally:
|
||||
await prisma_client.disconnect()
|
||||
|
||||
asyncio.run(_create())
|
||||
return server_id
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def proxy_server_url(tmp_path_factory: pytest.TempPathFactory, byok_server_id: str) -> typing.Iterator[str]:
|
||||
os.environ["DATABASE_URL"] = DATABASE_URL # restore if a sibling test cleared it
|
||||
os.environ["LITELLM_MASTER_KEY"] = MASTER_KEY
|
||||
|
||||
config = {
|
||||
"general_settings": {"master_key": MASTER_KEY},
|
||||
"model_list": [
|
||||
{
|
||||
"model_name": "fake-model",
|
||||
"litellm_params": {"model": "openai/fake", "api_key": "fake-key"},
|
||||
}
|
||||
],
|
||||
}
|
||||
config_path = tmp_path_factory.mktemp("byok_pkce") / "config.yaml"
|
||||
config_path.write_text(yaml.safe_dump(config))
|
||||
|
||||
server_url, server, thread, sock = _start_proxy_server(str(config_path))
|
||||
yield server_url
|
||||
|
||||
server.should_exit = True
|
||||
thread.join(timeout=10)
|
||||
sock.close()
|
||||
|
||||
|
||||
async def _complete_pkce_flow(proxy_server_url: str, server_id: str) -> str:
|
||||
"""Run authorize -> token and return the issued access token."""
|
||||
verifier, challenge = _pkce_pair()
|
||||
cookies = {"token": _session_cookie()}
|
||||
|
||||
async with httpx.AsyncClient(follow_redirects=False) as client:
|
||||
authorize = await client.post(
|
||||
f"{proxy_server_url}/v1/mcp/oauth/authorize",
|
||||
data={
|
||||
"redirect_uri": REDIRECT_URI,
|
||||
"code_challenge": challenge,
|
||||
"code_challenge_method": "S256",
|
||||
"state": "xyz",
|
||||
"server_id": server_id,
|
||||
"api_key": DEFAULT_OAUTH_ACCESS_TOKEN,
|
||||
"client_id": "byok-client",
|
||||
},
|
||||
cookies=cookies,
|
||||
)
|
||||
assert authorize.status_code == 302, authorize.text
|
||||
code = parse_qs(urlparse(authorize.headers["location"]).query)["code"][0]
|
||||
|
||||
token = await client.post(
|
||||
f"{proxy_server_url}/v1/mcp/oauth/token",
|
||||
data={
|
||||
"grant_type": "authorization_code",
|
||||
"code": code,
|
||||
"code_verifier": verifier,
|
||||
"redirect_uri": REDIRECT_URI,
|
||||
"client_id": "byok-client",
|
||||
},
|
||||
)
|
||||
assert token.status_code == 200, token.text
|
||||
return token.json()["access_token"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_byok_pkce_authorize_rejects_wrong_verifier(proxy_server_url: str, byok_server_id: str) -> None:
|
||||
"""PKCE enforcement: a token request with a mismatched verifier is rejected."""
|
||||
_verifier, challenge = _pkce_pair()
|
||||
cookies = {"token": _session_cookie()}
|
||||
|
||||
async with httpx.AsyncClient(follow_redirects=False) as client:
|
||||
authorize = await client.post(
|
||||
f"{proxy_server_url}/v1/mcp/oauth/authorize",
|
||||
data={
|
||||
"redirect_uri": REDIRECT_URI,
|
||||
"code_challenge": challenge,
|
||||
"code_challenge_method": "S256",
|
||||
"state": "xyz",
|
||||
"server_id": byok_server_id,
|
||||
"api_key": DEFAULT_OAUTH_ACCESS_TOKEN,
|
||||
"client_id": "byok-client",
|
||||
},
|
||||
cookies=cookies,
|
||||
)
|
||||
assert authorize.status_code == 302
|
||||
code = parse_qs(urlparse(authorize.headers["location"]).query)["code"][0]
|
||||
|
||||
token = await client.post(
|
||||
f"{proxy_server_url}/v1/mcp/oauth/token",
|
||||
data={
|
||||
"grant_type": "authorization_code",
|
||||
"code": code,
|
||||
"code_verifier": "wrong-verifier-that-will-not-match-the-challenge",
|
||||
"redirect_uri": REDIRECT_URI,
|
||||
"client_id": "byok-client",
|
||||
},
|
||||
)
|
||||
assert token.status_code == 400
|
||||
assert token.json()["error"] == "invalid_grant"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_byok_pkce_authorize_requires_session(proxy_server_url: str, byok_server_id: str) -> None:
|
||||
"""Without a UI session cookie the authorize endpoint must refuse."""
|
||||
_verifier, challenge = _pkce_pair()
|
||||
async with httpx.AsyncClient(follow_redirects=False) as client:
|
||||
response = await client.post(
|
||||
f"{proxy_server_url}/v1/mcp/oauth/authorize",
|
||||
data={
|
||||
"redirect_uri": REDIRECT_URI,
|
||||
"code_challenge": challenge,
|
||||
"code_challenge_method": "S256",
|
||||
"server_id": byok_server_id,
|
||||
"api_key": DEFAULT_OAUTH_ACCESS_TOKEN,
|
||||
"client_id": "byok-client",
|
||||
},
|
||||
)
|
||||
assert response.status_code == 401
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_byok_pkce_end_to_end_tool_call(proxy_server_url: str, byok_server_id: str) -> None:
|
||||
"""Full interactive PKCE flow: authorize -> token -> call the MCP tool with
|
||||
the issued access token. The proxy forwards the stored BYOK credential
|
||||
upstream, so the tool call succeeds."""
|
||||
access_token = await _complete_pkce_flow(proxy_server_url, byok_server_id)
|
||||
|
||||
headers = {
|
||||
"Authorization": f"Bearer {access_token}",
|
||||
"x-mcp-servers": "byok_pkce_server",
|
||||
}
|
||||
async with asyncio.timeout(20):
|
||||
async with streamablehttp_client(url=f"{proxy_server_url}/mcp", headers=headers) as (
|
||||
read,
|
||||
write,
|
||||
_get_session_id,
|
||||
):
|
||||
async with ClientSession(read, write) as session:
|
||||
await session.initialize()
|
||||
tools = await session.list_tools()
|
||||
assert any(tool.name.endswith("add") for tool in tools.tools)
|
||||
|
||||
result = await session.call_tool("add", arguments={"a": 5, "b": 6})
|
||||
assert result.content
|
||||
assert getattr(result.content[0], "text", None) == "11"
|
||||
Loading…
Add table
Reference in a new issue