#!/usr/bin/env python3
"""Standalone verifier for a CAIN-42 Evolution 4 DAG cluster evidence bundle.

No CAIN imports. Needs Python 3.8+, `cryptography`, and verify_pbft_qc_bundle.py
(the standalone PBFT certificate verifier) in the same directory.

    python3 verify_dag_bundle.py DAG_CLUSTER_EVIDENCE.json

Checks, from the bundle bytes only:
  * membership: configuration hash recomputed from member ids + Ed25519 keys;
  * every availability certificate in the last anchor's causal proof: the
    vertex is signed by its creator (a member), >= 3 distinct members signed
    the availability attestation for exactly that vertex and payload;
  * causal closure: every parent of every vertex is a certified vertex of the
    previous round inside the proof; round 0 has no parents;
  * one vertex per (creator, round) in the proof (no equivocation);
  * the anchor was decided by PBFT: the committed request carried inside the
    AuthorizationCertificate hashes to the digest of a valid COMMIT/FAST_COMMIT
    quorum certificate, and names exactly the anchor vertex with a valid
    availability certificate;
  * DETERMINISTIC ORDER, recomputed here: for the committed anchors in PBFT
    order, each anchor's certified history minus earlier anchors' history,
    sorted by (round, creator, digest), requests in batch order; the order
    hash of every anchor and the final ledger root must equal what EVERY node
    reported;
  * all four nodes report the same PBFT application-state hash and no DAG
    equivocation evidence;
  * negative controls: tampered attestations, a cut causal proof, and an
    anchor request naming another vertex must all be rejected.
"""
from __future__ import annotations

import copy
import hashlib
import json
import os
import sys
from typing import Any, Dict, List, Tuple

sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
import verify_pbft_qc_bundle as Q  # noqa: E402  (standalone sibling verifier, no CAIN code)

h, canon = Q.h, Q.canon
VERTEX_DOMAIN = "CAIN42/DAG/VERTEX/v1"
ATTEST_DOMAIN = "CAIN42/DAG/AVAILABILITY/v1"
CERT_DOMAIN = "CAIN42/DAG/AVAILABILITY-CERTIFICATE/v1"
CAUSAL_DOMAIN = "CAIN42/DAG/CAUSAL_PROOF/v1"
ORDER_DOMAIN = "CAIN42/DAG/ORDER/v1"


def vertex_body(v: Dict[str, Any]) -> Dict[str, Any]:
    return {"domain": VERTEX_DOMAIN, "cluster_id": v["cluster_id"], "epoch": v["epoch"], "round": v["round"],
            "creator": v["creator"], "parents": sorted(v["parents"]), "payload_digest": v["payload_digest"],
            "batch_size": v["batch_size"], "created_at": round(v["created_at"], 3),
            "protocol_version": v["protocol_version"]}


def vertex_digest(v: Dict[str, Any]) -> str:
    return h(vertex_body(v))


def verify_cert(cd: Dict[str, Any], mb: "Q.Membership") -> Tuple[bool, str]:
    try:
        v = cd["vertex"]
        d = vertex_digest(v)
        if v.get("vertex_digest") not in (None, d):
            return False, "stated vertex_digest does not match content"
        if (v["cluster_id"], v["epoch"]) != (mb.cluster_id, mb.epoch):
            return False, "wrong cluster or epoch"
        if v["creator"] not in mb.keys or not Q.ed25519_ok(mb.keys[v["creator"]], v["signature_b64"], d.encode()):
            return False, "vertex signature does not verify"
        good = set()
        for a, sig in cd["attestations"].items():
            body = {"domain": ATTEST_DOMAIN, "cluster_id": v["cluster_id"], "epoch": v["epoch"], "vertex_digest": d,
                    "creator": v["creator"], "round": v["round"], "payload_digest": v["payload_digest"], "attester": a}
            if a not in mb.keys or not Q.ed25519_ok(mb.keys[a], sig, h(body).encode()):
                return False, f"attestation of {a!r} does not verify"
            good.add(a)
        if len(good) < mb.quorum:
            return False, f"{len(good)} attestations, need {mb.quorum}"
        ch = h({"domain": CERT_DOMAIN, "vertex_digest": d, "attesters": sorted(cd["attestations"]),
                "signatures": [cd["attestations"][a] for a in sorted(cd["attestations"])]})
        if cd.get("certificate_hash") not in (None, ch):
            return False, "certificate_hash does not match content"
        return True, d
    except (KeyError, TypeError, ValueError) as e:
        return False, f"malformed: {e}"


def verify_causal(proof: Dict[str, Any], mb) -> Tuple[bool, str, Dict[str, Dict[str, Any]]]:
    verts: Dict[str, Dict[str, Any]] = {}
    slots = set()
    for cd in proof.get("certificates") or []:
        ok, d = verify_cert(cd, mb)
        if not ok:
            return False, d, {}
        v = cd["vertex"]
        if (v["creator"], v["round"]) in slots:
            return False, f"two vertices for ({v['creator']}, {v['round']})", {}
        slots.add((v["creator"], v["round"]))
        verts[d] = v
    if proof.get("vertex_digest") not in verts:
        return False, "anchor not in proof", {}
    for d, v in verts.items():
        if v["round"] == 0 and v["parents"]:
            return False, "round-0 vertex with parents", {}
        if v["round"] > 0 and len(set(v["parents"])) < mb.quorum:
            return False, "fewer than Q parents", {}
        for p in v["parents"]:
            if p not in verts or verts[p]["round"] != v["round"] - 1:
                return False, f"parent {p[:16]} of {d[:16]} missing or wrong round", {}
    body = {k: proof[k] for k in ("domain", "vertex_digest", "certificates")}
    if h(body) != proof.get("proof_hash"):
        return False, "proof_hash does not match content", {}
    return True, "", verts


def history(anchor: str, verts: Dict[str, Dict[str, Any]]) -> List[str]:
    out, stack, seen = [], [anchor], set()
    while stack:
        d = stack.pop()
        if d in seen:
            continue
        seen.add(d)
        out.append(d)
        stack.extend(verts[d]["parents"])
    return out


def tie_break_key(seed: str, digest: str) -> str:
    return h({"domain": "CAIN42/DAG/TIEBREAK/v1", "seed": seed, "vertex": digest})


def recompute_orders(anchors: List[Tuple[str, str]], verts, batches) -> Tuple[List[str], str]:
    """anchors: [(anchor digest, tie-break seed)] in PBFT order. Within a round
    vertices are ordered by H(seed, digest), the seed being the PBFT decision
    hash before the anchor (Evolution 5); legacy bundles have no seed."""
    ordered, ledger, hashes = set(), {}, []
    for a, seed in anchors:
        hist = [d for d in history(a, verts) if d not in ordered]
        if seed is None:
            hist.sort(key=lambda d: (verts[d]["round"], verts[d]["creator"], d))
        else:
            hist.sort(key=lambda d: (verts[d]["round"], tie_break_key(seed, d)))
        reqs: List[Dict[str, Any]] = []
        for d in hist:
            reqs.extend(batches[verts[d]["payload_digest"]])
        body = {"domain": ORDER_DOMAIN, "anchor": a, "vertices": hist, "requests_digest": h(reqs)}
        if seed is not None:
            body["tie_break_seed"] = seed
        hashes.append(h(body))
        ordered.update(hist)
        for r in reqs:
            op = r.get("operation") or {}
            if op.get("action") in ("write", "set", "update"):
                ledger[str(op.get("resource", "")).strip("/")] = op.get("data")
    return hashes, h({"domain": "CAIN42/DAG/LEDGER/v1", "state": ledger})


def verify_bundle(b: Dict[str, Any]) -> Dict[str, Any]:
    checks = []

    def check(name, ok, detail=""):
        checks.append({"check": name, "result": "PASS" if ok else "FAIL", "detail": detail})

    mb = Q.Membership(b["membership"])
    check("membership: 4 members, f=1, quorum 3, configuration_hash recomputed",
          (mb.n, mb.f, mb.quorum) == (4, 1, 3) and mb.configuration_hash == b["configuration_hash"])
    last = b["last_anchor"]
    ok, why, verts = verify_causal(last["causal_proof"], mb)
    check(f"causal proof of the last anchor: {len(verts)} certified vertices, closure, one per slot", ok, why)
    ac = last["authorization_certificate"]
    ok_ac, why_ac = Q.verify_auth_cert(ac, mb)
    check("anchor committed by PBFT: AuthorizationCertificate + its quorum certificate verify", ok_ac, why_ac)
    op = (ac.get("request") or {}).get("operation") or {}
    cert_in_op = (op.get("data") or {}).get("availability_certificate") or {}
    ok_b, dig = verify_cert(cert_in_op, mb)
    check("the committed request is a dag_anchor for exactly this vertex with a valid availability certificate",
          op.get("action") == "dag_anchor" and op.get("resource") == last["anchor"] and ok_b and dig == last["anchor"])
    ref = next(iter(sorted(b["nodes"])))
    ordered = [o for o in b["nodes"][ref]["dag_order"]["ordered"] if o.get("status") == "ORDERED"]
    anchors = [(o["anchor"], o.get("tie_break_seed")) for o in sorted(ordered, key=lambda o: o["sequence"])]
    missing = [a for a, _ in anchors if a not in verts]
    # The tie-break seed must be the PBFT decision hash of the sequence before each anchor.
    chain = {d["sequence"]: d["decision_hash"] for d in (b["nodes"][ref].get("decision_chain") or [])}
    seeded = [o for o in ordered if o.get("tie_break_seed") is not None]
    if seeded and chain:
        bad = [o["sequence"] for o in seeded if chain.get(o["sequence"] - 1, Q.genesis(mb.cluster_id, mb.epoch)
                                                          if o["sequence"] == 1 else None) != o["tie_break_seed"]]
        check("every tie-break seed is the PBFT decision hash before its anchor (not proposer-chosen)", not bad,
              f"{len(seeded)} anchors")
    check("every committed anchor is inside the last anchor's causal history", not missing, f"{len(anchors)} anchors")
    batches = b.get("batches") or {}
    if not missing and all(v["payload_digest"] in batches for v in verts.values()):
        hashes, root = recompute_orders(anchors, verts, batches)
        for n, nd in sorted(b["nodes"].items()):
            got = [o["order_hash"] for o in sorted((o for o in nd["dag_order"]["ordered"] if o.get("status") == "ORDERED"),
                                                   key=lambda o: o["sequence"])]
            detail = f"{len(got)} anchors"
            if got != hashes and got == hashes[:len(got)]:
                detail = f"LAGGING: ordered {len(got)} of {len(hashes)} anchors (a prefix, not a divergence)"
            elif got != hashes:
                detail = "DIVERGENT order"
            check(f"{n}: every anchor's order hash equals the order recomputed here", got == hashes, detail)
            check(f"{n}: ledger root equals the ledger recomputed here", nd["dag_order"]["ledger_root"] == root, root[:16])
    else:
        check("all batches present to recompute the order", False, "batches missing from bundle")
    apps = {nd["pbft"]["application_state_hash"] for nd in b["nodes"].values()}
    check("all nodes: identical PBFT application state hash", len(apps) == 1)
    eq = {n: nd["dag_status"]["equivocation_evidence"] for n, nd in b["nodes"].items()}
    check("no node holds DAG equivocation evidence (honest run, incl. a crash-restart)", not any(eq.values()), str(eq))
    # negative controls
    t = copy.deepcopy(last["causal_proof"])
    a0 = sorted(t["certificates"][0]["attestations"])[0]
    t["certificates"][0]["attestations"][a0] = t["certificates"][0]["attestations"][sorted(t["certificates"][0]["attestations"])[1]]
    check("negative control: swapped attestation signature -> rejected", not verify_causal(t, mb)[0])
    t = copy.deepcopy(last["causal_proof"])
    t["certificates"] = [c for c in t["certificates"] if c["vertex"]["round"] != 0]
    t["proof_hash"] = h({k: t[k] for k in ("domain", "vertex_digest", "certificates")})
    check("negative control: causal proof with round 0 cut away -> rejected", not verify_causal(t, mb)[0])
    t = copy.deepcopy(ac)
    t["request"]["operation"]["resource"] = "0" * 64
    check("negative control: anchor request naming another vertex -> rejected", not Q.verify_auth_cert(t, mb)[0])
    failed = [c for c in checks if c["result"] == "FAIL"]
    return {"verifier": "verify_dag_bundle.py (no CAIN imports)", "checks": checks, "passed": len(checks) - len(failed),
            "total": len(checks), "verdict": "ALL_CHECKS_PASSED" if not failed else "FAILED"}


def main(argv: List[str]) -> int:
    if len(argv) < 2:
        print(__doc__)
        return 2
    res = verify_bundle(Q.load(argv[1]))
    for c in res["checks"]:
        print(f"[{c['result']}] {c['check']}" + (f"  ({c['detail']})" if c["detail"] else ""))
    print(f"\n{res['verdict']}: {res['passed']}/{res['total']} checks")
    return 0 if res["verdict"] == "ALL_CHECKS_PASSED" else 1


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