diff --git a/tests/e2e/core/providers/test_catalog_oauth.py b/tests/e2e/core/providers/test_catalog_oauth.py new file mode 100644 index 0000000000..858938a110 --- /dev/null +++ b/tests/e2e/core/providers/test_catalog_oauth.py @@ -0,0 +1,324 @@ +"""Real-process E2E: OAuth / device-code providers against a loopback vendor. + +Every cell drives the real ``python -m hermes_cli.main`` in a hermetic HOME (fake HOME, HERMES_HOME +under it, no real credentials) against ``tests.fakes.providers.catalog_oauth.OAuthFake`` — the +vendor's OAuth authorization server and a bearer-checking inference server on 127.0.0.1. All other +egress goes through the ``CatalogFake`` sentinel proxy, which refuses and records any non-loopback +host. Cells assert user-visible outcomes: the device-code polling cadence the vendor sees, the +reply on stdout, the bearer on the next wire request, and the tokens persisted to auth.json. + +Open bugs are message-gated run-time xfails (``KNOWN`` + ``known_failure``): a cell XFAILs only +while it fails with that bug's signature and simply passes once the fix lands, in any merge order. + +Not redirectable to a loopback fake, so not covered here (explicit skips below): openai-codex and +qwen-oauth refresh (token URLs are module constants, no env/config override) and the Copilot token +exchange (hardcoded api.github.com). +""" + +from __future__ import annotations + +import json +import os +import subprocess +import sys +from datetime import datetime, timedelta, timezone +from pathlib import Path +from typing import Any + +import pytest + +from tests.e2e.core._pending_fixes import known_failure +from tests.fakes.providers.catalog_fake import CatalogFake +from tests.fakes.providers.catalog_oauth import NOUS_INVOKE_SCOPE, OAuthFake, make_jwt + +pytestmark = pytest.mark.skipif(sys.platform == "win32", reason="POSIX harness") + +REPO_ROOT = Path(__file__).resolve().parents[4] +PROD_NOUS_INFERENCE = "https://inference-api.nousresearch.com/v1" +TURN_TIMEOUT = 120.0 +# Public, credential-free model-metadata catalog (pricing/context lookups); never carries a vendor token. +CREDENTIAL_FREE_HOSTS = frozenset({"models.dev:443"}) + +# (pattern matched against the failing assertion text, reason). Delete an entry when its fix lands. +KNOWN: dict[str, tuple[str, str]] = { + "device_interval": ( + r"device poll gap .* < server interval", + "#121163 Nous device-code login polls at 1s, ignoring the server's interval"), + "nous_401_retry_route": ( + r"401 recovery retry left NOUS_INFERENCE_BASE_URL", + "#121323 Nous 401 pool recovery retries on the stored production host, " + "dropping the NOUS_INFERENCE_BASE_URL override"), +} + +_PASSTHROUGH_ENV = frozenset({"PATH", "LANG", "LANGUAGE", "USER", "LOGNAME", "SHELL", "TZ"}) +_SECRET_SUFFIXES = ("_API_KEY", "_TOKEN", "_SECRET", "_ACCESS_KEY") + + +# --- hermetic home --------------------------------------------------------------------------- + + +class Home: + def __init__(self, root: Path, sentinel: CatalogFake, provider: str, model: str, base_url: str = "") -> None: + self.home = root / "home" + self.hermes_home = self.home / ".hermes" + self.hermes_home.mkdir(parents=True) + self.sentinel = sentinel + (self.hermes_home / "config.yaml").write_text( + f"model:\n provider: {provider}\n default: {model}\n" + # The startup cost guard probes ``model.base_url``/models for pricing (credential-free), + # else the provider's production host; point it at the fake so any production-host + # egress the sentinel records is the credential-bearing turn itself. + + (f" base_url: {base_url}\n" if base_url else "") + + "updates:\n check: false\n" + "agent:\n api_max_retries: 1\n auto_recovery_cycles: 0\n") + + @property + def auth_path(self) -> Path: + return self.hermes_home / "auth.json" + + def seed_auth(self, store: dict[str, Any]) -> None: + self.auth_path.write_text(json.dumps(store, indent=2)) + + def auth(self) -> dict[str, Any]: + return json.loads(self.auth_path.read_text()) + + def env(self, extra: dict[str, str] | None = None) -> dict[str, str]: + env = {k: v for k, v in os.environ.items() + if (k in _PASSTHROUGH_ENV or k.startswith("LC_")) and not k.endswith(_SECRET_SUFFIXES)} + env.update({ + "HOME": str(self.home), "HERMES_HOME": str(self.hermes_home), "PYTHONPATH": str(REPO_ROOT), + "PYTHONUNBUFFERED": "1", "NO_COLOR": "1", "TERM": "dumb", + "TMPDIR": str(self.home), "HERMES_SHARED_AUTH_DIR": str(self.home / "shared"), + "CODEX_HOME": str(self.home / ".codex"), + # Child HOME is the fixture home, so its state.db is tmp_path's (guard's documented escape). + "HERMES_STATE_DB_GUARD_BYPASS": "1", + **self.sentinel.proxy_env()}) + env.update(extra or {}) + return env + + def run(self, argv: list[str], extra_env: dict[str, str] | None = None, + timeout: float = TURN_TIMEOUT) -> subprocess.CompletedProcess[str]: + return subprocess.run( + [sys.executable, "-m", "hermes_cli.main", *argv], cwd=str(self.home), env=self.env(extra_env), + stdin=subprocess.DEVNULL, capture_output=True, text=True, timeout=timeout) + + +def _iso(delta_s: float) -> str: + return (datetime.now(timezone.utc) + timedelta(seconds=delta_s)).isoformat() + + +def _nous_state(fake: OAuthFake, access: str, refresh: str, ttl_s: int) -> dict[str, Any]: + return {"access_token": access, "refresh_token": refresh, "token_type": "Bearer", + "scope": NOUS_INVOKE_SCOPE, "client_id": "hermes-cli", "portal_base_url": fake.origin, + "inference_base_url": PROD_NOUS_INFERENCE, "obtained_at": _iso(-60), "expires_at": _iso(ttl_s), + "agent_key": access, "agent_key_expires_at": _iso(ttl_s), + "tls": {"insecure": False, "ca_bundle": None}} + + +OTHER_NOUS_ROW = "nous-manual-2" +OPENROUTER_ROW = {"id": "or-1", "label": "openrouter-key", "auth_type": "api_key", "priority": 0, + "source": "manual", "access_token": "sk-or-oauth-e2e-untouched", + "base_url": "https://openrouter.ai/api/v1"} + + +def _nous_store(fake: OAuthFake, access: str, refresh: str, ttl_s: int) -> dict[str, Any]: + state = _nous_state(fake, access, refresh, ttl_s) + other_access = make_jwt("other-row", ttl_s=7200) + other = {**_nous_state(fake, other_access, "rt-other-row-untouched", 7200), "id": OTHER_NOUS_ROW, + "label": "second-login", "auth_type": "oauth", "priority": 1, "source": "manual:device_code"} + return {"version": 1, "active_provider": "nous", "providers": {"nous": state}, + "credential_pool": { + "nous": [{**state, "id": "nous-dc-1", "label": "seed", "auth_type": "oauth", "priority": 0, + "source": "device_code"}, other], + "openrouter": [OPENROUTER_ROW]}} + + +def _row(store: dict[str, Any], provider: str, row_id: str) -> dict[str, Any]: + rows = [r for r in store.get("credential_pool", {}).get(provider, []) if r.get("id") == row_id] + assert rows, f"credential_pool.{provider} lost row {row_id}: {store.get('credential_pool', {}).get(provider)}" + return rows[0] + + +def _assert_untouched(before: dict[str, Any], after: dict[str, Any]) -> None: + for provider, row_id in (("nous", OTHER_NOUS_ROW), ("openrouter", "or-1")): + want, got = _row(before, provider, row_id), _row(after, provider, row_id) + for key in ("access_token", "refresh_token"): + assert got.get(key) == want.get(key), ( + f"persisting the rotation clobbered credential_pool.{provider}[{row_id}].{key}: " + f"{want.get(key)!r} -> {got.get(key)!r}") + + +def _describe(proc: subprocess.CompletedProcess[str], fake: OAuthFake, sentinel: CatalogFake) -> str: + wire = [(r.method, r.path, r.bearer[-12:]) for r in fake.requests] + return (f"rc={proc.returncode}\nstdout={proc.stdout[-1500:]}\nstderr={proc.stderr[-2500:]}\n" + f"wire={wire}\negress={sentinel.egress_hosts()}") + + +def _vendor_egress(sentinel: CatalogFake) -> list[str]: + """Non-loopback hosts the child tried to reach, minus the credential-free metadata catalog.""" + return [h for h in sentinel.egress_hosts() if h not in CREDENTIAL_FREE_HOSTS] + + +@pytest.fixture +def sentinel(): + with CatalogFake() as s: + yield s + + +# --- Nous device-code login ------------------------------------------------------------------ + +# Two pendings around one slow_down: the gaps before slow_down show the base interval the client +# honours, the gaps after it show the +1s RFC 8628 growth. +DEVICE_SCRIPT = ["authorization_pending", "slow_down", "authorization_pending"] +DEVICE_INTERVAL = 2 + + +@pytest.fixture(scope="module") +def device_login(tmp_path_factory): + """One real ``hermes auth add nous --type oauth`` device-code login against the fake Portal.""" + root = tmp_path_factory.mktemp("nous-device") + with CatalogFake() as sentinel, OAuthFake(device_interval=DEVICE_INTERVAL, poll_script=DEVICE_SCRIPT) as fake: + home = Home(root, sentinel, "nous", "oauth-e2e/model") + seed = {"version": 1, "credential_pool": {"openrouter": [OPENROUTER_ROW]}} + home.seed_auth(seed) + proc = home.run(["auth", "add", "nous", "--type", "oauth", "--no-browser", "--portal-url", fake.origin], + extra_env={"NOUS_INFERENCE_BASE_URL": f"{fake.origin}/v1"}, timeout=90) + yield {"proc": proc, "fake": fake, "home": home, "seed": seed, "sentinel": sentinel, + "gaps": [b.t - a.t for a, b in zip(fake.device_polls(), fake.device_polls()[1:])]} + + +def test_nous_device_login_persists_tokens(device_login) -> None: + proc, fake, home = device_login["proc"], device_login["fake"], device_login["home"] + info = _describe(proc, fake, device_login["sentinel"]) + assert proc.returncode == 0, f"device-code login failed\n{info}" + assert len(fake.device_polls()) == len(DEVICE_SCRIPT) + 1, f"login stopped polling early\n{info}" + assert "OAUT-HE2E" in proc.stdout, f"user code never shown to the user\n{info}" + issued = fake.issued[0] + store = home.auth() + state = store.get("providers", {}).get("nous") or {} + assert state.get("refresh_token") == issued["refresh_token"], ( + f"providers.nous did not persist the login's refresh token: {state.get('refresh_token')!r}\n{info}") + assert state.get("access_token") == issued["access_token"], f"providers.nous access token not persisted\n{info}" + pooled = [r.get("refresh_token") for r in store.get("credential_pool", {}).get("nous", [])] + assert issued["refresh_token"] in pooled, f"credential_pool.nous missing the login: {pooled}\n{info}" + assert _row(store, "openrouter", "or-1")["access_token"] == OPENROUTER_ROW["access_token"], ( + "login clobbered an unrelated credential_pool row") + assert not _vendor_egress(device_login["sentinel"]), f"login leaked egress: {_vendor_egress(device_login['sentinel'])}" + + +def test_nous_device_login_slow_down_grows_interval(device_login) -> None: + gaps = device_login["gaps"] + assert len(gaps) == len(DEVICE_SCRIPT), f"unexpected poll count, gaps={gaps}" + before, after = gaps[0], gaps[1:] + assert all(g >= before + 0.9 for g in after), ( + f"slow_down did not grow the poll interval by >=1s: gap before slow_down {before:.2f}s, " + f"after {[round(g, 2) for g in after]}") + + +def test_nous_device_login_honors_server_interval(device_login) -> None: + gaps = device_login["gaps"] + assert len(gaps) == len(DEVICE_SCRIPT), f"unexpected poll count, gaps={gaps}" + with known_failure(*KNOWN["device_interval"]): + assert gaps[0] >= DEVICE_INTERVAL - 0.1, ( + f"device poll gap {gaps[0]:.2f}s < server interval {DEVICE_INTERVAL}s (gaps={[round(g, 2) for g in gaps]})") + + +# --- Nous token refresh ---------------------------------------------------------------------- + + +def _nous_turn(tmp_path: Path, sentinel: CatalogFake, fake: OAuthFake, access: str, ttl_s: int): + home = Home(tmp_path, sentinel, "nous", "oauth-e2e/model", base_url=f"{fake.origin}/v1") + seed = _nous_store(fake, access, "rt-seed-0", ttl_s) + home.seed_auth(seed) + proc = home.run(["-z", "Say hi", "--provider", "nous", "-m", "oauth-e2e/model"], + extra_env={"NOUS_INFERENCE_BASE_URL": f"{fake.origin}/v1"}) + return home, seed, proc + + +def _assert_rotation_persisted(home: Home, seed: dict[str, Any], fake: OAuthFake, info: str) -> str: + refreshes = fake.refreshes() + assert refreshes, f"no refresh-token exchange reached the Portal\n{info}" + assert refreshes[0].headers.get("x-nous-refresh-token") == "rt-seed-0", ( + f"refresh redeemed the wrong token: {refreshes[0].headers.get('x-nous-refresh-token')!r}\n{info}") + rotated = fake.issued[-1] + store = home.auth() + assert store["providers"]["nous"].get("refresh_token") == rotated["refresh_token"], ( + f"providers.nous kept a spent refresh token {store['providers']['nous'].get('refresh_token')!r} " + f"instead of the rotation {rotated['refresh_token']!r}\n{info}") + assert _row(store, "nous", "nous-dc-1").get("refresh_token") == rotated["refresh_token"], ( + f"credential_pool.nous[nous-dc-1] kept a spent refresh token " + f"{_row(store, 'nous', 'nous-dc-1').get('refresh_token')!r}\n{info}") + _assert_untouched(seed, store) + return rotated["access_token"] + + +def test_nous_expired_access_token_refreshes_before_turn(tmp_path, sentinel) -> None: + stale = make_jwt("expired", ttl_s=-60) + with OAuthFake(valid_refresh={"rt-seed-0"}, revoked={stale}) as fake: + home, seed, proc = _nous_turn(tmp_path, sentinel, fake, stale, ttl_s=-60) + info = _describe(proc, fake, sentinel) + assert proc.returncode == 0 and fake.reply in proc.stdout, f"turn with an expired token failed\n{info}" + fresh = _assert_rotation_persisted(home, seed, fake, info) + bearers = {r.bearer for r in fake.inference()} + assert bearers == {fresh}, f"inference used {bearers}, expected only the refreshed token\n{info}" + assert len(fake.refreshes()) == 1, f"refresh token redeemed {len(fake.refreshes())}x\n{info}" + assert not _vendor_egress(sentinel), f"turn leaked egress: {_vendor_egress(sentinel)}" + + +def test_nous_inference_401_refreshes_rotates_and_retries(tmp_path, sentinel) -> None: + revoked = make_jwt("revoked-by-server", ttl_s=7200) + with OAuthFake(valid_refresh={"rt-seed-0"}, revoked={revoked}) as fake: + home, seed, proc = _nous_turn(tmp_path, sentinel, fake, revoked, ttl_s=7200) + info = _describe(proc, fake, sentinel) + assert any(r.bearer == revoked for r in fake.inference()), f"the stale bearer was never tried\n{info}" + fresh = _assert_rotation_persisted(home, seed, fake, info) + with known_failure(*KNOWN["nous_401_retry_route"]): + assert not _vendor_egress(sentinel), ( + f"401 recovery retry left NOUS_INFERENCE_BASE_URL: egress to {_vendor_egress(sentinel)}\n{info}") + assert any(r.bearer == fresh for r in fake.inference()), f"no retry with the refreshed token\n{info}" + assert proc.returncode == 0 and fake.reply in proc.stdout, f"turn failed after refresh\n{info}" + + +# --- MiniMax OAuth refresh ------------------------------------------------------------------- + + +def test_minimax_oauth_expired_token_refreshes_and_persists(tmp_path, sentinel) -> None: + with OAuthFake(valid_refresh={"rt-mm-seed"}, revoked={"mm-stale-access"}) as fake: + home = Home(tmp_path, sentinel, "minimax-oauth", "MiniMax-M2") + seed = {"version": 1, "active_provider": "minimax-oauth", "providers": {"minimax-oauth": { + "portal_base_url": fake.origin, "inference_base_url": f"{fake.origin}/anthropic", + "client_id": "oauth-e2e-client", "access_token": "mm-stale-access", "refresh_token": "rt-mm-seed", + "expires_at": _iso(-60), "region": "global", "token_type": "Bearer", "scope": "group_id profile"}}, + "credential_pool": {"openrouter": [OPENROUTER_ROW]}} + home.seed_auth(seed) + proc = home.run(["-z", "Say hi", "--provider", "minimax-oauth", "-m", "MiniMax-M2"]) + info = _describe(proc, fake, sentinel) + assert proc.returncode == 0 and fake.reply in proc.stdout, f"MiniMax turn failed\n{info}" + assert fake.refreshes() and fake.refreshes()[0].form.get("refresh_token") == "rt-mm-seed", ( + f"MiniMax refresh did not redeem the stored refresh token\n{info}") + rotated = fake.issued[-1] + bearers = {r.bearer for r in fake.inference()} + assert bearers == {rotated["access_token"]}, f"inference used {bearers}, not the refreshed token\n{info}" + state = home.auth()["providers"]["minimax-oauth"] + assert state.get("refresh_token") == rotated["refresh_token"], ( + f"providers.minimax-oauth kept spent refresh token {state.get('refresh_token')!r}\n{info}") + assert state.get("access_token") == rotated["access_token"], f"MiniMax access token not persisted\n{info}" + assert _row(home.auth(), "openrouter", "or-1")["access_token"] == OPENROUTER_ROW["access_token"] + assert not _vendor_egress(sentinel), f"turn leaked egress: {_vendor_egress(sentinel)}" + + +# --- not redirectable ------------------------------------------------------------------------ + +_UNREDIRECTABLE = { + "openai-codex": "token refresh URL is the module constant CODEX_OAUTH_TOKEN_URL (auth.openai.com); " + "no env/config override, so a refresh cannot reach a loopback fake", + "qwen-oauth": "token refresh URL is the module constant QWEN_OAUTH_TOKEN_URL (chat.qwen.ai); " + "no env/config override", + "copilot": "token exchange URL is hardcoded to api.github.com", +} + + +@pytest.mark.parametrize("provider", sorted(_UNREDIRECTABLE)) +def test_oauth_refresh_not_redirectable(provider: str) -> None: + pytest.skip(f"{provider}: {_UNREDIRECTABLE[provider]}") + diff --git a/tests/fakes/providers/catalog_fake.py b/tests/fakes/providers/catalog_fake.py new file mode 100644 index 0000000000..a34794f42a --- /dev/null +++ b/tests/fakes/providers/catalog_fake.py @@ -0,0 +1,363 @@ +"""Multi-dialect recording loopback provider for the provider-catalog E2E matrix. + +One real HTTP server on 127.0.0.1 that answers every wire dialect the bundled +model-provider plugins speak, routed by request path: + +* ``POST …/chat/completions`` — OpenAI Chat Completions (JSON or SSE) +* ``POST …/messages`` — Anthropic Messages (JSON or SSE) +* ``POST …/responses`` — OpenAI Responses (SSE event stream) +* ``GET …/models`` — model listing (OpenAI ``data`` shape), or a + scripted 404 / hang so picker fallbacks can be exercised + +Every request (method, path, headers, body) is recorded, so a test can assert on +exactly which credential reached which host in which header. + +The model script is stateless and dialect-neutral: a main-turn request (one that +offers tools) whose history carries no tool result gets ONE tool call +(``tool_name``/``tool_args``); once a tool result is present the model answers +``final_text``. Requests without tools are auxiliary (titles, summaries) and get +a short plain answer. ``fail_status`` turns every inference request into that +HTTP error (a dead provider for fallback tests). +""" + +from __future__ import annotations + +import json +import threading +import time +import uuid +from dataclasses import dataclass, field +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from typing import Any + +USAGE_IN = 1234 +USAGE_OUT = 56 + + +@dataclass +class Recorded: + method: str + path: str + headers: dict[str, str] + body: Any + t: float = field(default_factory=time.time) + + @property + def dialect(self) -> str: + return path_dialect(self.path) if self.method == "POST" else "listing" + + +def path_dialect(path: str) -> str: + p = path.split("?", 1)[0].rstrip("/") + for suffix, name in (("/chat/completions", "chat"), ("/messages", "anthropic"), ("/responses", "responses")): + if p.endswith(suffix): + return name + return "unknown" + + +class CatalogFake: + """Threaded loopback server. Use as a context manager.""" + + def __init__( + self, + *, + tool_name: str = "read_file", + tool_args: dict[str, Any] | None = None, + final_text: str = "CATALOG-TURN-COMPLETE", + models: list[str] | None = None, + models_status: int = 200, + models_hang_s: float = 0.0, + fail_status: int | None = None, + ) -> None: + self.tool_name = tool_name + self.tool_args = tool_args or {} + self.final_text = final_text + self.models = list(models or ["catalog-model-a", "catalog-model-b"]) + self.models_status = models_status + self.models_hang_s = models_hang_s + self.fail_status = fail_status + self.requests: list[Recorded] = [] + # Egress sentinel: when a child process gets ``HTTPS_PROXY``/``HTTP_PROXY`` = ``origin`` + # (see ``proxy_env``), every request aimed at a NON-loopback host lands here instead of + # the real vendor and is refused with 403 — ``egress`` names each host it tried to reach. + self.egress: list[Recorded] = [] + self._lock = threading.Lock() + self._stop = threading.Event() + self._server: ThreadingHTTPServer | None = None + + def __enter__(self) -> "CatalogFake": + server = ThreadingHTTPServer(("127.0.0.1", 0), _handler_for(self)) + server.daemon_threads = True + self._server = server + threading.Thread(target=server.serve_forever, name="catalog-fake", daemon=True).start() + return self + + def __exit__(self, *_exc: object) -> None: + self._stop.set() + if self._server is not None: + self._server.shutdown() + self._server.server_close() + + @property + def origin(self) -> str: + assert self._server is not None, "server not started" + return f"http://127.0.0.1:{self._server.server_address[1]}" + + def proxy_env(self) -> dict[str, str]: + """Env that routes every non-loopback request of a child through the egress sentinel.""" + return {"HTTPS_PROXY": self.origin, "HTTP_PROXY": self.origin, "https_proxy": self.origin, + "http_proxy": self.origin, "ALL_PROXY": "", "all_proxy": "", + "NO_PROXY": "127.0.0.1,localhost", "no_proxy": "127.0.0.1,localhost"} + + def egress_hosts(self) -> list[str]: + return sorted({r.path for r in self.egress}) + + def inference(self) -> list[Recorded]: + return [r for r in self.requests if r.method == "POST"] + + def listings(self) -> list[Recorded]: + return [r for r in self.requests if r.method == "GET"] + + def _record(self, rec: Recorded) -> None: + with self._lock: + self.requests.append(rec) + + +# --- dialect-neutral script --------------------------------------------------- + + +def _has_tool_result(dialect: str, body: dict[str, Any]) -> bool: + if dialect == "chat": + return any(m.get("role") == "tool" for m in body.get("messages") or []) + if dialect == "anthropic": + return any( + isinstance(m.get("content"), list) and any( + isinstance(b, dict) and b.get("type") == "tool_result" for b in m["content"]) + for m in body.get("messages") or []) + items = body.get("input") if isinstance(body.get("input"), list) else [] + return any(isinstance(i, dict) and i.get("type") == "function_call_output" for i in items) + + +def _plan(fake: CatalogFake, dialect: str, body: dict[str, Any]) -> tuple[str, str | None]: + """(``text``, ``tool_call_id`` or None) for the next model turn.""" + if not body.get("tools"): + return "Catalog aux answer", None + if _has_tool_result(dialect, body): + return fake.final_text, None + return "", f"call_{uuid.uuid4().hex[:10]}" + + +# --- wire renderers ------------------------------------------------------------- + + +def _chat_payloads(fake: CatalogFake, text: str, call_id: str | None, stream: bool) -> list[dict[str, Any]] | dict: + usage = {"prompt_tokens": USAGE_IN, "completion_tokens": USAGE_OUT, "total_tokens": USAGE_IN + USAGE_OUT} + base = {"id": "chatcmpl-catalog", "created": int(time.time()), "model": "catalog-model-a"} + args = json.dumps(fake.tool_args) + tool_calls = [{"id": call_id, "type": "function", "function": {"name": fake.tool_name, "arguments": args}}] + finish = "tool_calls" if call_id else "stop" + if not stream: + msg: dict[str, Any] = {"role": "assistant", "content": text or None} + if call_id: + msg["tool_calls"] = tool_calls + return {**base, "object": "chat.completion", "usage": usage, + "choices": [{"index": 0, "message": msg, "finish_reason": finish}]} + + def chunk(delta: dict[str, Any], fin: str | None = None) -> dict[str, Any]: + return {**base, "object": "chat.completion.chunk", "choices": [{"index": 0, "delta": delta, "finish_reason": fin}]} + + out = [chunk({"role": "assistant", "content": ""})] + if call_id: + out.append(chunk({"tool_calls": [{"index": 0, "id": call_id, "type": "function", + "function": {"name": fake.tool_name, "arguments": ""}}]})) + out.append(chunk({"tool_calls": [{"index": 0, "function": {"arguments": args}}]})) + else: + out.append(chunk({"content": text})) + last = chunk({}, finish) + last["usage"] = usage + out.append(last) + return out + + +def _anthropic_events(fake: CatalogFake, text: str, call_id: str | None) -> list[tuple[str, dict[str, Any]]]: + msg = {"id": "msg_catalog", "type": "message", "role": "assistant", "model": "catalog-model-a", + "content": [], "stop_reason": None, "stop_sequence": None, + "usage": {"input_tokens": USAGE_IN, "output_tokens": 1}} + ev: list[tuple[str, dict[str, Any]]] = [("message_start", {"type": "message_start", "message": msg})] + if call_id: + ev += [ + ("content_block_start", {"type": "content_block_start", "index": 0, "content_block": { + "type": "tool_use", "id": call_id.replace("call_", "toolu_"), "name": fake.tool_name, "input": {}}}), + ("content_block_delta", {"type": "content_block_delta", "index": 0, "delta": { + "type": "input_json_delta", "partial_json": json.dumps(fake.tool_args)}}), + ] + else: + ev += [ + ("content_block_start", {"type": "content_block_start", "index": 0, + "content_block": {"type": "text", "text": ""}}), + ("content_block_delta", {"type": "content_block_delta", "index": 0, + "delta": {"type": "text_delta", "text": text}}), + ] + ev += [ + ("content_block_stop", {"type": "content_block_stop", "index": 0}), + ("message_delta", {"type": "message_delta", "delta": { + "stop_reason": "tool_use" if call_id else "end_turn", "stop_sequence": None}, + "usage": {"output_tokens": USAGE_OUT}}), + ("message_stop", {"type": "message_stop"}), + ] + return ev + + +def _anthropic_json(fake: CatalogFake, text: str, call_id: str | None) -> dict[str, Any]: + content: list[dict[str, Any]] = ( + [{"type": "tool_use", "id": call_id.replace("call_", "toolu_"), "name": fake.tool_name, "input": fake.tool_args}] + if call_id else [{"type": "text", "text": text}]) + return {"id": "msg_catalog", "type": "message", "role": "assistant", "model": "catalog-model-a", + "content": content, "stop_reason": "tool_use" if call_id else "end_turn", "stop_sequence": None, + "usage": {"input_tokens": USAGE_IN, "output_tokens": USAGE_OUT}} + + +def _responses_events(fake: CatalogFake, text: str, call_id: str | None) -> list[dict[str, Any]]: + rid = f"resp_{uuid.uuid4().hex[:10]}" + usage = {"input_tokens": USAGE_IN, "output_tokens": USAGE_OUT, "total_tokens": USAGE_IN + USAGE_OUT, + "input_tokens_details": {"cached_tokens": 0}, "output_tokens_details": {"reasoning_tokens": 0}} + shell = {"id": rid, "object": "response", "created_at": int(time.time()), "model": "catalog-model-a", + "status": "in_progress", "output": []} + if call_id: + item = {"type": "function_call", "id": f"fc_{call_id}", "call_id": call_id, "name": fake.tool_name, + "arguments": json.dumps(fake.tool_args), "status": "completed"} + middle = [ + {"type": "response.output_item.added", "output_index": 0, "item": {**item, "arguments": "", "status": "in_progress"}}, + {"type": "response.function_call_arguments.delta", "output_index": 0, "item_id": item["id"], + "delta": item["arguments"]}, + {"type": "response.output_item.done", "output_index": 0, "item": item}, + ] + else: + item = {"type": "message", "id": f"msg_{rid}", "role": "assistant", "status": "completed", + "content": [{"type": "output_text", "text": text, "annotations": []}]} + middle = [ + {"type": "response.output_item.added", "output_index": 0, + "item": {**item, "content": [], "status": "in_progress"}}, + {"type": "response.output_text.delta", "output_index": 0, "content_index": 0, "item_id": item["id"], + "delta": text}, + {"type": "response.output_item.done", "output_index": 0, "item": item}, + ] + done = {**shell, "status": "completed", "output": [item], "usage": usage} + events = [{"type": "response.created", "response": shell}, *middle, + {"type": "response.completed", "response": done}] + for seq, e in enumerate(events): + e["sequence_number"] = seq + return events + + +def _handler_for(fake: CatalogFake) -> type[BaseHTTPRequestHandler]: + class Handler(BaseHTTPRequestHandler): + protocol_version = "HTTP/1.1" + + def log_message(self, *_a: object) -> None: + pass + + def _headers(self) -> dict[str, str]: + return {k.lower(): v for k, v in self.headers.items()} + + def _json(self, status: int, payload: Any) -> None: + body = json.dumps(payload).encode() + self.send_response(status) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", str(len(body))) + self.end_headers() + self.wfile.write(body) + + def _sse_open(self) -> None: + self.send_response(200) + self.send_header("Content-Type", "text/event-stream") + self.send_header("Cache-Control", "no-cache") + self.send_header("Connection", "close") + self.end_headers() + self.close_connection = True + + def _sse(self, data: Any, event: str | None = None) -> None: + prefix = f"event: {event}\n" if event else "" + self.wfile.write(f"{prefix}data: {json.dumps(data)}\n\n".encode()) + self.wfile.flush() + + def _refuse_egress(self, method: str) -> bool: + """Proxy-form requests (CONNECT host:port, or an absolute-URI GET/POST) are egress.""" + if method != "CONNECT" and not self.path.startswith(("http://", "https://")): + return False + target = self.path.split("://", 1)[-1].split("/", 1)[0] + with fake._lock: + fake.egress.append(Recorded(method, target, self._headers(), None)) + self._json(403, {"error": {"message": f"catalog egress sentinel refused {target}"}}) + self.close_connection = True + return True + + def do_CONNECT(self) -> None: # noqa: N802 + self._refuse_egress("CONNECT") + + def do_GET(self) -> None: # noqa: N802 + if self._refuse_egress("GET"): + return + fake._record(Recorded("GET", self.path, self._headers(), None)) + if not self.path.split("?", 1)[0].rstrip("/").endswith("/models"): + self._json(404, {"error": {"message": "not found"}}) + return + if fake.models_hang_s: + fake._stop.wait(fake.models_hang_s) + self.close_connection = True + return + if fake.models_status != 200: + self._json(fake.models_status, {"error": {"message": "no listing here"}}) + return + self._json(200, {"object": "list", "data": [ + {"id": m, "object": "model", "type": "model", "display_name": m, "created": 1, + "created_at": "2026-01-01T00:00:00Z", "owned_by": "catalog", "context_length": 131072} + for m in fake.models]}) + + def do_POST(self) -> None: # noqa: N802 + raw = self.rfile.read(int(self.headers.get("Content-Length", 0) or 0)) + if self._refuse_egress("POST"): + return + try: + body = json.loads(raw or b"{}") + except json.JSONDecodeError: + body = {"_raw": raw.decode("utf-8", "replace")} + fake._record(Recorded("POST", self.path, self._headers(), body)) + dialect = path_dialect(self.path) + if dialect == "unknown": + self._json(404, {"error": {"message": f"unsupported path {self.path}"}}) + return + if fake.fail_status is not None: + self._json(fake.fail_status, {"error": { + "message": "catalog fake: provider down", "type": "server_error"}}) + return + text, call_id = _plan(fake, dialect, body) + stream = bool(body.get("stream")) + if dialect == "chat": + payload = _chat_payloads(fake, text, call_id, stream) + if not stream: + self._json(200, payload) + return + self._sse_open() + for c in payload: + self._sse(c) + self.wfile.write(b"data: [DONE]\n\n") + self.wfile.flush() + return + if dialect == "anthropic": + if not stream: + self._json(200, _anthropic_json(fake, text, call_id)) + return + self._sse_open() + for name, data in _anthropic_events(fake, text, call_id): + self._sse(data, event=name) + return + events = _responses_events(fake, text, call_id) + if not stream: + self._json(200, events[-1]["response"]) + return + self._sse_open() + for e in events: + self._sse(e, event=e["type"]) + + return Handler diff --git a/tests/fakes/providers/catalog_oauth.py b/tests/fakes/providers/catalog_oauth.py new file mode 100644 index 0000000000..e2df49327c --- /dev/null +++ b/tests/fakes/providers/catalog_oauth.py @@ -0,0 +1,278 @@ +"""Loopback OAuth authorization server + bearer-checking inference server for the OAuth E2E cells. + +One real HTTP server on 127.0.0.1 that plays the vendor side of the OAuth providers: + +* Nous Portal (RFC 8628 device flow + single-use refresh-token rotation): + ``POST /api/oauth/device/code`` and ``POST /api/oauth/token`` (``grant_type`` device_code or + refresh_token; the refresh token rides in the ``x-nous-refresh-token`` header). The device poll + answers the scripted error codes in ``poll_script`` (one per poll) and then issues tokens; every + poll's arrival time is recorded so a test can measure the client's real polling cadence. +* MiniMax OAuth refresh: ``POST /oauth/token`` (form ``refresh_token``), MiniMax's + ``{"status": "success", "expired_in": ...}`` shape. +* Inference: ``POST …/chat/completions`` (JSON or SSE) and ``POST …/messages`` (Anthropic JSON or + SSE). Every inference request is bearer-checked: a token in ``revoked`` gets HTTP 401 like a + real gateway rejecting a revoked/expired key; any other token gets a plain answer ``reply``. + +Refresh tokens are single-use, as at the real vendors: redeeming one retires it, and redeeming a +retired one returns ``invalid_grant``. Every request is recorded (method, path, headers, form or +JSON body, arrival time). +""" + +from __future__ import annotations + +import base64 +import json +import threading +import time +import urllib.parse +from dataclasses import dataclass, field +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from typing import Any + +NOUS_INVOKE_SCOPE = "inference:invoke" + + +def b64url(obj: dict[str, Any]) -> str: + return base64.urlsafe_b64encode(json.dumps(obj).encode()).rstrip(b"=").decode() + + +def make_jwt(tag: str, *, ttl_s: int = 3600, scope: str = NOUS_INVOKE_SCOPE) -> str: + """Unsigned JWT-shaped bearer: the client only decodes claims (``scope``/``exp``).""" + claims = {"sub": "oauth-e2e-user", "scope": scope, "exp": int(time.time()) + ttl_s, "tag": tag} + return ".".join([b64url({"alg": "none", "typ": "JWT"}), b64url(claims), "sig"]) + + +@dataclass +class Req: + method: str + path: str + headers: dict[str, str] + form: dict[str, str] + body: Any + t: float = field(default_factory=time.monotonic) + + @property + def bearer(self) -> str: + auth = self.headers.get("authorization", "") + return auth[7:] if auth.lower().startswith("bearer ") else self.headers.get("x-api-key", "") + + +class OAuthFake: + """Threaded loopback vendor. Use as a context manager.""" + + def __init__(self, *, reply: str = "OAUTH-TURN-COMPLETE", device_interval: int = 2, + poll_script: list[str] | None = None, valid_refresh: set[str] | None = None, + revoked: set[str] | None = None) -> None: + self.reply = reply + self.device_interval = device_interval + self.poll_script = list(poll_script or []) + self.valid_refresh = set(valid_refresh or ()) + self.revoked = set(revoked or ()) + self.requests: list[Req] = [] + self.issued: list[dict[str, Any]] = [] # every token response, in order + self._lock = threading.Lock() + self._server: ThreadingHTTPServer | None = None + + def __enter__(self) -> "OAuthFake": + server = ThreadingHTTPServer(("127.0.0.1", 0), _handler_for(self)) + server.daemon_threads = True + self._server = server + threading.Thread(target=server.serve_forever, name="oauth-fake", daemon=True).start() + return self + + def __exit__(self, *_exc: object) -> None: + if self._server is not None: + self._server.shutdown() + self._server.server_close() + + @property + def origin(self) -> str: + assert self._server is not None, "server not started" + return f"http://127.0.0.1:{self._server.server_address[1]}" + + # --- views ----------------------------------------------------------------------------- + def device_polls(self) -> list[Req]: + return [r for r in self.requests + if r.path == "/api/oauth/token" and r.form.get("grant_type", "").endswith("device_code")] + + def refreshes(self) -> list[Req]: + return [r for r in self.requests if r.path in ("/api/oauth/token", "/oauth/token") + and r.form.get("grant_type") == "refresh_token"] + + def inference(self) -> list[Req]: + return [r for r in self.requests + if r.method == "POST" and r.path.split("?")[0].endswith(("/chat/completions", "/messages"))] + + # --- token issuance -------------------------------------------------------------------- + def _issue(self, kind: str) -> dict[str, Any]: + with self._lock: + n = len(self.issued) + 1 + rt = f"rt-{kind}-{n}" + self.valid_refresh.add(rt) + tok = {"access_token": make_jwt(f"{kind}-{n}"), "refresh_token": rt, "token_type": "Bearer", + "expires_in": 3600, "scope": NOUS_INVOKE_SCOPE} + self.issued.append(tok) + return tok + + def _redeem(self, refresh_token: str) -> bool: + with self._lock: + if refresh_token not in self.valid_refresh: + return False + self.valid_refresh.discard(refresh_token) + return True + + +def _handler_for(fake: OAuthFake) -> type[BaseHTTPRequestHandler]: + class Handler(BaseHTTPRequestHandler): + protocol_version = "HTTP/1.1" + + def log_message(self, *_a: object) -> None: + pass + + def _json(self, status: int, payload: Any) -> None: + data = json.dumps(payload).encode() + self.send_response(status) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", str(len(data))) + self.end_headers() + self.wfile.write(data) + + def _sse(self, frames: list[tuple[str | None, Any]], done: bool) -> None: + self.send_response(200) + self.send_header("Content-Type", "text/event-stream") + self.send_header("Connection", "close") + self.end_headers() + self.close_connection = True + for event, data in frames: + prefix = f"event: {event}\n" if event else "" + self.wfile.write(f"{prefix}data: {json.dumps(data)}\n\n".encode()) + if done: + self.wfile.write(b"data: [DONE]\n\n") + self.wfile.flush() + + def do_GET(self) -> None: # noqa: N802 + headers = {k.lower(): v for k, v in self.headers.items()} + with fake._lock: + fake.requests.append(Req("GET", self.path, headers, {}, None)) + if self.path.split("?")[0].rstrip("/").endswith("/models"): + self._json(200, {"object": "list", "data": [{"id": "oauth-e2e/model", "object": "model"}]}) + return + self._json(404, {"error": "not_found"}) + + def do_POST(self) -> None: # noqa: N802 + raw = self.rfile.read(int(self.headers.get("Content-Length") or 0)).decode("utf-8", "replace") + headers = {k.lower(): v for k, v in self.headers.items()} + ctype = headers.get("content-type", "") + form = ({k: v[0] for k, v in urllib.parse.parse_qs(raw).items()} + if "x-www-form-urlencoded" in ctype else {}) + try: + body = json.loads(raw) if raw.startswith(("{", "[")) else None + except json.JSONDecodeError: + body = None + req = Req("POST", self.path, headers, form, body) + with fake._lock: + fake.requests.append(req) + path = self.path.split("?")[0].rstrip("/") + route = _ROUTES.get(path) or next( + (fn for suffix, fn in _SUFFIX_ROUTES if path.endswith(suffix)), None) + if route is None: + self._json(404, {"error": {"message": f"unsupported path {self.path}"}}) + return + route(self, fake, req) + + return Handler + + +# --- routes ----------------------------------------------------------------------------------- + + +def _device_code(h: Any, fake: OAuthFake, _req: Req) -> None: + h._json(200, {"device_code": "dc-oauth-e2e", "user_code": "OAUT-HE2E", + "verification_uri": f"{fake.origin}/device", + "verification_uri_complete": f"{fake.origin}/device?code=OAUT-HE2E", + "expires_in": 120, "interval": fake.device_interval}) + + +def _nous_token(h: Any, fake: OAuthFake, req: Req) -> None: + grant = req.form.get("grant_type", "") + if grant.endswith("device_code"): + with fake._lock: + step = fake.poll_script.pop(0) if fake.poll_script else None + if step is not None: + h._json(400, {"error": step, "error_description": f"scripted {step}"}) + return + h._json(200, fake._issue("login")) + return + if grant == "refresh_token": + if not fake._redeem(req.headers.get("x-nous-refresh-token", "")): + h._json(400, {"error": "invalid_grant", "error_description": "refresh token already used or unknown"}) + return + h._json(200, fake._issue("nous-rot")) + return + h._json(400, {"error": "unsupported_grant_type"}) + + +def _minimax_token(h: Any, fake: OAuthFake, req: Req) -> None: + if req.form.get("grant_type") != "refresh_token" or not fake._redeem(req.form.get("refresh_token", "")): + h._json(400, {"status": "error", "error": "invalid_grant"}) + return + tok = fake._issue("mm-rot") + h._json(200, {"status": "success", "access_token": tok["access_token"], + "refresh_token": tok["refresh_token"], "expired_in": 3600, "token_type": "Bearer"}) + + +def _reject_revoked(h: Any, fake: OAuthFake, req: Req) -> bool: + if req.bearer and req.bearer not in fake.revoked: + return False + h._json(401, {"error": {"message": "invalid or revoked access token", "type": "authentication_error", + "code": "invalid_api_key"}}) + return True + + +def _chat(h: Any, fake: OAuthFake, req: Req) -> None: + if _reject_revoked(h, fake, req): + return + body = req.body or {} + usage = {"prompt_tokens": 10, "completion_tokens": 2, "total_tokens": 12} + base = {"id": "chatcmpl-oauth", "created": int(time.time()), "model": body.get("model") or "m"} + if not body.get("stream"): + h._json(200, {**base, "object": "chat.completion", "usage": usage, "choices": [ + {"index": 0, "message": {"role": "assistant", "content": fake.reply}, "finish_reason": "stop"}]}) + return + chunk = {**base, "object": "chat.completion.chunk"} + h._sse([(None, {**chunk, "choices": [{"index": 0, "delta": {"role": "assistant", "content": fake.reply}, + "finish_reason": None}]}), + (None, {**chunk, "choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}], "usage": usage})], + done=True) + + +def _messages(h: Any, fake: OAuthFake, req: Req) -> None: + if _reject_revoked(h, fake, req): + return + body = req.body or {} + msg = {"id": "msg_oauth", "type": "message", "role": "assistant", "model": body.get("model") or "m", + "stop_sequence": None} + if not body.get("stream"): + h._json(200, {**msg, "content": [{"type": "text", "text": fake.reply}], "stop_reason": "end_turn", + "usage": {"input_tokens": 10, "output_tokens": 2}}) + return + h._sse([ + ("message_start", {"type": "message_start", "message": { + **msg, "content": [], "stop_reason": None, "usage": {"input_tokens": 10, "output_tokens": 1}}}), + ("content_block_start", {"type": "content_block_start", "index": 0, + "content_block": {"type": "text", "text": ""}}), + ("content_block_delta", {"type": "content_block_delta", "index": 0, + "delta": {"type": "text_delta", "text": fake.reply}}), + ("content_block_stop", {"type": "content_block_stop", "index": 0}), + ("message_delta", {"type": "message_delta", "delta": {"stop_reason": "end_turn", "stop_sequence": None}, + "usage": {"output_tokens": 2}}), + ("message_stop", {"type": "message_stop"}), + ], done=False) + + +_ROUTES = { + "/api/oauth/device/code": _device_code, + "/api/oauth/token": _nous_token, + "/oauth/token": _minimax_token, +} +_SUFFIX_ROUTES = (("/chat/completions", _chat), ("/messages", _messages))