net-kingdom/sso-mfa/k8s/authelia/tests/probe_claims.py

215 lines
8.1 KiB
Python
Raw Normal View History

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