diff --git a/tests/e2e/CLAUDE.md b/tests/e2e/CLAUDE.md index e2d13c80e35..aa05208dd8c 100644 --- a/tests/e2e/CLAUDE.md +++ b/tests/e2e/CLAUDE.md @@ -13,7 +13,7 @@ Each subdirectory under `tests/e2e/` is one suite, scoped to an endpoint family - `realtime/` - realtime websocket sessions, including the pipecat audio path - `quota_management/` - quota enforcement and accounting, one subfolder per behavior: `ratelimit/` (rpm/tpm blocks, window reset, pacing headers on live traffic), `budgets/` (budget definition, enforcement, and reset windows: key, team, tag, soft, multi-window), and `spend_tracking/` (spend logging and cost attribution on `/spend/*`) - `management/` - key/team/user/organization management routes: create/update/delete persistence via the info routes, team membership, and llm-only-key route denials; also the dashboard UI behavior on top of them, driven through the proxy-served UI at /ui with playwright (optional dep behind importorskip) -- `mcp/` - the MCP gateway: server registration via `/v1/mcp/server`, tool listing/calling over the streamable-http protocol under both auth headers, and per-server enforcement such as `max_concurrent_requests`; `stub/` holds the deterministic upstream MCP server the compose stack runs for it +- `mcp/` - the MCP gateway: server registration via `/v1/mcp/server`, tool listing/calling over the streamable-http protocol under both auth headers, the interactive (authorization_code) OAuth flow driven with the mcp SDK's own OAuth client, and per-server enforcement such as `max_concurrent_requests`; `stub/` holds the deterministic upstream MCP mounts and stub OAuth2 IdP the compose stack runs for it - `logging/` - logging-integration delivery (datadog and friends) - `security/` - secret handling and log-leak protection - `router/` - routing and reliability behavior (fallbacks, cooldowns) diff --git a/tests/e2e/coverage_registry/mcp.yaml b/tests/e2e/coverage_registry/mcp.yaml index 4a8854660b6..1233abe83df 100644 --- a/tests/e2e/coverage_registry/mcp.yaml +++ b/tests/e2e/coverage_registry/mcp.yaml @@ -135,3 +135,19 @@ assertions: [succeeds] source: "server.py:1089" rationale: Smoke; rarely used; same auth model as tools +- id: mcp.list_tools.oauth.completes_authorization_code_flow + module: mcp + tier: P0 + operation: list_tools + auth_family: oauth + assertions: [completes_authorization_code_flow] + source: "auth/user_api_key_auth_mcp.py preemptive 401 challenge + discoverable_endpoints.py gateway AS" + rationale: The interactive per-user flow real MCP hosts run (401 challenge, discovery, DCR, PKCE authorize dance, token exchange); the path desktop-client re-auth incidents live on +- id: mcp.call_tool.oauth.uses_per_user_token + module: mcp + tier: P0 + operation: call_tool + auth_family: oauth + assertions: [uses_per_user_token] + source: "v2 authorization_code egress arm (per-user DB token)" + rationale: Calls on an interactive server must present the user's upstream token, never the caller's virtual key or IdP JWT diff --git a/tests/e2e/docker-compose.yml b/tests/e2e/docker-compose.yml index 89430c9b758..1f5f196d75a 100644 --- a/tests/e2e/docker-compose.yml +++ b/tests/e2e/docker-compose.yml @@ -170,6 +170,8 @@ services: mcp-stub: build: ./mcp/stub + ports: + - "8765:8765" healthcheck: test: ["CMD", "python", "-c", "import socket; socket.create_connection(('localhost', 8765), timeout=2)"] interval: 3s diff --git a/tests/e2e/e2e_config.py b/tests/e2e/e2e_config.py index 8864efd752f..85f735c9002 100644 --- a/tests/e2e/e2e_config.py +++ b/tests/e2e/e2e_config.py @@ -42,6 +42,26 @@ DD_SINK_URL = os.environ.get("E2E_DD_SINK_URL", "http://localhost:9915").rstrip( # test runs somewhere the compose stub is not visible from. MCP_STUB_URL = os.environ.get("E2E_MCP_STUB_URL", "http://mcp-stub:8765/mcp") +# The stub's sibling surfaces (see mcp/stub/stub_server.py): a Bearer-guarded +# upstream for the interactive OAuth flow, plus the stub IdP's token endpoint. +# Derived from MCP_STUB_URL so one override relocates the whole stub. +_MCP_STUB_BASE = MCP_STUB_URL.removesuffix("/mcp") +MCP_STUB_OAUTHUSER_URL = f"{_MCP_STUB_BASE}/oauthuser/mcp" +MCP_STUB_TOKEN_URL = f"{_MCP_STUB_BASE}/oauth/token" + +# The stub's authorization endpoint is dereferenced by the browser (pytest +# playing that role from the host), never by the proxy, so unlike the URLs +# above it needs the host-visible address of the stub: the compose file +# publishes mcp-stub on host port 8765 for exactly this hop. +MCP_STUB_AUTHORIZE_BROWSER_URL = os.environ.get( + "E2E_MCP_STUB_AUTHORIZE_BROWSER_URL", "http://localhost:8765/oauth/authorize" +) + +# Deterministic test-only credentials; must mirror mcp/stub/stub_server.py. +MCP_STUB_OAUTH_USER_CLIENT_ID = "e2e-stub-user-client-id" +MCP_STUB_OAUTH_USER_CLIENT_SECRET = "e2e-stub-user-client-secret" +MCP_STUB_OAUTH_USER_ACCESS_TOKEN = "e2e-stub-user-access-token" + # Writes on the proxy are eventually consistent (e.g. spend rows flush on # proxy_batch_write_at, ~60s). Read-backs poll to this deadline, never sleep-once. POLL_TIMEOUT = float(os.environ.get("E2E_POLL_TIMEOUT", "120")) diff --git a/tests/e2e/mcp/mcp_client.py b/tests/e2e/mcp/mcp_client.py index a6ac16fb3e9..6efea3016dc 100644 --- a/tests/e2e/mcp/mcp_client.py +++ b/tests/e2e/mcp/mcp_client.py @@ -7,10 +7,10 @@ production MCP hosts use, aimed at the gateway's per-server URL namespace {PROXY}/{alias}/mcp. Every protocol method takes the request headers as a plain dict, built inside -the test body, so the exact wire format is visible where it is asserted. The -gateway accepts the LiteLLM virtual key as either -`x-litellm-api-key: Bearer sk-...` or `Authorization: Bearer sk-...` (both -Bearer-prefixed on the MCP routes, matching the docs). +the test body, so the exact wire format is visible where it is asserted. The gateway accepts the LiteLLM +virtual key as either `x-litellm-api-key: Bearer sk-...` or +`Authorization: Bearer sk-...` (both Bearer-prefixed on the MCP routes, +matching the docs). """ from __future__ import annotations @@ -19,13 +19,16 @@ import asyncio import time from dataclasses import dataclass from typing import Mapping, cast +from urllib.parse import parse_qsl, urljoin import httpx import pytest from mcp import ClientSession +from mcp.client.auth import OAuthClientProvider from mcp.client.streamable_http import streamable_http_client +from mcp.shared.auth import OAuthClientInformationFull, OAuthClientMetadata, OAuthToken from mcp.types import CallToolResult, TextContent -from pydantic import BaseModel +from pydantic import BaseModel, RootModel from e2e_config import PROXY_BASE_URL, REQUEST_TIMEOUT from e2e_gateway import Gateway, build_gateway @@ -67,6 +70,11 @@ class StubToolStats(BaseModel): completed: int +class StubRecordedHeaders(RootModel[dict[str, str]]): + """JSON the stub's `recorded_headers` tool returns: the (lowercased) header + map of the most recent request its auth guard let through.""" + + def _mcp_url(alias: str) -> str: return f"{PROXY_BASE_URL}/{alias}/mcp" @@ -113,6 +121,137 @@ async def _call_tool(url: str, headers: dict[str, str], tool: str, arguments: To return McpToolText(text=_first_text(result), is_error=bool(result.isError)) +# ---------- interactive (authorization_code) OAuth: the MCP-host side ---------- + +# Where the "browser" lands at the end of the authorize dance. Nothing listens +# here on purpose: the redirect chaser stops as soon as the chain points at this +# URL and reads the code/state off the query string, exactly like a desktop MCP +# host that intercepts its loopback redirect. +OAUTH_CLIENT_REDIRECT_URI = "http://127.0.0.1:53682/e2e/callback" + + +class InMemoryTokenStorage: + """The mcp SDK's TokenStorage protocol, held in memory for one test: the + DCR-registered client and the tokens the gateway minted for it.""" + + def __init__(self) -> None: + self._tokens: OAuthToken | None = None + self._client_info: OAuthClientInformationFull | None = None + + async def get_tokens(self) -> OAuthToken | None: + return self._tokens + + async def set_tokens(self, tokens: OAuthToken) -> None: + self._tokens = tokens + + async def get_client_info(self) -> OAuthClientInformationFull | None: + return self._client_info + + async def set_client_info(self, client_info: OAuthClientInformationFull) -> None: + self._client_info = client_info + + +async def _follow_authorize_redirects(start_url: str) -> tuple[str, str | None]: + """Play the browser's role in the authorize dance: chase the redirect chain + (gateway authorize -> upstream IdP authorize -> gateway callback -> MCP + host redirect_uri) with plain GETs, and return the code/state from the + final hop's query string without dereferencing it.""" + async with httpx.AsyncClient(timeout=httpx.Timeout(REQUEST_TIMEOUT), follow_redirects=False) as browser: + url = start_url + for _ in range(10): + if url.startswith(OAUTH_CLIENT_REDIRECT_URI): + params = dict(parse_qsl(httpx.URL(url).query.decode())) + assert "code" in params, f"authorize chain reached the client redirect_uri without a code: {url}" + return params["code"], params.get("state") + response = await browser.get(url) + assert response.status_code in (302, 303, 307), ( + f"authorize chain broke at {url}: {response.status_code} {response.text[:300]}" + ) + url = urljoin(url, response.headers["location"]) + raise AssertionError(f"authorize chain never reached {OAUTH_CLIENT_REDIRECT_URI}; last url: {url}") + + +def _oauth_provider(url: str, storage: InMemoryTokenStorage) -> OAuthClientProvider: + """The SDK's real OAuth machinery (RFC 9728/8414 discovery, RFC 7591 DCR, + PKCE, token exchange) with the browser leg replaced by the redirect chaser.""" + code_holder: dict[str, str | None] = {} # mutable-ok: hand-off between the two SDK callbacks + + async def redirect_handler(authorize_url: str) -> None: + code, state = await _follow_authorize_redirects(authorize_url) + code_holder["code"] = code + code_holder["state"] = state + + async def callback_handler() -> tuple[str, str | None]: + code = code_holder.get("code") + assert code is not None, "callback_handler ran before the authorize redirect completed" + return code, code_holder.get("state") + + return OAuthClientProvider( + server_url=url, + client_metadata=OAuthClientMetadata.model_validate( + { + "redirect_uris": [OAUTH_CLIENT_REDIRECT_URI], + "token_endpoint_auth_method": "none", + "grant_types": ["authorization_code", "refresh_token"], + "response_types": ["code"], + "client_name": "e2e-mcp-host", + } + ), + storage=storage, + redirect_handler=redirect_handler, + callback_handler=callback_handler, + ) + + +class _HeaderInjectingTransport(httpx.AsyncBaseTransport): + """Adds the caller's LiteLLM key header to every outgoing request, + including the ones the SDK's OAuth machinery builds internally (discovery, + DCR, token exchange). Those bypass httpx client-default headers, but the + gateway resolves which user to store the upstream token for from the key + on the token exchange, exactly like a production MCP host configured with + a LiteLLM key header on the server URL.""" + + def __init__(self, inner: httpx.AsyncBaseTransport, headers: dict[str, str]) -> None: + self._inner = inner + self._headers = headers + + async def handle_async_request(self, request: httpx.Request) -> httpx.Response: + for name, value in self._headers.items(): + if name not in request.headers: + request.headers[name] = value + return await self._inner.handle_async_request(request) + + +def _oauth_http_client(headers: dict[str, str], auth: OAuthClientProvider) -> httpx.AsyncClient: + return httpx.AsyncClient( + headers=headers, + auth=auth, + timeout=httpx.Timeout(REQUEST_TIMEOUT), + follow_redirects=True, + transport=_HeaderInjectingTransport(httpx.AsyncHTTPTransport(), headers), + ) + + +async def _oauth_list_tool_names(url: str, headers: dict[str, str], storage: InMemoryTokenStorage) -> tuple[str, ...]: + async with _oauth_http_client(headers, _oauth_provider(url, storage)) as http_client: + async with streamable_http_client(url, http_client=http_client) as (read, write, _): + async with ClientSession(read, write) as session: + await session.initialize() + listed = await session.list_tools() + return tuple(sorted(tool.name for tool in listed.tools)) + + +async def _oauth_call_tool( + url: str, headers: dict[str, str], storage: InMemoryTokenStorage, tool: str, arguments: ToolArguments +) -> McpToolText: + async with _oauth_http_client(headers, _oauth_provider(url, storage)) as http_client: + async with streamable_http_client(url, http_client=http_client) as (read, write, _): + async with ClientSession(read, write) as session: + await session.initialize() + result = await session.call_tool(tool, dict(arguments)) + return McpToolText(text=_first_text(result), is_error=bool(result.isError)) + + @dataclass(frozen=True, slots=True) class McpClient: gateway: Gateway @@ -173,10 +312,45 @@ class McpClient: behave like independent clients.""" return asyncio.run(_call_tool(_mcp_url(alias), headers, tool, arguments)) + def poll_oauth_tool_names( + self, alias: str, headers: dict[str, str], storage: InMemoryTokenStorage + ) -> tuple[str, ...]: + """The full interactive flow (discovery, DCR, authorize dance, token + exchange, then tools/list) retried to the shared deadline, since the + just-created server record and key propagate asynchronously. Tokens + land in `storage`, so later calls skip the dance like a real host.""" + deadline = time.monotonic() + self.gateway.poll_timeout + last_error: Exception | None = None + while time.monotonic() < deadline: + try: + return asyncio.run(_oauth_list_tool_names(_mcp_url(alias), headers, storage)) + except Exception as exc: # noqa: BLE001 - retried to the deadline; the last error is surfaced below + last_error = exc + time.sleep(self.gateway.poll_interval) + pytest.fail( + f"interactive OAuth flow for {alias!r} never completed within " + f"{self.gateway.poll_timeout}s; last error: {last_error!r}" + ) + + def oauth_call_tool( + self, alias: str, headers: dict[str, str], storage: InMemoryTokenStorage, tool: str, arguments: ToolArguments + ) -> McpToolText: + """One tools/call authenticated by the tokens in `storage` (minted by a + prior poll_oauth_tool_names dance), over its own fresh MCP session.""" + return asyncio.run(_oauth_call_tool(_mcp_url(alias), headers, storage, tool, arguments)) + def stub_stats(self, alias: str, headers: dict[str, str], stats_tool: str, marker: str) -> StubToolStats: outcome = self.call_tool(alias, headers, stats_tool, {"marker": marker}) return StubToolStats.model_validate_json(outcome.text) + def stub_recorded_headers(self, alias: str, headers: dict[str, str], headers_tool: str) -> dict[str, str]: + """The headers the stub's auth guard recorded for its most recent + authorized request: by construction, the tools/call this method just + made, i.e. exactly what the gateway sent upstream on the caller's behalf.""" + outcome = self.call_tool(alias, headers, headers_tool, {}) + assert outcome.is_error is False, f"recorded_headers call errored: {outcome.text[:300]}" + return StubRecordedHeaders.model_validate_json(outcome.text).root + def build_client() -> McpClient: return McpClient(gateway=build_gateway()) diff --git a/tests/e2e/mcp/stub/stub_server.py b/tests/e2e/mcp/stub/stub_server.py index 89d44a36566..36d1079a4f4 100644 --- a/tests/e2e/mcp/stub/stub_server.py +++ b/tests/e2e/mcp/stub/stub_server.py @@ -1,27 +1,67 @@ -"""Deterministic MCP upstream for the mcp e2e suite. +"""Deterministic MCP upstreams for the mcp e2e suite. -Serves the streamable-http MCP protocol with three tools. `echo` answers -immediately so auth tests can assert an exact round-trip. `slow_echo` holds the -request open for `sleep_seconds` while per-`marker` in-flight and max-in-flight -counters track how many calls the proxy let through simultaneously; that is the -observable a per-server `max_concurrent_requests` cap must bound. `stats` reads -those counters back, so tests observe upstream concurrency through the proxy -itself and the stub needs no side-channel port. +One process, one port, two streamable-http MCP mounts plus a deterministic +OAuth2 IdP, so the compose stack keeps a single `mcp-stub` service: -Counter updates are plain attribute mutations between awaits, so asyncio's -single-threaded scheduling makes them atomic; markers come from +- `/mcp` — anonymous. `echo` answers immediately so auth tests can assert an + exact round-trip; `slow_echo` holds the request open for `sleep_seconds` + while per-`marker` in-flight and max-in-flight counters track how many calls + the proxy let through simultaneously (the observable a per-server + `max_concurrent_requests` cap must bound); `stats` reads those counters back, + so tests observe upstream concurrency through the proxy itself and the stub + needs no side-channel port. +- `/oauthuser/mcp` — the interactive (authorization_code) upstream: rejects + anything but `Bearer OAUTH_USER_ACCESS_TOKEN`, which only the + authorization_code grant hands out, so a served request proves the whole + browser dance (authorize redirect, code, PKCE-verified token exchange) ran. +- `/oauth/authorize` — the auto-approving authorization endpoint: validates + client_id/response_type, records the one-time code with its PKCE challenge, + and 302s straight back to the caller's redirect_uri with code and state (the + "user" of this IdP always consents instantly). +- `/oauth/token` — the token endpoint: authorization_code validates the code, + redirect_uri, client credentials, and (when a challenge was recorded) the + S256 code_verifier, then answers with OAUTH_USER_ACCESS_TOKEN + + OAUTH_USER_REFRESH_TOKEN; refresh_token re-issues OAUTH_USER_ACCESS_TOKEN + for the known refresh token. + +The guarded mount records the headers of the most recent authorized request; +its `recorded_headers` tool reads them back through the proxy, so tests can +assert exactly which credentials the gateway attached upstream (and that the +caller's LiteLLM virtual key never left the gateway). + +The credential constants are mirrored in tests/e2e/e2e_config.py; keep the two +in sync. Counter updates are plain attribute mutations between awaits, so +asyncio's single-threaded scheduling makes them atomic; markers come from `unique_marker()` so concurrent test runs never share a counter. """ from __future__ import annotations import asyncio +import base64 +import contextlib +import hashlib import json +import uuid +from collections.abc import AsyncGenerator from dataclasses import dataclass +import uvicorn from mcp.server.fastmcp import FastMCP +from starlette.applications import Starlette +from starlette.datastructures import FormData, Headers +from starlette.requests import Request +from starlette.responses import JSONResponse, RedirectResponse, Response +from starlette.routing import Mount, Route +from starlette.types import ASGIApp, Receive, Scope, Send -mcp = FastMCP("e2e-stub", host="0.0.0.0", port=8765, stateless_http=True) +OAUTH_USER_CLIENT_ID = "e2e-stub-user-client-id" +OAUTH_USER_CLIENT_SECRET = "e2e-stub-user-client-secret" +OAUTH_USER_ACCESS_TOKEN = "e2e-stub-user-access-token" +OAUTH_USER_REFRESH_TOKEN = "e2e-stub-user-refresh-token" + +main_mcp = FastMCP("e2e-stub", host="0.0.0.0", port=8765, stateless_http=True) +oauthuser_mcp = FastMCP("e2e-stub-oauthuser", host="0.0.0.0", port=8765, stateless_http=True) @dataclass @@ -33,14 +73,16 @@ class _MarkerStats: _stats: dict[str, _MarkerStats] = {} +_last_authorized_headers: dict[str, dict[str, str]] = {} -@mcp.tool() + +@main_mcp.tool() def echo(text: str) -> str: """Return `text` unchanged.""" return text -@mcp.tool() +@main_mcp.tool() async def slow_echo(text: str, marker: str, sleep_seconds: float) -> str: """Return `text` after `sleep_seconds`, recording concurrency under `marker`.""" stats = _stats.setdefault(marker, _MarkerStats()) @@ -54,7 +96,7 @@ async def slow_echo(text: str, marker: str, sleep_seconds: float) -> str: return text -@mcp.tool() +@main_mcp.tool() def stats(marker: str) -> str: """Return the JSON stats recorded for `marker`.""" recorded = _stats.get(marker, _MarkerStats()) @@ -67,5 +109,149 @@ def stats(marker: str) -> str: ) +def _register_guarded_tools(server: FastMCP, mount: str) -> None: + def echo(text: str) -> str: + """Return `text` unchanged.""" + return text + + def recorded_headers() -> str: + """Return the headers of the most recent authorized request as JSON.""" + return json.dumps(_last_authorized_headers.get(mount, {})) + + _ = server.tool()(echo) + _ = server.tool()(recorded_headers) + + +_register_guarded_tools(oauthuser_mcp, "oauthuser") + + +def _require_header(app: ASGIApp, *, mount: str, header: str, expected: str) -> ASGIApp: + """Serve `app` only to requests carrying `header: expected`; 401 otherwise. + Authorized requests have their full header map recorded under `mount`.""" + + async def guard(scope: Scope, receive: Receive, send: Send) -> None: + if scope["type"] != "http": + await app(scope, receive, send) + return + headers = Headers(scope=scope) + if headers.get(header) != expected: + response = JSONResponse({"error": "unauthorized", "detail": f"missing or wrong {header}"}, status_code=401) + await response(scope, receive, send) + return + _last_authorized_headers[mount] = dict(headers.items()) + await app(scope, receive, send) + + return guard + + +_issued_codes: dict[str, dict[str, str]] = {} + + +async def oauth_authorize(request: Request) -> Response: + """The auto-approving authorization endpoint: the resource owner of this + IdP consents instantly, so a headless test can drive the browser leg of + the authorization_code flow with plain HTTP redirects.""" + params = request.query_params + redirect_uri = params.get("redirect_uri", "") + valid = params.get("response_type") == "code" and params.get("client_id") == OAUTH_USER_CLIENT_ID and redirect_uri + if not valid: + return JSONResponse( + {"error": "invalid_request", "detail": "need response_type=code, the user client_id, and a redirect_uri"}, + status_code=400, + ) + code = f"e2e-stub-code-{uuid.uuid4().hex}" + _issued_codes[code] = { + "redirect_uri": redirect_uri, + "code_challenge": params.get("code_challenge", ""), + } + state = params.get("state", "") + separator = "&" if "?" in redirect_uri else "?" + location = f"{redirect_uri}{separator}code={code}" + (f"&state={state}" if state else "") + return RedirectResponse(location, status_code=302) + + +def _s256(code_verifier: str) -> str: + return base64.urlsafe_b64encode(hashlib.sha256(code_verifier.encode()).digest()).rstrip(b"=").decode() + + +def _authorization_code_grant(form: FormData) -> JSONResponse: + code_value = form.get("code") + record = _issued_codes.pop(code_value, None) if isinstance(code_value, str) else None + granted = ( + record is not None + and form.get("client_id") == OAUTH_USER_CLIENT_ID + and form.get("client_secret") == OAUTH_USER_CLIENT_SECRET + and form.get("redirect_uri") == record["redirect_uri"] + ) + if not granted: + return JSONResponse({"error": "invalid_grant"}, status_code=400) + assert record is not None + if record["code_challenge"]: + verifier = form.get("code_verifier") + if not isinstance(verifier, str) or _s256(verifier) != record["code_challenge"]: + return JSONResponse({"error": "invalid_grant", "detail": "PKCE verification failed"}, status_code=400) + return JSONResponse( + { + "access_token": OAUTH_USER_ACCESS_TOKEN, + "token_type": "Bearer", + "expires_in": 3600, + "refresh_token": OAUTH_USER_REFRESH_TOKEN, + } + ) + + +def _refresh_token_grant(form: FormData) -> JSONResponse: + granted = ( + form.get("refresh_token") == OAUTH_USER_REFRESH_TOKEN + and form.get("client_id") == OAUTH_USER_CLIENT_ID + and form.get("client_secret") == OAUTH_USER_CLIENT_SECRET + ) + if not granted: + return JSONResponse({"error": "invalid_grant"}, status_code=400) + return JSONResponse({"access_token": OAUTH_USER_ACCESS_TOKEN, "token_type": "Bearer", "expires_in": 3600}) + + +async def oauth_token(request: Request) -> JSONResponse: + """The token endpoint for all three grants. Every form field is matched + exactly so a failure points at the precise field the proxy sent wrong.""" + form = await request.form() + grant_type = form.get("grant_type") + if grant_type == "authorization_code": + return _authorization_code_grant(form) + if grant_type == "refresh_token": + return _refresh_token_grant(form) + return JSONResponse({"error": "unsupported_grant_type"}, status_code=400) + + +def build_app() -> Starlette: + servers = (main_mcp, oauthuser_mcp) + apps = {server.name: server.streamable_http_app() for server in servers} + + @contextlib.asynccontextmanager + async def lifespan(_: Starlette) -> AsyncGenerator[None]: + async with contextlib.AsyncExitStack() as stack: + for server in servers: + await stack.enter_async_context(server.session_manager.run()) + yield + + return Starlette( + routes=[ + Route("/oauth/token", oauth_token, methods=["POST"]), + Route("/oauth/authorize", oauth_authorize, methods=["GET"]), + Mount( + "/oauthuser", + app=_require_header( + apps["e2e-stub-oauthuser"], + mount="oauthuser", + header="authorization", + expected=f"Bearer {OAUTH_USER_ACCESS_TOKEN}", + ), + ), + Mount("/", app=apps["e2e-stub"]), + ], + lifespan=lifespan, + ) + + if __name__ == "__main__": - mcp.run(transport="streamable-http") + uvicorn.run(build_app(), host="0.0.0.0", port=8765) diff --git a/tests/e2e/mcp/test_mcp_oauth_interactive_e2e.py b/tests/e2e/mcp/test_mcp_oauth_interactive_e2e.py new file mode 100644 index 00000000000..92ce7a8b037 --- /dev/null +++ b/tests/e2e/mcp/test_mcp_oauth_interactive_e2e.py @@ -0,0 +1,117 @@ +"""Live e2e: the interactive (authorization_code) OAuth MCP flow. + +Covers mcp.list_tools.oauth.completes_authorization_code_flow and +mcp.call_tool.oauth.uses_per_user_token: a gateway-managed oauth2 server in +the authorization_code flow must challenge an unauthorized MCP session with +401 + WWW-Authenticate, let a real MCP host complete the whole interactive +dance (RFC 9728/8414 discovery against the gateway, RFC 7591 dynamic client +registration, the authorize redirect through the upstream IdP, the +PKCE-verified token exchange), and then serve tools/list and tools/call using +the per-user upstream token the gateway obtained; the caller's virtual key +must never leave the gateway. + +The MCP-host side is the official mcp SDK's own OAuth machinery +(OAuthClientProvider), the same code path desktop MCP hosts run; only the +browser leg is replaced by a redirect chaser that follows the authorize chain +(gateway -> stub IdP -> gateway callback -> host redirect_uri) with plain +GETs. The upstream is the /oauthuser stub mount, which 401s anything but the +exact access token the stub IdP's authorization_code grant hands out, so a +served call proves the gateway completed the code exchange (the stub verifies +client credentials, redirect_uri, one-time code, and the S256 code_verifier) +rather than forwarding anything it already had. The recorded_headers +read-back then makes the injected token explicit and adds the leak check. +""" + +from __future__ import annotations + +import pytest + +from e2e_config import ( + MCP_STUB_AUTHORIZE_BROWSER_URL, + MCP_STUB_OAUTH_USER_ACCESS_TOKEN, + MCP_STUB_OAUTH_USER_CLIENT_ID, + MCP_STUB_OAUTH_USER_CLIENT_SECRET, + MCP_STUB_OAUTHUSER_URL, + MCP_STUB_TOKEN_URL, + MCP_STUB_URL, + unique_marker, +) +from lifecycle import ResourceManager +from mcp_client import InMemoryTokenStorage, McpClient, McpDenied, StubRecordedHeaders +from models import KeyGenerateBody, McpServerCreateBody, McpServerCredentials + +pytestmark = pytest.mark.e2e + +GUARDED_STUB_TOOLS = ("echo", "recorded_headers") + + +class TestMcpOauthAuthorizationCode: + """An authorization_code server challenges unauthenticated sessions, runs + the full interactive dance with a real MCP host, and serves tools with the + per-user token the gateway obtained from the IdP.""" + + @pytest.mark.covers("mcp.list_tools.oauth.completes_authorization_code_flow") + @pytest.mark.covers("mcp.call_tool.oauth.uses_per_user_token") + def test_challenge_then_authorization_code_dance_reaches_tools( + self, client: McpClient, resources: ResourceManager + ) -> None: + marker = unique_marker() + alias = f"e2emcpauthcode{marker}" + created = client.create_server( + McpServerCreateBody( + alias=alias, + url=MCP_STUB_OAUTHUSER_URL, + allow_all_keys=True, + auth_type="oauth2", + oauth2_flow="authorization_code", + authorization_url=MCP_STUB_AUTHORIZE_BROWSER_URL, + token_url=MCP_STUB_TOKEN_URL, + credentials=McpServerCredentials( + client_id=MCP_STUB_OAUTH_USER_CLIENT_ID, + client_secret=MCP_STUB_OAUTH_USER_CLIENT_SECRET, + ), + ) + ) + resources.defer(lambda: client.delete_server(created.server_id)) + + stored = client.server_info(created.server_id) + assert stored.auth_type == "oauth2" + assert stored.oauth2_flow == "authorization_code" + assert stored.authorization_url == MCP_STUB_AUTHORIZE_BROWSER_URL + assert stored.token_url == MCP_STUB_TOKEN_URL + assert stored.credentials is None, f"client secret must be redacted on read-back, got {stored.credentials}" + + key = client.gateway.generate_key(KeyGenerateBody(user_id="e2e-test-user")) + resources.defer(lambda: client.gateway.delete_key(key)) + headers = {"x-litellm-api-key": f"Bearer {key}"} + + control_alias = f"e2emcpauthcodectl{marker}" + control = client.create_server( + McpServerCreateBody(alias=control_alias, url=MCP_STUB_URL, allow_all_keys=True) + ) + resources.defer(lambda: client.delete_server(control.server_id)) + _ = client.poll_tool_names(control_alias, headers) + + challenged = client.list_tools_once(alias, headers) + assert isinstance(challenged, McpDenied), f"session without a user token was served tools: {challenged}" + assert challenged.status_code == 401, f"expected the 401 OAuth challenge, got {challenged}" + + storage = InMemoryTokenStorage() + names = client.poll_oauth_tool_names(alias, headers, storage) + expected = tuple(sorted(f"{alias}-{tool}" for tool in GUARDED_STUB_TOOLS)) + assert names == expected, f"post-dance listing was {names}, expected exactly {expected}" + + payload = f"e2e-{marker}" + result = client.oauth_call_tool(alias, headers, storage, f"{alias}-echo", {"text": payload}) + assert result.is_error is False, f"echo through the user-token upstream errored: {result.text[:300]}" + assert result.text == payload + + recorded = client.oauth_call_tool(alias, headers, storage, f"{alias}-recorded_headers", {}) + assert recorded.is_error is False, f"recorded_headers call errored: {recorded.text[:300]}" + upstream_headers = StubRecordedHeaders.model_validate_json(recorded.text).root + assert upstream_headers.get("authorization") == f"Bearer {MCP_STUB_OAUTH_USER_ACCESS_TOKEN}", ( + "upstream must receive exactly the per-user token the stub IdP hands out for the code exchange, " + f"got {upstream_headers.get('authorization')!r}" + ) + leaked = sorted(name for name, value in upstream_headers.items() if key in value) + assert leaked == [], f"caller's virtual key crossed the gateway boundary in header(s) {leaked}" diff --git a/tests/e2e/models.py b/tests/e2e/models.py index 4a0981c4d6d..f54674515f1 100644 --- a/tests/e2e/models.py +++ b/tests/e2e/models.py @@ -98,6 +98,18 @@ class KeyInfoResponse(BaseModel): # ---------- mcp servers ---------- +class McpServerCredentials(BaseModel): + """The `credentials` blob on /v1/mcp/server: `auth_value` is the static + secret for the shared-key auth types (api_key, bearer_token, ...); + `client_id`/`client_secret`/`scopes` drive the OAuth2 client_credentials + exchange. Stored encrypted and redacted (nulled) in every read-back.""" + + auth_value: str | None = None + client_id: str | None = None + client_secret: str | None = None + scopes: list[str] | None = None + + class McpServerCreateBody(BaseModel): """POST /v1/mcp/server. `allow_all_keys` opts the server out of per-key object_permission grants so any virtual key on the proxy may use it.""" @@ -107,6 +119,11 @@ class McpServerCreateBody(BaseModel): transport: str = "http" allow_all_keys: bool = True max_concurrent_requests: int | None = None + auth_type: str | None = None + credentials: McpServerCredentials | None = None + authorization_url: str | None = None + token_url: str | None = None + oauth2_flow: Literal["client_credentials", "authorization_code"] | None = None class McpServerInfo(BaseModel): @@ -118,6 +135,11 @@ class McpServerInfo(BaseModel): transport: str | None = None allow_all_keys: bool | None = None max_concurrent_requests: int | None = None + auth_type: str | None = None + credentials: McpServerCredentials | None = None + authorization_url: str | None = None + token_url: str | None = None + oauth2_flow: str | None = None # ---------- customers ----------