test: provider-catalog OAuth E2E (device-code cadence, refresh rotation) + multi-dialect loopback fake
This commit is contained in:
324
tests/e2e/core/providers/test_catalog_oauth.py
Normal file
324
tests/e2e/core/providers/test_catalog_oauth.py
Normal 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]}")
|
||||
|
||||
363
tests/fakes/providers/catalog_fake.py
Normal file
363
tests/fakes/providers/catalog_fake.py
Normal 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
|
||||
278
tests/fakes/providers/catalog_oauth.py
Normal file
278
tests/fakes/providers/catalog_oauth.py
Normal 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))
|
||||
Reference in New Issue
Block a user