test: PKCE plugin helper against a local fake IdP
Two invariants: (1) the full add/status/refresh/logout lifecycle through the real credential pool against an in-process IdP on 127.0.0.1 that checks the S256 verifier and revokes refresh tokens on rotation, including peer-rotation adoption and a forged-state callback that never reaches the token endpoint; (2) a token_url off the authorize host allowlist is refused before any HTTP request on both login and refresh, and a non-https endpoint is rejected. (cherry picked from commit 0dcf2c7c22ceae2eb847dbda15fb2c6a65ae5c3d)
This commit is contained in:
99
tests/providers/fake_pkce_idp.py
Normal file
99
tests/providers/fake_pkce_idp.py
Normal file
@@ -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)))
|
||||
'''
|
||||
107
tests/providers/test_oauth_pkce_plugin.py
Normal file
107
tests/providers/test_oauth_pkce_plugin.py
Normal file
@@ -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"))
|
||||
Reference in New Issue
Block a user