#!/usr/bin/env python3
"""Offline validator for Domain Name Proof proofs.

    pip install cryptography
    python3 verify_proof.py proof.json keys.json [expected_domain] [expected_audience] [expected_nonce]

keys.json is https://domainnameproof.online/.well-known/domain-name-proof-key.json (cache or pin it).
No call to the issuer is needed to validate a proof.
"""
import base64
import hashlib
import json
import sys
from datetime import datetime, timezone

from cryptography.exceptions import InvalidSignature
from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PublicKey

ISSUER = "domainnameproof.online"


def canonicalize(value):
    """RFC 8785 JCS for proofs: strings, integers, booleans, null, lists, objects with ASCII keys.
    Proofs never contain floats; for this subset json.dumps with sorted keys matches JCS."""
    def check(v):
        if isinstance(v, float):
            raise ValueError("floats are not allowed in proofs")
        if isinstance(v, dict):
            for k, x in v.items():
                check(x)
        elif isinstance(v, list):
            for x in v:
                check(x)
    check(value)
    return json.dumps(value, sort_keys=True, separators=(",", ":"), ensure_ascii=False).encode("utf-8")


def b64u_decode(s):
    return base64.urlsafe_b64decode(s + "=" * (-len(s) % 4))


def thumbprint(x):
    return base64.urlsafe_b64encode(
        hashlib.sha256(canonicalize({"crv": "Ed25519", "kty": "OKP", "x": x})).digest()
    ).rstrip(b"=").decode()


def parse_time(s):
    return datetime.strptime(s, "%Y-%m-%dT%H:%M:%SZ").replace(tzinfo=timezone.utc)


def validate_proof(proof, keys, expected_domain=None, expected_audience=None, expected_nonce=None, now=None):
    now = now or datetime.now(timezone.utc)
    errors = []
    if proof.get("proof_version") != "1":
        return {"valid": False, "errors": ["proof_version_unsupported"]}
    if proof.get("issuer") != ISSUER:
        errors.append("issuer_mismatch")
    key = next((k for k in keys if k.get("kid") == proof.get("key_id") and k.get("crv") == "Ed25519"), None)
    signature_valid = False
    if key is None or thumbprint(key["x"]) != key["kid"]:
        errors.append("unknown_key")
    else:
        unsigned = {k: v for k, v in proof.items() if k != "signature"}
        try:
            sig = b64u_decode(proof["signature"])
            if len(sig) != 64:
                raise InvalidSignature()
            Ed25519PublicKey.from_public_bytes(b64u_decode(key["x"])).verify(sig, canonicalize(unsigned))
            signature_valid = True
        except (InvalidSignature, ValueError, KeyError):
            errors.append("signature_invalid")
    verified_at, valid_until = parse_time(proof["verified_at"]), parse_time(proof["valid_until"])
    if not verified_at < valid_until or (valid_until - verified_at).total_seconds() > 30 * 86400:
        errors.append("timestamps_inconsistent")
    if now >= valid_until:
        errors.append("proof_expired")
    if expected_domain is not None:
        try:  # IDNA2008/UTS 46 non-transitional, same as the issuer (pip install idna)
            import idna
            want = idna.encode(expected_domain.rstrip("."), uts46=True).decode("ascii").lower()
        except ImportError:  # stdlib codec is IDNA2003 (maps ß to ss); pass punycode to be exact
            want = expected_domain.rstrip(".").encode("idna").decode("ascii").lower()
        if want != proof["domain"]:
            errors.append("domain_mismatch")
    if expected_audience is not None and expected_audience != proof.get("audience"):
        errors.append("audience_mismatch")
    if expected_nonce is not None and expected_nonce != proof.get("nonce"):
        errors.append("nonce_mismatch")
    return {"valid": not errors, "signature_valid": signature_valid, "errors": errors}


if __name__ == "__main__":
    if len(sys.argv) < 3:
        print(__doc__)
        sys.exit(2)
    with open(sys.argv[1], encoding="utf-8") as f:
        proof = json.load(f)
    proof = proof.get("proof", proof)
    with open(sys.argv[2], encoding="utf-8") as f:
        keys = json.load(f)["keys"]
    extra = sys.argv[3:] + [None] * 3
    result = validate_proof(proof, keys, extra[0], extra[1], extra[2])
    print(json.dumps(result, indent=2))
    print("Note: a valid proof shows control at verified_at, not current control, legal ownership or identity.")
    sys.exit(0 if result["valid"] else 1)
