#!/usr/bin/env python3
"""Verify a TrustEvidence evidence package on your own computer, without trusting TrustEvidence.

    python3 verify.py TrustEvidence_EvidencePackage_<...>.zip
    python3 verify.py <package.zip> --published-keys https://<trustevidence host>/.well-known/trustevidence-keys.json

Checks that every file is exactly as exported (manifest), that each evidence record's content and metadata are
unchanged since collection (content hash, hash chain), that each record is part of a sealed batch (Merkle proof),
that the batches form an unbroken chain, and that the batch statements and the manifest carry valid signatures.

Needs Python 3.9 or later. Checking signatures also needs the `cryptography` package (pip install cryptography);
without it every other check still runs and signatures are reported as not checked.

Exit status: 0 verified, 1 verified with limits (see the output), 2 failed or unreadable.
"""
from __future__ import annotations

import argparse
import base64
import hashlib
import io
import json
import sys
import urllib.request
import zipfile

PROOFS_FORMAT = "trustevidence.evidence-proofs.v1"
STATEMENT_TYPE = "trustevidence.evidence-batch.v1"
UNLISTED_OK = {"manifest.json", "SHA256SUMS.txt", "manifest.sig.json"}

try:  # optional: signatures are reported as not checked without it
    from cryptography.exceptions import InvalidSignature
    from cryptography.hazmat.primitives import hashes
    from cryptography.hazmat.primitives.asymmetric import ec, ed25519
    from cryptography.hazmat.primitives.asymmetric.utils import encode_dss_signature
    HAVE_CRYPTO = True
except ImportError:  # pragma: no cover - exercised by hand
    HAVE_CRYPTO = False


def _sha256(data: bytes) -> bytes:
    return hashlib.sha256(data).digest()


def _b64(value: str) -> bytes:
    value = value.strip()
    return base64.urlsafe_b64decode(value.replace("+", "-").replace("/", "_") + "=" * (-len(value) % 4))


# ------------------------------------------------------------------------------------------------ Merkle ----------
def _node(left: bytes, right: bytes) -> bytes:
    return _sha256(b"\x01" + left + right)


def verify_inclusion(leaf: bytes, index: int, size: int, path: list[bytes], root: bytes) -> bool:
    """RFC 9162 section 2.1.3.2."""
    if not 0 <= index < size:
        return False
    fn, sn, r = index, size - 1, leaf
    for p in path:
        if sn == 0:
            return False
        if fn & 1 or fn == sn:
            r = _node(p, r)
            if not fn & 1:
                while not fn & 1 and fn != 0:
                    fn >>= 1
                    sn >>= 1
        else:
            r = _node(r, p)
        fn >>= 1
        sn >>= 1
    return sn == 0 and r == root


# ------------------------------------------------------------------------------------------------ keys ------------
def key_material(key: dict) -> tuple[bytes, str] | None:
    """(raw public key, SHA-256 fingerprint) of a key entry, or None when it carries no usable key."""
    try:
        if key.get("alg") == "Ed25519" and key.get("public_key"):
            raw = _b64(key["public_key"])
            return (raw, hashlib.sha256(raw).hexdigest()) if len(raw) == 32 else None
        if key.get("alg") == "ES256" and key.get("jwk"):
            x, y = _b64(key["jwk"]["x"]), _b64(key["jwk"]["y"])
            raw = b"\x04" + x + y
            return (raw, hashlib.sha256(raw).hexdigest()) if len(x) == len(y) == 32 else None
    except (ValueError, KeyError, TypeError):
        return None
    return None


def check_signature(key: dict | None, message: bytes, signature_b64: str) -> str:
    """'valid', 'invalid', 'no key' or 'not checked'."""
    material = key_material(key or {})
    if material is None:
        return "no key"
    if not HAVE_CRYPTO:
        return "not checked"
    raw, _ = material
    try:
        signature = base64.b64decode(signature_b64)
        if key["alg"] == "Ed25519":
            ed25519.Ed25519PublicKey.from_public_bytes(raw).verify(signature, message)
            return "valid"
        if len(signature) != 64:
            return "invalid"
        public = ec.EllipticCurvePublicNumbers(int.from_bytes(raw[1:33], "big"), int.from_bytes(raw[33:], "big"),
                                               ec.SECP256R1()).public_key()
        der = encode_dss_signature(int.from_bytes(signature[:32], "big"), int.from_bytes(signature[32:], "big"))
        public.verify(der, message, ec.ECDSA(hashes.SHA256()))
        return "valid"
    except (InvalidSignature, ValueError, TypeError):
        return "invalid"


# ------------------------------------------------------------------------------------------------ verification ----
def verify_package(data: bytes, published: list[dict] | None = None) -> dict:
    """Verify a package (the ZIP's bytes). `published`: the keys TrustEvidence publishes, to compare fingerprints."""
    problems: list[str] = []
    limits: list[str] = []
    try:
        with zipfile.ZipFile(io.BytesIO(data)) as z:
            files = {n: z.read(n) for n in z.namelist() if not n.endswith("/")}
    except zipfile.BadZipFile:
        return {"verdict": "failed", "problems": ["not a ZIP file"], "limits": [], "notes": []}

    def load(name: str):
        try:
            return json.loads(files[name])
        except KeyError:
            problems.append(f"{name} is missing")
        except ValueError:
            problems.append(f"{name} is not valid JSON")
        return None

    # 1. Every file is exactly as exported, and nothing was added.
    manifest = load("manifest.json") or {}
    listed = manifest.get("files") or {}
    for name, meta in sorted(listed.items()):
        if name not in files:
            problems.append(f"{name}: listed in the manifest but missing")
        elif hashlib.sha256(files[name]).hexdigest() != meta.get("sha256"):
            problems.append(f"{name}: changed since export (SHA-256 differs from the manifest)")
    for name in sorted(set(files) - set(listed) - UNLISTED_OK):
        problems.append(f"{name}: not part of the exported package")

    proofs = load("proofs.json") or {}
    if proofs and proofs.get("format") != PROOFS_FORMAT:
        problems.append(f"proofs.json has unknown format {proofs.get('format')!r}")
    tenant = proofs.get("tenant_id")

    # 2. Keys: recompute fingerprints and compare them with the published ones when given.
    keys: dict[str, dict] = {}
    key_report = []
    published_fps = None if published is None else {k.get("fingerprint_sha256") for k in published}
    for key in proofs.get("keys") or []:
        keys[key.get("key_id")] = key
        material = key_material(key)
        fingerprint = material[1] if material else None
        if material and key.get("alg") == "Ed25519" and key.get("key_id") != "ed25519:" + fingerprint[:16]:
            problems.append(f"key {key.get('key_id')}: key id does not match the public key")
        is_published = None if published_fps is None or fingerprint is None else fingerprint in published_fps
        key_report.append({"key_id": key.get("key_id"), "alg": key.get("alg"), "fingerprint_sha256": fingerprint,
                           "published": is_published})

    # 3. Manifest signature.
    sigdoc = files.get("manifest.sig.json")
    if sigdoc is None:
        manifest_signature = "unsigned"
        limits.append("the package manifest is not signed, so it does not prove the package came from TrustEvidence")
    else:
        try:
            sig = json.loads(sigdoc)
            manifest_signature = check_signature(keys.get(sig.get("key_id")), files.get("manifest.json", b""), sig["signature"])
        except (ValueError, KeyError, TypeError):
            manifest_signature = "invalid"
        if manifest_signature == "invalid":
            problems.append("manifest signature is invalid")
        elif manifest_signature != "valid":
            limits.append(f"manifest signature {manifest_signature}")

    # 4. Batch statements: format, tenant, signatures, and the chain of roots.
    statements: dict[int, dict] = {}
    sig_counts = {"valid": 0, "invalid": 0, "unsigned": 0, "no key": 0, "not checked": 0}
    for b in proofs.get("batches") or []:
        index = b.get("batch_index")
        try:
            st = json.loads(b["statement"])
        except (KeyError, TypeError, ValueError):
            problems.append(f"batch {index}: statement unreadable")
            continue
        if st.get("type") != STATEMENT_TYPE or st.get("batch_index") != index or st.get("tenant_id") != tenant:
            problems.append(f"batch {index}: statement does not match the package (type, index or organisation)")
        statements[index] = st
        state = check_signature(keys.get(b.get("signing_key_id")), b["statement"].encode("utf-8"), b["signature"]) \
            if b.get("signature") else "unsigned"
        sig_counts[state] += 1
        if state == "invalid":
            problems.append(f"batch {index}: signature is invalid")
    for index in sorted(statements):
        if index - 1 in statements and statements[index].get("prev_root_sha256") != statements[index - 1].get("root_sha256"):
            problems.append(f"batch {index}: does not link to batch {index - 1} (a batch was altered, removed or reordered)")
    if sig_counts["unsigned"]:
        limits.append(f"{sig_counts['unsigned']} batch(es) are sealed but not signed")
    if sig_counts["no key"] or sig_counts["not checked"]:
        limits.append(f"{sig_counts['no key'] + sig_counts['not checked']} batch signature(s) could not be checked"
                      + ("" if HAVE_CRYPTO else " (install the 'cryptography' package)"))

    # 5. Records: content, chain, Merkle inclusion.
    records = sorted(proofs.get("records") or [], key=lambda e: e.get("sequence") or 0)
    unsealed = proven = 0
    previous = None
    for e in records:
        label = f"record #{e.get('sequence')}"
        problems_before = len(problems)
        raw, content = files.get(f"{e.get('path')}/record.json"), files.get(f"{e.get('path')}/content.json")
        if raw is None or content is None:
            problems.append(f"{label}: record.json or content.json is missing")
            continue
        try:
            rec = json.loads(raw)
        except ValueError:
            problems.append(f"{label}: record.json is not valid JSON")
            continue
        if rec.get("id") != e.get("evidence_id") or rec.get("sequence") != e.get("sequence") or rec.get("tenant_id") != tenant:
            problems.append(f"{label}: record does not match its proof entry")
        if hashlib.sha256(content).hexdigest() != rec.get("content_sha256"):
            problems.append(f"{label}: content changed (SHA-256 differs from content_sha256)")
        link = f"{rec.get('prev_chain_hash') or 'GENESIS'}|{rec.get('content_sha256')}|{rec.get('id')}|{rec.get('collected_at')}"
        if hashlib.sha256(link.encode("utf-8")).hexdigest() != rec.get("chain_hash"):
            problems.append(f"{label}: chain hash does not match the record")
        if previous and previous.get("sequence") == rec.get("sequence", 0) - 1 and rec.get("prev_chain_hash") != previous.get("chain_hash"):
            problems.append(f"{label}: does not link to record #{previous.get('sequence')}")
        previous = rec
        if e.get("batch_index") is None:
            unsealed += 1
            continue
        st = statements.get(e["batch_index"])
        if st is None:
            problems.append(f"{label}: its batch {e['batch_index']} is not in the package")
            continue
        first, size = st.get("first_sequence"), st.get("leaf_count")
        leaf = _sha256(b"\x00" + _sha256(raw))
        try:
            path = [bytes.fromhex(p) for p in e.get("audit_path") or []]
            ok = (e.get("leaf_index") == rec.get("sequence", -1) - first
                  and verify_inclusion(leaf, e["leaf_index"], size, path, bytes.fromhex(st.get("root_sha256", ""))))
        except (TypeError, ValueError):
            ok = False
        if not ok:
            problems.append(f"{label}: not proven to be in batch {e['batch_index']} (record altered or proof invalid)")
        elif len(problems) == problems_before:
            proven += 1  # every check on this record passed
    if unsealed:
        limits.append(f"{unsealed} record(s) were not sealed yet when exported; export again in a few minutes for full proofs")

    if any(k["published"] is False for k in key_report):
        limits.append("signed with a key TrustEvidence does not publish")
    notes = []
    if published is None and key_report:
        notes.append("compare the key fingerprints below with the ones TrustEvidence publishes "
                     "(https://<trustevidence host>/.well-known/trustevidence-keys.json)")
    return {
        "verdict": "failed" if problems else "verified with limits" if limits else "verified",
        "problems": problems, "limits": limits, "notes": notes,
        "records": {"total": len(records), "proven": proven, "unsealed": unsealed},
        "batches": {"total": len(statements), **{k.replace(" ", "_"): v for k, v in sig_counts.items()}},
        "manifest": {"files": len(listed), "signature": manifest_signature},
        "keys": key_report,
        "organisation": manifest.get("organization"), "period": manifest.get("period"),
        "generated_at": manifest.get("generated_at"),
    }


def _published(source: str) -> list[dict]:
    if source.startswith("https://"):
        with urllib.request.urlopen(source, timeout=20) as resp:  # noqa: S310 - https only
            return json.load(resp).get("keys", [])
    with open(source, encoding="utf-8") as fh:
        return json.load(fh).get("keys", [])


def render(r: dict) -> str:
    lines = [f"TrustEvidence evidence package: {r['verdict'].upper()}"]
    if r.get("organisation"):
        lines.append(f"Organisation: {r['organisation']}   period: {(r.get('period') or {}).get('start')} to "
                     f"{(r.get('period') or {}).get('end')}   exported: {r.get('generated_at')}")
    if "records" in r:
        lines.append(f"Records: {r['records']['total']} ({r['records']['proven']} fully verified, "
                     f"{r['records']['unsealed']} not yet sealed)")
        b = r["batches"]
        lines.append(f"Batches: {b['total']} (signatures valid {b['valid']}, invalid {b['invalid']}, unsigned {b['unsigned']}, "
                     f"unchecked {b['no_key'] + b['not_checked']})")
        lines.append(f"Files: {r['manifest']['files']} in the manifest; manifest signature: {r['manifest']['signature']}")
    for k in r.get("keys", []):
        state = {True: "published by TrustEvidence", False: "NOT published by TrustEvidence", None: "compare with the published key"}[k["published"]]
        lines.append(f"Key {k['alg']} {k['fingerprint_sha256']}: {state}")
    for title, items in (("Problems", r["problems"]), ("Limits", r["limits"]), ("Notes", r["notes"])):
        if items:
            lines.append(f"{title}:")
            lines += [f"  - {item}" for item in items[:200]]
            if len(items) > 200:
                lines.append(f"  ... and {len(items) - 200} more")
    return "\n".join(lines)


def main(argv: list[str] | None = None) -> int:
    parser = argparse.ArgumentParser(description="Verify a TrustEvidence evidence package.")
    parser.add_argument("package", help="the evidence package ZIP")
    parser.add_argument("--published-keys", help="URL (https) or file of TrustEvidence's published keys")
    parser.add_argument("--json", action="store_true", help="print the result as JSON")
    args = parser.parse_args(argv)
    try:
        with open(args.package, "rb") as fh:
            data = fh.read()
        published = _published(args.published_keys) if args.published_keys else None
    except (OSError, ValueError) as exc:
        print(f"Cannot read input: {exc}", file=sys.stderr)
        return 2
    result = verify_package(data, published)
    print(json.dumps(result, indent=2) if args.json else render(result))
    return {"verified": 0, "verified with limits": 1}.get(result["verdict"], 2)


if __name__ == "__main__":
    sys.exit(main())
