From 3d22e0197f24f5c8dc4135cabccc01d58a34ba8d Mon Sep 17 00:00:00 2001 From: mubashir1osmani Date: Sun, 7 Jun 2026 17:39:58 -0700 Subject: [PATCH] add mcp tests --- tests/mcp_tests/mcp_server.py | 130 ++++++- .../test_configs/test_oauth2_mcp_config.yaml | 74 ++++ tests/mcp_tests/test_mcp_logging.py | 121 ++++++ tests/mcp_tests/test_proxy_mcp_auth_e2e.py | 343 ++++++++++++++++++ .../mcp_tests/test_proxy_mcp_byok_pkce_e2e.py | 257 +++++++++++++ 5 files changed, 913 insertions(+), 12 deletions(-) create mode 100644 tests/mcp_tests/test_configs/test_oauth2_mcp_config.yaml create mode 100644 tests/mcp_tests/test_proxy_mcp_auth_e2e.py create mode 100644 tests/mcp_tests/test_proxy_mcp_byok_pkce_e2e.py diff --git a/tests/mcp_tests/mcp_server.py b/tests/mcp_tests/mcp_server.py index bc6accbb721..bfbb29931c1 100644 --- a/tests/mcp_tests/mcp_server.py +++ b/tests/mcp_tests/mcp_server.py @@ -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 diff --git a/tests/mcp_tests/test_configs/test_oauth2_mcp_config.yaml b/tests/mcp_tests/test_configs/test_oauth2_mcp_config.yaml new file mode 100644 index 00000000000..87134ce01ea --- /dev/null +++ b/tests/mcp_tests/test_configs/test_oauth2_mcp_config.yaml @@ -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" diff --git a/tests/mcp_tests/test_mcp_logging.py b/tests/mcp_tests/test_mcp_logging.py index 55b49aa0d29..d1ce6ea59b9 100644 --- a/tests/mcp_tests/test_mcp_logging.py +++ b/tests/mcp_tests/test_mcp_logging.py @@ -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" diff --git a/tests/mcp_tests/test_proxy_mcp_auth_e2e.py b/tests/mcp_tests/test_proxy_mcp_auth_e2e.py new file mode 100644 index 00000000000..82072de9eac --- /dev/null +++ b/tests/mcp_tests/test_proxy_mcp_auth_e2e.py @@ -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" diff --git a/tests/mcp_tests/test_proxy_mcp_byok_pkce_e2e.py b/tests/mcp_tests/test_proxy_mcp_byok_pkce_e2e.py new file mode 100644 index 00000000000..4b8060af673 --- /dev/null +++ b/tests/mcp_tests/test_proxy_mcp_byok_pkce_e2e.py @@ -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"