test: provider-catalog OAuth E2E (device-code cadence, refresh rotation) + multi-dialect loopback fake

This commit is contained in:
teknium1
2026-09-24 02:55:28 -07:00
committed by Teknium
parent a8a549ebb4
commit 063cdcaa06
3 changed files with 965 additions and 0 deletions

View File

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

View File

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

View File

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