"""Integration regression for NK-WP-0041; needs requests, PyYAML, cryptography.""" import argparse import base64 import json import pathlib import socket import subprocess import tempfile import time import urllib.parse import requests import yaml from cryptography.hazmat.primitives import hashes, serialization from cryptography.hazmat.primitives.asymmetric import padding, rsa def decode(value): return base64.urlsafe_b64decode(value + "=" * (-len(value) % 4)) parser = argparse.ArgumentParser( description="Exercise real Authelia ID tokens with and without the KeyCape claims policy; scratch data only." ) parser.add_argument("--authelia-bin", required=True) binary = str(pathlib.Path(parser.parse_args().authelia_bin).resolve()) key = rsa.generate_private_key(public_exponent=65537, key_size=2048) pem = key.private_bytes( serialization.Encoding.PEM, serialization.PrivateFormat.PKCS8, serialization.NoEncryption(), ).decode() source = yaml.safe_load( yaml.safe_load( (pathlib.Path(__file__).resolve().parents[1] / "configmap.yaml").read_text() )["data"]["configuration.yml"] )["identity_providers"]["oidc"] for enabled in (False, True): with tempfile.TemporaryDirectory(prefix="nk-claims-") as tmp: p = pathlib.Path(tmp) with socket.socket() as s: s.bind(("127.0.0.1", 0)) port = s.getsockname()[1] password = "scratch-password-only" digest = ( subprocess.check_output( [ binary, "crypto", "hash", "generate", "argon2", "--password", password, ], text=True, ) .strip() .split("Digest: ")[1] ) (p / "users.yml").write_text( yaml.safe_dump( { "users": { "nk-probe": { "displayname": "Scratch User", "password": digest, "email": "nk-probe+x@example.com", "groups": [], } } } ) ) client = dict(source["clients"][0]) client.update( secret="scratch-client-secret-only", redirect_uris=["https://client.example.com/callback"], ) if not enabled: client.pop("claims_policy") oidc = { "hmac_secret": "h" * 64, "jwks": [{"key": pem, "algorithm": "RS256", "use": "sig"}], "clients": [client], } if enabled: oidc["claims_policies"] = source["claims_policies"] cfg = { "ntp": {"disable_startup_check": True}, "server": {"address": f"tcp://127.0.0.1:{port}/"}, "log": {"level": "info"}, "authentication_backend": {"file": {"path": str(p / "users.yml")}}, "session": { "secret": "s" * 64, "cookies": [ { "domain": "example.com", "authelia_url": "https://auth.example.com", } ], }, "storage": { "encryption_key": "e" * 64, "local": {"path": str(p / "db.sqlite3")}, }, "notifier": {"filesystem": {"filename": str(p / "notifications")}}, "access_control": {"default_policy": "one_factor"}, "identity_validation": {"reset_password": {"jwt_secret": "j" * 64}}, "identity_providers": {"oidc": oidc}, } (p / "config.yml").write_text(yaml.safe_dump(cfg)) with (p / "log").open("w") as log: proc = subprocess.Popen( [binary, "--config", str(p / "config.yml")], stdout=log, stderr=log ) try: session = requests.Session() session.trust_env = False base = f"http://127.0.0.1:{port}" session.headers.update( { "Host": "auth.example.com", "X-Forwarded-Proto": "https", "X-Forwarded-Host": "auth.example.com", } ) for _ in range(100): if proc.poll() is not None: raise RuntimeError((p / "log").read_text()) try: if ( session.get(base + "/api/health", timeout=1).status_code == 200 ): break except requests.ConnectionError: pass time.sleep(0.1) else: raise RuntimeError("Scratch Authelia did not become healthy") r = session.post( base + "/api/firstfactor", json={ "username": "nk-probe", "password": password, "keepMeLoggedIn": False, }, timeout=5, ) assert r.status_code == 200, (r.status_code, r.text) session.headers["Cookie"] = "; ".join( c.name + "=" + c.value for c in session.cookies ) r = session.get( base + "/api/oidc/authorization", params={ "client_id": "keycape", "redirect_uri": "https://client.example.com/callback", "response_type": "code", "scope": "openid profile email groups", "state": "scratch-state-long-enough", "nonce": "scratch-nonce-long-enough", }, allow_redirects=False, timeout=5, ) loc = r.headers.get("Location", "") query = urllib.parse.parse_qs(urllib.parse.urlparse(loc).query) assert "code" in query, ( r.status_code, loc, r.text, (p / "log").read_text(), ) r = session.post( base + "/api/oidc/token", auth=("keycape", "scratch-client-secret-only"), data={ "grant_type": "authorization_code", "code": query["code"][0], "redirect_uri": "https://client.example.com/callback", }, timeout=5, ) assert r.status_code == 200, (r.status_code, r.text) parts = r.json()["id_token"].split(".") key.public_key().verify( decode(parts[2]), (".".join(parts[:2])).encode(), padding.PKCS1v15(), hashes.SHA256(), ) claims = json.loads(decode(parts[1])) assert claims["iss"] == "https://auth.example.com" assert "keycape" in claims["aud"] assert claims["nonce"] == "scratch-nonce-long-enough" assert claims["exp"] > time.time() assert (claims.get("preferred_username") == "nk-probe") == enabled assert claims["sub"] != "nk-probe" print( json.dumps( { "claims_policy": enabled, "signed_id_token_verified": True, "preferred_username_present": "preferred_username" in claims, "subject_is_directory_username": False, } ) ) finally: proc.terminate() proc.wait(timeout=10)