diff --git a/tests/providers/fake_pkce_idp.py b/tests/providers/fake_pkce_idp.py new file mode 100644 index 0000000000..fbb1aca4a6 --- /dev/null +++ b/tests/providers/fake_pkce_idp.py @@ -0,0 +1,99 @@ +"""Local fake OAuth IdP on 127.0.0.1 for the PKCE plugin helper (tests + live evidence). + +/authorize → 302 to redirect_uri?code=…&state=… (records code_challenge) +/token → grant_type=authorization_code: verifies S256(code_verifier) == challenge, issues tokens + grant_type=refresh_token: rotates (old refresh token becomes invalid → 400) +""" +from __future__ import annotations + +import base64 +import hashlib +import json +import secrets +import threading +from http.server import BaseHTTPRequestHandler, HTTPServer +from urllib.parse import parse_qs, urlencode, urlparse + + +class FakeIdP: + def __init__(self) -> None: + self.codes: dict[str, str] = {} # code -> code_challenge + self.refresh_tokens: set[str] = set() + self.token_requests: list[dict] = [] + self.issued = 0 + idp = self + + class Handler(BaseHTTPRequestHandler): + def log_message(self, *a): # noqa: A003 + return + + def do_GET(self): # noqa: N802 + parsed = urlparse(self.path) + if parsed.path != "/authorize": + self.send_response(404); self.end_headers(); return + q = parse_qs(parsed.query) + assert q["code_challenge_method"] == ["S256"], q + code = secrets.token_urlsafe(16) + idp.codes[code] = q["code_challenge"][0] + state = idp.override_state if idp.override_state is not None else q["state"][0] + self.send_response(302) + self.send_header("Location", f"{q['redirect_uri'][0]}?{urlencode({'code': code, 'state': state})}") + self.end_headers() + + def do_POST(self): # noqa: N802 + if urlparse(self.path).path != "/token": + self.send_response(404); self.end_headers(); return + body = parse_qs(self.rfile.read(int(self.headers.get("Content-Length", 0))).decode()) + form = {k: v[0] for k, v in body.items()} + idp.token_requests.append(form) + grant = form.get("grant_type") + if grant == "authorization_code": + challenge = idp.codes.pop(form.get("code", ""), None) + digest = hashlib.sha256(form.get("code_verifier", "").encode()).digest() + if challenge is None or base64.urlsafe_b64encode(digest).decode().rstrip("=") != challenge: + return self._json(400, {"error": "invalid_grant"}) + elif grant == "refresh_token": + if form.get("refresh_token") not in idp.refresh_tokens: + return self._json(400, {"error": "invalid_grant"}) + idp.refresh_tokens.discard(form["refresh_token"]) + else: + return self._json(400, {"error": "unsupported_grant_type"}) + idp.issued += 1 + refresh = f"rt-{idp.issued}-{secrets.token_hex(4)}" + idp.refresh_tokens.add(refresh) + self._json(200, {"access_token": f"at-{idp.issued}-{secrets.token_hex(4)}", + "refresh_token": refresh, "expires_in": 3600, "token_type": "Bearer"}) + + def _json(self, status, payload): + 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) + + self.override_state: str | None = None + self.server = HTTPServer(("127.0.0.1", 0), Handler) + self.base = f"http://127.0.0.1:{self.server.server_address[1]}" + self._thread = threading.Thread(target=self.server.serve_forever, kwargs={"poll_interval": 0.05}, daemon=True) + + def start(self) -> "FakeIdP": + self._thread.start() + return self + + def stop(self) -> None: + self.server.shutdown() + self.server.server_close() + + +PLUGIN_TEMPLATE = ''' +from providers import register_provider +from providers.base import ProviderProfile +from hermes_cli.auth_oauth_pkce_plugin import OAuthPKCEConfig, pkce_auth_handler, pkce_refresh_credential + +cfg = OAuthPKCEConfig(client_id="hermes-example", authorize_url="{base}/authorize", token_url="{base}/token", + scopes=("inference",), timeout_seconds=20) +register_provider(ProviderProfile(name="example-pkce", auth_type="oauth_external", + base_url="https://example.invalid/v1", fallback_models=("example-model",), + auth_handler=pkce_auth_handler(cfg), refresh_credential=pkce_refresh_credential(cfg))) +''' diff --git a/tests/providers/test_oauth_pkce_plugin.py b/tests/providers/test_oauth_pkce_plugin.py new file mode 100644 index 0000000000..89df7aa519 --- /dev/null +++ b/tests/providers/test_oauth_pkce_plugin.py @@ -0,0 +1,107 @@ +"""Invariants for the declarative OAuth PKCE helper an out-of-tree provider plugin plugs into +``ProviderProfile.auth_handler`` / ``refresh_credential``.""" + +from __future__ import annotations + +import json +from dataclasses import replace +from types import SimpleNamespace +from unittest.mock import Mock +from urllib.request import urlopen + +import pytest + +import hermes_cli.auth +import providers +from hermes_cli import auth_oauth_pkce_plugin as pkce +from hermes_cli.auth_constants import AuthError +from providers import register_provider +from providers.base import ProviderProfile +from tests.providers.fake_pkce_idp import FakeIdP + +PROVIDER = "example-pkce" + + +@pytest.fixture +def idp(monkeypatch, tmp_path): + monkeypatch.setenv("HERMES_HOME", str(tmp_path / "hermes")) + server = FakeIdP().start() + yield server + server.stop() + + +def _config(idp: FakeIdP, **overrides) -> pkce.OAuthPKCEConfig: + base = pkce.OAuthPKCEConfig(client_id="hermes-example", authorize_url=f"{idp.base}/authorize", + token_url=f"{idp.base}/token", scopes=("inference",), timeout_seconds=5) + return replace(base, **overrides) + + +def _browser_hits(url: str) -> bool: + with urlopen(url, timeout=5) as response: # the fake IdP 302s straight to the loopback callback + assert response.status == 200 + return True + + +def test_pkce_handler_add_status_refresh_logout_against_fake_idp(idp, monkeypatch, tmp_path, request): + from agent.credential_pool import load_pool + + monkeypatch.setattr(pkce.webbrowser, "open", _browser_hits) + monkeypatch.setattr("hermes_cli.auth_device_flow._can_open_graphical_browser", lambda: True) + cfg = _config(idp) + handler, refresh = pkce.pkce_auth_handler(cfg), pkce.pkce_refresh_credential(cfg) + # Registered like a real plugin so the credential pool finds ``refresh_credential`` through the seam. + register_provider(ProviderProfile(name=PROVIDER, auth_type="oauth_external", base_url="https://example.invalid/v1", + auth_handler=handler, refresh_credential=refresh)) + request.addfinalizer(lambda: (providers._REGISTRY.pop(PROVIDER, None), + hermes_cli.auth.PROVIDER_REGISTRY.pop(PROVIDER, None))) + args = SimpleNamespace(provider=PROVIDER, no_browser=False) + + assert handler("add", args) is True + exchange = idp.token_requests[-1] + assert exchange["grant_type"] == "authorization_code" and exchange["code_verifier"] + assert exchange["redirect_uri"].startswith("http://127.0.0.1:") + rows = json.loads((tmp_path / "hermes" / "auth.json").read_text())["credential_pool"][PROVIDER] + assert [(r["auth_type"], r["source"], r["oauth_pkce"]["client_id"]) for r in rows] == [ + ("oauth", pkce.POOL_SOURCE, "hermes-example")] + assert rows[0]["refresh_token"] in idp.refresh_tokens and rows[0]["expires_at_ms"] + assert handler("status", args) is True + assert handler("refresh", args) is False # declined → the pool's generic refresh owns it + + # Pool-driven rotation through the real refresh path; the old single-use refresh token is gone. + before = rows[0]["refresh_token"] + pool = load_pool(PROVIDER) + rotated = pool.try_refresh_matching(credential_id=rows[0]["id"]) + assert rotated is not None and rotated.refresh_token != before and rotated.access_token != rows[0]["access_token"] + assert idp.token_requests[-1] == {"grant_type": "refresh_token", "refresh_token": before, "scope": "inference", + "client_id": "hermes-example"} + assert before not in idp.refresh_tokens + # A peer that already rotated on disk is adopted without spending the refresh token again. + stale = load_pool(PROVIDER).entries()[0] + stale.refresh_token = before + assert refresh(stale)["refresh_token"] == rotated.refresh_token + assert idp.token_requests[-1]["refresh_token"] == before # no new POST + + assert handler("logout", args) is True + assert load_pool(PROVIDER).entries() == [] + + # Negative: a callback whose state is not ours is rejected and never reaches the token endpoint. + idp.override_state = "forged" + posts = len(idp.token_requests) + with pytest.raises(AuthError) as excinfo: + handler("add", args) + assert excinfo.value.code == "oauth_state_mismatch" and len(idp.token_requests) == posts + + +def test_token_url_off_authorize_allowlist_refused_before_any_request(idp, monkeypatch): + from hermes_cli.auth_constants import httpx + + monkeypatch.setattr(httpx, "post", Mock(side_effect=AssertionError("token endpoint must not be contacted"))) + cfg = _config(idp, token_url="https://token.attacker.example/token") + with pytest.raises(AuthError) as login_err: + pkce.login(PROVIDER, cfg, open_browser=False) + entry = SimpleNamespace(provider=PROVIDER, id="abc123", refresh_token="rt", access_token="at", expires_at_ms=None) + with pytest.raises(AuthError) as refresh_err: + pkce.pkce_refresh_credential(cfg)(entry) + assert {login_err.value.code, refresh_err.value.code} == {"oauth_token_host_rejected"} + with pytest.raises(AuthError, match="HTTPS"): + pkce.validate_config(PROVIDER, _config(idp, authorize_url="http://auth.example.com/authorize"))