"""Client for the mcp chat-completion OAuth e2e suite. Registers a gateway-managed OAuth (authorization_code) MCP server, seeds the per-user upstream token by driving the interactive authorize dance with the official mcp SDK's OAuthClientProvider (the browser leg is a headless Chromium primed with a human's saved Linear session), then exercises the server through /chat/completions, where the gateway lists and executes its tools with the stored per-user token. Management routes (/v1/mcp/server CRUD, /chat/completions) go through the shared ProxyClient transport. The MCP protocol used to seed the token goes through the mcp SDK, the same library production MCP hosts run. """ from __future__ import annotations import asyncio import re import time from dataclasses import dataclass from typing import TYPE_CHECKING, Final from urllib.parse import parse_qsl import httpx import httpx2 import pytest from e2e_config import PROXY_BASE_URL, REQUEST_TIMEOUT from e2e_http import AuthHeaders, NoBody, unwrap from idp import Identity from mcp import ClientSession from mcp.client.auth import OAuthClientProvider from mcp.client.streamable_http import streamable_http_client from mcp.shared.auth import AuthorizationCodeResult, OAuthClientInformationFull, OAuthClientMetadata, OAuthToken from mcp.types import TextContent from models import ( ChatBody, ChatResponse, McpOauthUserCredentialStatus, McpServerCreateBody, McpServerInfo, McpServerUserCredentialListResponse, McpServerUserCredentialRow, ) from proxy_client import ProxyClient if TYPE_CHECKING: from playwright.async_api import Route # Where the "browser" lands at the end of the authorize dance. Nothing listens # here: the route interceptor short-circuits the final redirect and reads the # code/state off its query string, exactly like a desktop MCP host intercepting # its loopback redirect. OAUTH_CLIENT_REDIRECT_URI = "http://127.0.0.1:53682/e2e/callback" BROWSER_CONSENT_TIMEOUT = 60.0 def _mcp_url(alias: str, base_url: str = PROXY_BASE_URL) -> str: return f"{base_url.rstrip('/')}/{alias}/mcp" class InMemoryTokenStorage: """The mcp SDK's TokenStorage protocol, in memory for one dance: the DCR-registered client and the gateway tokens 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 _browser_follow_authorize( start_url: str, storage_state_path: str, identity: Identity | None = None, server_alias: str | None = None, allow_upstream_consent: bool = True, ) -> tuple[str, str | None]: """Play the browser's role for a real upstream whose authorize endpoint serves an interactive consent page (Linear). A headless Chromium primed with a human's saved Linear session opens the gateway authorize URL and clicks through Linear's consent screens (the mcp.linear.app Approve form, then the linear.app workspace-selection page), riding the rest of the chain (Linear -> gateway callback -> host redirect_uri). The final hop is intercepted and short-circuited, since nothing listens there, and its code/state are read off the query string.""" from playwright.async_api import async_playwright captured: dict[str, str] = {} # mutable-ok: hand-off from the request listener trail: list[str] = [] # mutable-ok: navigation diagnostics for a failed dance def _note_request(request: object) -> None: url = getattr(request, "url", "") host = httpx.URL(url).host if not allow_upstream_consent and (host == "linear.app" or host.endswith(".linear.app")): captured["upstream_consent"] = "seen" if url.startswith(OAUTH_CLIENT_REDIRECT_URI) and "url" not in captured: captured["url"] = url async def _swallow_redirect(route: Route) -> None: await route.fulfill(status=200, content_type="text/plain", body="ok") async with async_playwright() as playwright: browser = await playwright.chromium.launch(headless=True) context = await browser.new_context(storage_state=storage_state_path) await context.route(re.compile(re.escape(OAUTH_CLIENT_REDIRECT_URI) + r".*"), _swallow_redirect) page = await context.new_page() context.on("request", _note_request) page.on("framenavigated", lambda frame: trail.append(frame.url.split("?", 1)[0])) await page.goto(start_url, wait_until="domcontentloaded") deadline = time.monotonic() + BROWSER_CONSENT_TIMEOUT while "url" not in captured and time.monotonic() < deadline: try: await page.wait_for_load_state("networkidle", timeout=8000) except Exception: # noqa: BLE001 - a busy consent page never idles; fall through and try to advance it pass if "upstream_consent" in captured or "url" in captured: break if await page.locator("#username").count() and identity is not None: await page.locator("#username").fill(identity.username) await page.locator("#password").fill(identity.password) await page.locator("#kc-login").click() continue if "/ui/connect" in page.url and server_alias is not None: card = page.locator("div.cursor-pointer").filter(has=page.get_by_text(server_alias, exact=True)) if await card.count() != 1: await asyncio.sleep(0.5) continue connect = card.get_by_text("Connect", exact=True) if await connect.count(): await connect.click() continue if not await card.locator("svg.text-success").count(): await asyncio.sleep(0.5) continue finish = page.get_by_role("button", name="Finish connecting", exact=True) if await finish.count() and await finish.is_enabled(): await finish.click() continue control = page.locator( 'button[name="action"][value="approve"], button:has-text("Authorize"), ' 'button:has-text("Allow"), button:has-text("@"), a:has-text("@")' ).first try: await control.click(timeout=5000) except Exception: # noqa: BLE001 - nothing to advance yet; loop and re-check await asyncio.sleep(0.5) final_url = page.url await browser.close() # A redirect chain can finish inside goto/networkidle before the loop checks the page. assert "upstream_consent" not in captured, "cold reconnect required upstream consent" landing = captured.get("url") assert landing is not None, ( f"consent flow never reached {OAUTH_CLIENT_REDIRECT_URI}; " f"final={final_url.split('?', 1)[0]!r}; trail={trail[-6:]}" ) params = dict(parse_qsl(httpx.URL(landing).query.decode())) assert "code" in params, "client redirect_uri carried no authorization code" return params["code"], params.get("state") def _oauth_provider( url: str, storage: InMemoryTokenStorage, storage_state_path: str | None, identity: Identity | None = None, server_alias: str | None = None, allow_upstream_consent: bool = True, ) -> OAuthClientProvider: """The SDK's real OAuth machinery (RFC 9728/8414 discovery, RFC 7591 DCR, PKCE, token exchange) with the browser leg driven by Playwright against the upstream's consent screen.""" code_holder: dict[str, str | None] = {} # mutable-ok: hand-off between the two SDK callbacks async def _reject_redirect(_: str) -> None: raise AssertionError("gateway demanded a fresh upstream consent; stored per-user token was not reused") async def _follow_redirect(authorize_url: str) -> None: assert storage_state_path is not None code, state = await _browser_follow_authorize( authorize_url, storage_state_path, identity, server_alias, allow_upstream_consent ) code_holder["code"] = code code_holder["state"] = state redirect_handler: Final = _reject_redirect if storage_state_path is None else _follow_redirect async def callback_handler() -> AuthorizationCodeResult: code = code_holder.get("code") assert code is not None, "callback_handler ran before the authorize redirect completed" return AuthorizationCodeResult(code=code, state=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(httpx2.AsyncBaseTransport): """Adds the caller's LiteLLM key header to every outgoing SDK request (discovery, DCR, token exchange), so 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.""" def __init__(self, inner: httpx2.AsyncBaseTransport, headers: dict[str, str], gateway_url: str) -> None: self._inner = inner self._headers = headers self._gateway_url = httpx2.URL(gateway_url) @staticmethod def _port(url: httpx2.URL) -> int | None: if url.port is not None: return url.port return {"http": 80, "https": 443}.get(url.scheme) async def handle_async_request(self, request: httpx2.Request) -> httpx2.Response: same_origin: Final = ( request.url.scheme == self._gateway_url.scheme and request.url.host == self._gateway_url.host and self._port(request.url) == self._port(self._gateway_url) ) if same_origin: for name, value in self._headers.items(): if name not in request.headers: request.headers[name] = value else: for name, value in self._headers.items(): if request.headers.get(name) == value: del request.headers[name] return await self._inner.handle_async_request(request) async def aclose(self) -> None: await self._inner.aclose() def _oauth_http_client( headers: dict[str, str], auth: OAuthClientProvider, gateway_url: str = PROXY_BASE_URL ) -> httpx2.AsyncClient: return httpx2.AsyncClient( auth=auth, timeout=httpx2.Timeout(REQUEST_TIMEOUT), follow_redirects=True, transport=_HeaderInjectingTransport(httpx2.AsyncHTTPTransport(), headers, gateway_url), ) async def _seed_via_dance( url: str, headers: dict[str, str], storage: InMemoryTokenStorage, storage_state_path: str ) -> tuple[str, ...]: async with _oauth_http_client(headers, _oauth_provider(url, storage, storage_state_path)) 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)) @dataclass(frozen=True, slots=True) class OauthToolRun: tools: tuple[str, ...] is_error: bool text: str async def _list_and_call( url: str, headers: dict[str, str], storage: InMemoryTokenStorage, storage_state_path: str | None, tool: str, arguments: dict[str, str], gateway_url: str = PROXY_BASE_URL, identity: Identity | None = None, server_alias: str | None = None, allow_upstream_consent: bool = True, ) -> OauthToolRun: async with _oauth_http_client( headers, _oauth_provider(url, storage, storage_state_path, identity, server_alias, allow_upstream_consent), gateway_url, ) 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: Final = await session.list_tools() result: Final = await session.call_tool(tool, arguments) text: Final = "".join(content.text for content in result.content if isinstance(content, TextContent)) return OauthToolRun( tools=tuple(sorted(tool_item.name for tool_item in listed.tools)), is_error=result.is_error, text=text, ) @dataclass(frozen=True, slots=True) class ChatMcpClient: proxy: ProxyClient def create_server(self, body: McpServerCreateBody) -> McpServerInfo: return unwrap( self.proxy.transport.post( "/v1/mcp/server", headers=self.proxy.transport.master, json=body, response_type=McpServerInfo, ) ) def server_info(self, server_id: str) -> McpServerInfo: return unwrap( self.proxy.transport.get( f"/v1/mcp/server/{server_id}", headers=self.proxy.transport.master, params=NoBody(), response_type=McpServerInfo, ) ) def delete_server(self, server_id: str) -> None: _ = self.proxy.transport.delete( f"/v1/mcp/server/{server_id}", headers=self.proxy.transport.master, json=NoBody(), response_type=NoBody, ) def seed_user_token(self, alias: str, key: str, storage_state_path: str) -> tuple[str, ...]: """Drive the interactive authorize dance for `key`'s user so the gateway stores their upstream token, retried to the shared deadline since the just-created server and key propagate asynchronously. The LiteLLM key rides x-litellm-api-key so the gateway binds the token to that user. Returns the upstream tool names the dance listed, proof the token works.""" headers = {"x-litellm-api-key": f"Bearer {key}"} storage = InMemoryTokenStorage() deadline = time.monotonic() + self.proxy.poll_timeout last_error: Exception | None = None while time.monotonic() < deadline: try: return asyncio.run(_seed_via_dance(_mcp_url(alias), headers, storage, storage_state_path)) except Exception as exc: # noqa: BLE001 - retried to the deadline; the last error surfaces below last_error = exc time.sleep(self.proxy.poll_interval) pytest.fail( f"authorize dance for {alias!r} never completed within {self.proxy.poll_timeout}s; " f"last error: {last_error!r}" ) def list_and_call( self, alias: str, headers: dict[str, str], storage: InMemoryTokenStorage, storage_state_path: str | None, tool: str, arguments: dict[str, str], base_url: str = PROXY_BASE_URL, identity: Identity | None = None, allow_upstream_consent: bool = True, ) -> OauthToolRun: return asyncio.run( _list_and_call( f"{base_url.rstrip('/')}/mcp" if identity is not None else _mcp_url(alias, base_url), headers, storage, storage_state_path, tool, arguments, base_url, identity, alias, allow_upstream_consent, ) ) def server_user_credentials(self, server_id: str) -> tuple[McpServerUserCredentialRow, ...]: return unwrap( self.proxy.transport.get( f"/v1/mcp/server/{server_id}/user-credentials", headers=self.proxy.transport.master, params=NoBody(), response_type=McpServerUserCredentialListResponse, ) ).root def revoke_user_token(self, server_id: str, headers: AuthHeaders) -> None: _ = unwrap( self.proxy.transport.delete( f"/v1/mcp/server/{server_id}/oauth-user-credential", headers=headers, json=NoBody(), response_type=McpOauthUserCredentialStatus, ) ) def chat_with_mcp(self, headers: AuthHeaders, body: ChatBody) -> ChatResponse: """POST /chat/completions carrying the LiteLLM key in `headers` (either ingress form) with an MCP server attached in `body.tools`. The gateway resolves the user from the key and lists/executes the server's tools with that user's stored upstream token.""" return unwrap( self.proxy.transport.post( "/chat/completions", headers=headers, json=body, response_type=ChatResponse, ) ) def build_chat_client(proxy: ProxyClient) -> ChatMcpClient: return ChatMcpClient(proxy=proxy)