test(e2e): cover the interactive authorization_code MCP OAuth flow end to end

This commit is contained in:
Tin Chi Lo 2026-07-14 21:01:52 -07:00
parent b1442d0398
commit 1b42f655b3
8 changed files with 558 additions and 21 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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