add mcp tests

This commit is contained in:
mubashir1osmani 2026-06-07 17:39:58 -07:00
parent 3448bf79f8
commit 3d22e0197f
No known key found for this signature in database
GPG key ID: AB055FF67D0B4D9A
5 changed files with 913 additions and 12 deletions

View file

@ -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

View 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"

View file

@ -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"

View 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"

View 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"