mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
test(e2e): cover the interactive authorization_code MCP OAuth flow end to end
This commit is contained in:
parent
b1442d0398
commit
1b42f655b3
8 changed files with 558 additions and 21 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"))
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
117
tests/e2e/mcp/test_mcp_oauth_interactive_e2e.py
Normal file
117
tests/e2e/mcp/test_mcp_oauth_interactive_e2e.py
Normal 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}"
|
||||
|
|
@ -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 ----------
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue