#!/usr/bin/env python3
"""CAIN-42 Public Proof Fabric verifier -- CLEAN ROOM: imports nothing from CAIN (stdlib + `cryptography`).

    python3 verify_proof_fabric.py DIR_OR_URL                 verify a proof-fabric bundle (local dir or https URL)
    python3 verify_proof_fabric.py DIR_OR_URL --claims        also verify every signed claim against all 3 sites
    python3 verify_proof_fabric.py --pack CAIN42_EVIDENCE_PACK.zip   verify the downloadable enterprise pack

It does not ask CAIN whether the evidence is valid. It recomputes:
  * MANIFEST.json: Ed25519 signature by the evidence-root key (and that key == the one in the signed claims registry
    with --claims), and the SHA-256 of every listed file
  * BUILD_MANIFEST / ARTIFACT_FILES / DEPLOYMENT_ATTESTATION: the artifact tree hash from the per-file list, the
    deployment id and configuration hash from their inputs, and the attestation's artifact == the build's
  * SBOM: well-formed CycloneDX, component count == the build manifest's
  * TEST_MANIFEST: totals == per-test statuses, and per-suite counts == the published JUnit XML (parsed here)
  * TEST_VECTORS: every vector with its own implementations of canonical JSON, domain digests, Ed25519, Merkle,
    policy precedence, authority intersection, blast radius, governance quorum and lease validity
  * PROVENANCE_GRAPH: every edge joins existing nodes, the artifact node == the build manifest
  * NEGATIVE_EVIDENCE: counts == entries, every fix names a full commit id
Result: machine-readable JSON (PASS / FAIL per check); exit 0 only if everything passes.
"""
from __future__ import annotations

import base64
import hashlib
import io
import json
import sys
import urllib.request
import xml.etree.ElementTree as ET
import zipfile
from pathlib import Path

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

D_PROOF = "CAIN42/PUBLIC-PROOF-FABRIC/v1"
SITES = {"cainstudio.online": "https://cainstudio.online/proof/bundle/",
         "mcpgate.online": "https://mcpgate.online/proof/bundle/", "clawx.click": "https://clawx.click/evidence/"}
POLICY_LEVELS = ("CONSTITUTION", "SYSTEM_SAFETY", "ORGANIZATION", "ENVIRONMENT", "RESOURCE", "AGENT", "TRAJECTORY", "ACTION")
EFFECTS = ("DENY", "REQUIRE_APPROVAL", "ALLOW")
BLAST = ("LOCAL", "SERVICE", "MULTI-SERVICE", "ORGANIZATIONAL", "CROSS-ENVIRONMENT", "SYSTEM-WIDE", "CRITICAL")


# ------------------------------------------------------------------ primitives (independent implementations)
def canon(o) -> bytes:
    return json.dumps(o, sort_keys=True, separators=(",", ":"), ensure_ascii=True).encode()


def H(b: bytes) -> str:
    return hashlib.sha256(b).hexdigest()


def digest(domain, fields) -> str:
    return H(canon({"domain": domain, **fields}))


def ed_ok(pub_b64, sig_b64, msg: bytes) -> bool:
    try:
        ed25519.Ed25519PublicKey.from_public_bytes(base64.b64decode(pub_b64)).verify(base64.b64decode(sig_b64), msg)
        return True
    except (InvalidSignature, ValueError, TypeError):
        return False


def within(child, pat) -> bool:
    return pat == "*" or child == pat or (pat.endswith("*") and child.startswith(pat[:-1]))


def merkle(leaves):
    if not leaves:
        return H(b"")
    layer = [bytes.fromhex(x) for x in leaves]
    while len(layer) > 1:
        nxt = []
        for i in range(0, len(layer), 2):
            nxt.append(hashlib.sha256(b"\x01" + layer[i] + layer[i + 1]).digest() if i + 1 < len(layer) else layer[i])
        layer = nxt
    return layer[0].hex()


def policy(rules, req):
    hits = []
    for r in rules:
        if r["level"] not in POLICY_LEVELS or r["effect"] not in EFFECTS:
            return "DENY"
        if r.get("max_sensitivity") is not None and req["sensitivity"] <= r["max_sensitivity"]:
            continue
        if all(any(within(req[k], p) for p in r[f]) for k, f in (("agent", "agents"), ("capability", "capabilities"),
                                                                 ("resource", "resources"), ("environment", "environments"))):
            hits.append(r["effect"])
    return next((e for e in EFFECTS if e in hits), "DENY")


def authority(layers, cap, res):
    missing = [n for n, L in layers.items() if not (L and any(c == cap and within(res, p) for c, p in L))]
    return not missing, missing


def blast(nodes, edges, target, rev):
    dep = {}
    for a, b in edges:
        dep.setdefault(b, []).append(a)
    seen, frontier, order = {target}, [target], [target]
    while frontier:
        nxt = []
        for n in frontier:
            for d in sorted(dep.get(n, [])):
                if d not in seen:
                    seen.add(d)
                    order.append(d)
                    nxt.append(d)
        frontier = nxt
    svc = [n for n in order if nodes.get(n, {}).get("kind") == "SERVICE"]
    envs = {nodes.get(n, {}).get("environment", "?") for n in order}
    orgs = {nodes.get(n, {}).get("organization", "?") for n in order}
    total = len(nodes)
    c = ("SYSTEM-WIDE" if len(orgs) > 1 or (total > 3 and len(order) >= total) else "CROSS-ENVIRONMENT" if len(envs) > 1
         else "ORGANIZATIONAL" if len(order) >= 10 else "MULTI-SERVICE" if len(svc) >= 2 else "SERVICE" if len(order) > 1
         else "LOCAL")
    if rev in ("IRREVERSIBLE", "UNKNOWN") and BLAST.index(c) >= BLAST.index("MULTI-SERVICE"):
        c = "CRITICAL"
    return c


def quorum(node_keys, votes):
    q = 2 * ((len(node_keys) - 1) // 3) + 1
    by = {}
    for v in votes:
        pub = node_keys.get(v.get("node_id"))
        body = {"node_id": v.get("node_id"), "state_digest": v.get("state_digest")}
        if pub and ed_ok(pub, v.get("signature", ""), digest("CAIN42/CAG-L5-GOVERNANCE-STATE/v1", body).encode()):
            by.setdefault(v["state_digest"], set()).add(v["node_id"])
    if len(by) > 1:
        return False, "CONFLICTING_GOVERNANCE_STATE"
    if not by:
        return False, "NO_VALID_VOTES"
    return (True, None) if len(next(iter(by.values()))) >= q else (False, "BELOW_QUORUM")


def lease_reasons(L, r, max_ttl, bands):
    why = []
    if L["revocation_state"] != "ACTIVE":
        why.append("LEASE_REVOKED")
    if not L["expires_at"] or L["expires_at"] <= L["issued_at"]:
        why.append("LEASE_UNBOUNDED")
    elif L["expires_at"] - L["issued_at"] > max_ttl:
        why.append("LEASE_TTL_EXCEEDS_MAXIMUM")
    if L["expires_at"] and r["now"] > L["expires_at"]:
        why.append("LEASE_EXPIRED")
    if r["now"] < L["issued_at"]:
        why.append("LEASE_NOT_YET_VALID")
    if not L["policy_version"] or not r["policy_version"]:
        why.append("POLICY_UNBOUND")
    elif r["policy_version"] != L["policy_version"]:
        why.append("POLICY_CHANGED")
    if r["trust"] < L["trust_floor"]:
        why.append("TRUST_BELOW_FLOOR")
    if r["risk_band"] not in bands or L["risk_ceiling"] not in bands:
        why.append("RISK_BAND_UNKNOWN")
    elif bands.index(r["risk_band"]) > bands.index(L["risk_ceiling"]):
        why.append("RISK_CEILING_EXCEEDED")
    if not L["context_hash"] or not r["context_hash"]:
        why.append("CONTEXT_UNBOUND")
    elif r["context_hash"] != L["context_hash"]:
        why.append("CONTEXT_CHANGED")
    if not L["action_scope"] or not L["resource_scope"]:
        why.append("LEASE_SCOPE_EMPTY")
    else:
        if r["action"] not in L["action_scope"]:
            why.append("ACTION_SCOPE_EXCEEDED")
        if not any(within(r["resource"], p) for p in L["resource_scope"]):
            why.append("RESOURCE_SCOPE_EXCEEDED")
    return sorted(why)


# ------------------------------------------------------------------ sources
class Source:
    def __init__(self, loc):
        self.loc = loc.rstrip("/") + "/" if loc.startswith("http") else loc
        self.zip = None

    def get(self, name: str) -> bytes:
        if self.zip is not None:
            return self.zip.read(self.prefix + name)
        if self.loc.startswith("http"):
            with urllib.request.urlopen(self.loc + name, timeout=60) as r:
                return r.read()
        return (Path(self.loc) / name).read_bytes()

    def json(self, name):
        return json.loads(self.get(name))


def fetch(url):
    with urllib.request.urlopen(url, timeout=60) as r:
        return r.read()


# ------------------------------------------------------------------ checks
def verify(src: Source, with_claims: bool) -> dict:
    checks = []

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

    man = src.json("MANIFEST.json")
    sig = man.get("signature", {})
    body = {k: v for k, v in man.items() if k != "signature"}
    dg = H(canon({"domain": D_PROOF, **body}))
    check("manifest.digest", dg == sig.get("digest_sha256"))
    check("manifest.signature", ed_ok(sig.get("public_key_b64", ""), sig.get("signature_b64", ""), dg.encode()))
    root_key = sig.get("public_key_b64")
    bad = []
    for f, h in man.get("files", {}).items():
        try:
            if H(src.get(f)) != h:
                bad.append(f)
        except Exception:  # noqa: BLE001
            bad.append(f + " (missing)")
    check("manifest.file_hashes", not bad, f"{len(man.get('files', {}))} files; bad: {bad[:5]}")

    bm, af, att = src.json("BUILD_MANIFEST.json"), src.json("ARTIFACT_FILES.json"), src.json("DEPLOYMENT_ATTESTATION.json")
    tree = H(canon(af["entries"]))
    check("artifact.tree_hash_recomputed", tree == af["tree_sha256"] == bm["artifact"]["tree_sha256"], tree[:16])
    check("artifact.file_count", len(af["entries"]) == bm["artifact"]["files"] == bm["artifact_vs_commit"]["files"])
    d = bm["artifact_vs_commit"]
    check("artifact.commit_accounting", d["identical_to_commit"] + len(d["modified_vs_commit"]) + len(d["untracked"]) == d["files"],
          f"{d['identical_to_commit']} identical, {len(d['modified_vs_commit'])} modified, {len(d['untracked'])} untracked")
    check("attestation.binds_build", att["artifact_hash"] == bm["artifact"]["tree_sha256"] and att["git_commit"] == bm["git_commit"])
    import calendar, time as _t  # noqa: E401
    started = calendar.timegm(_t.strptime(att["deployment_timestamp"], "%Y-%m-%dT%H:%M:%SZ"))
    dep_id = H(canon({"commit": att["git_commit"], "artifact": att["artifact_hash"], "started": started}))[:16]
    check("attestation.deployment_id_recomputed", dep_id == att["deployment_id"], dep_id)
    check("attestation.configuration_hash_recomputed",
          H(canon(att["configuration"]["public_switches"])) == att["configuration"]["configuration_hash"])
    check("attestation.no_hardware_claim", str(att.get("HARDWARE_ATTESTATION", "")).startswith("NOT AVAILABLE"))

    sb = src.json("SBOM.cdx.json")
    check("sbom.cyclonedx", sb.get("bomFormat") == "CycloneDX" and sb.get("specVersion") == "1.5")
    check("sbom.count_matches_build", len(sb["components"]) == bm["dependencies"]["count"], len(sb["components"]))

    tm = src.json("TEST_MANIFEST.json")
    st = {"PASS": "passed", "FAIL": "failed", "ERROR": "error", "SKIPPED": "skipped", "XFAIL": "xfailed"}
    tot = {k: 0 for k in st.values()}
    for t in tm["tests"]:
        tot[st[t["status"]]] += 1
    check("tests.totals_match_entries", all(tot[k] == tm["totals"][k] for k in tot) and len(tm["tests"]) == tm["totals"]["tests"],
          tot)
    for s in tm["suites"]:
        x = src.get(f"junit/{s['suite']}.xml")
        ok_hash = H(x) == s["junit_sha256"]
        root = ET.fromstring(x)
        n = sum(1 for _ in root.iter("testcase"))
        fails = sum(1 for tc in root.iter("testcase") for c in tc if c.tag in ("failure", "error"))
        check(f"tests.{s['suite']}.junit", ok_hash and n == sum(s["counts"].values()) and fails == s["counts"]["failed"] + s["counts"]["error"],
              f"{n} testcases, {fails} failures in the XML")
        check(f"tests.{s['suite']}.nonempty", n > 0, n)

    V = src.json("TEST_VECTORS.json")
    vr = []
    for c in V["canonical_json"]:
        vr.append(canon(c["input"]).decode() == c["canonical_utf8"] and H(canon(c["input"])) == c["sha256"])
    for s in V["signatures"]:
        dgs = digest(s["domain"], s["fields"])
        if "digest_sha256" in s:
            vr.append(dgs == s["digest_sha256"])
        vr.append(ed_ok(s["public_key_b64"], s["signature_b64"], dgs.encode()) == s["expected_valid"])
    for m in V["merkle"]["cases"]:
        vr.append(merkle(m["leaves"]) == m["root"])
    for c in V["policy_precedence"]["cases"]:
        vr.append(policy(c["rules"], c["request"]) == c["expected_effect"])
    for c in V["authority_intersection"]["cases"]:
        ok, missing = authority(c["layers"], c["capability"], c["resource"])
        vr.append(ok == c["allowed"] and sorted(missing) == sorted(c["outside"]))
    for c in V["blast_radius"]["cases"]:
        vr.append(blast(c["nodes"], c["edges"], c["target"], c["reversibility"]) == c["expected_class"])
    for c in V["governance_quorum"]["cases"]:
        ok, reason = quorum(V["governance_quorum"]["node_public_keys_b64"], c["votes"])
        vr.append(ok == c["authoritative"] and (ok or reason == c.get("reason")))
    lv = V["lease_validity"]
    for c in lv["cases"]:
        vr.append(lease_reasons(c["lease"], c["request"], lv["max_ttl_seconds"], lv["risk_bands"]) == c["expected_reasons"])
    check("vectors.all_recomputed", all(vr), f"{sum(vr)}/{len(vr)} vectors match an independent implementation")
    check("vectors.no_private_key", V.get("private_keys_published") is False and "BEGIN" not in json.dumps(V))

    pg = src.json("PROVENANCE_GRAPH.json")
    ids = {n["id"] for n in pg["nodes"]}
    dangling = [e for e in pg["edges"] if e["from"] not in ids or e["to"] not in ids]
    check("provenance.edges_join_nodes", not dangling, dangling[:3])
    arts = [n for n in pg["nodes"] if n["kind"] == "ARTIFACT"]
    check("provenance.artifact_is_build", arts and arts[0]["tree_sha256"] == bm["artifact"]["tree_sha256"])
    kinds = {n["kind"] for n in pg["nodes"]}
    check("provenance.full_chain", {"SOURCE", "COMMIT", "BUILD", "ARTIFACT", "DEPLOYMENT", "TEST", "RESULT", "EVIDENCE",
                                    "CLAIM", "VERIFIER"} <= kinds, sorted(kinds))

    ne = src.json("NEGATIVE_EVIDENCE.json")
    check("negative_evidence.counts", ne["counts"]["failed_claims"] + ne["counts"]["defects_fixed"] == len(ne["entries"]))
    check("negative_evidence.commits", all(len(e.get("commit", "")) == 40 for e in ne["entries"] if e["kind"] == "DEFECT_FIXED"))

    claims_out = []
    if with_claims:
        reg_raw = fetch(SITES["clawx.click"] + "claims/CAIN42_FINAL_PUBLIC_CLAIMS.json")
        reg = json.loads(reg_raw)
        rb = {k: v for k, v in reg.items() if k not in ("registry_digest", "signature_b64")}
        rdg = H(canon(rb))
        reg_ok = rdg == reg["registry_digest"] and ed_ok(reg["evidence_root_public_key_b64"], reg["signature_b64"], rdg.encode())
        check("claims.registry_signature", reg_ok)
        check("claims.same_root_key", reg["evidence_root_public_key_b64"] == root_key)
        cache = {}
        for c in reg["claims"]:
            probs = []
            for a in c["artifacts"]:
                for site, base in SITES.items():
                    url = base + f"{a['bundle']}/{a['file']}"
                    try:
                        if url not in cache:
                            cache[url] = H(fetch(url))
                        if cache[url] != a["sha256"]:
                            probs.append(f"{site}:{a['file']} hash mismatch")
                    except Exception as e:  # noqa: BLE001
                        probs.append(f"{site}:{a['file']} {type(e).__name__}")
            claims_out.append({"claim": c["claim_id"], "status": c["status"], "state": (c.get("status_model") or {}).get("state"),
                               "label": c.get("public_label"), "commit": reg.get("git_commit"),
                               "artifacts": [f"{a['bundle']}/{a['file']} sha256:{a['sha256'][:16]}" for a in c["artifacts"]],
                               "signature": "VALID" if reg_ok else "INVALID",
                               "evidence": "VALID" if not probs else "INVALID", "problems": probs,
                               "result": ("VERIFIED" if not probs and reg_ok and c["status"] == "VERIFIED" else
                                          "EVIDENCE INTACT (claim status: %s)" % c["status"] if not probs and reg_ok else "FAIL")})
        check("claims.all_artifacts_intact_on_3_sites", all(x["evidence"] == "VALID" for x in claims_out),
              f"{sum(x['evidence'] == 'VALID' for x in claims_out)}/{len(claims_out)} claims")
    failed = [c["check"] for c in checks if c["result"] == "FAIL"]
    return {"verifier": "verify_proof_fabric.py", "clean_room": True, "imports_cain": False, "source": src.loc,
            "result": "PASS" if not failed else "FAIL", "passed": len(checks) - len(failed), "failed": failed,
            "checks": checks, "claims": claims_out}


def main(argv) -> int:
    if "--pack" in argv:
        p = argv[argv.index("--pack") + 1]
        data = fetch(p) if p.startswith("http") else Path(p).read_bytes()
        z = zipfile.ZipFile(io.BytesIO(data))
        names = z.namelist()
        prefix = names[0].split("/")[0] + "/"
        src = Source("pack:" + p)
        src.zip, src.prefix = z, prefix + "proof-fabric/"
        res = verify(src, False)
        # pack-level: every file listed in the signed PACK_MANIFEST hashes as listed
        pm = json.loads(z.read(prefix + "PACK_MANIFEST.json"))
        body = {k: v for k, v in pm.items() if k != "signature"}
        dg = H(canon({"domain": D_PROOF, **body}))
        ok_sig = dg == pm["signature"]["digest_sha256"] and ed_ok(pm["signature"]["public_key_b64"], pm["signature"]["signature_b64"], dg.encode())
        bad = [f for f, h in pm["files"].items() if H(z.read(prefix + f)) != h]
        res["checks"] += [{"check": "pack.manifest_signature", "result": "PASS" if ok_sig else "FAIL", "detail": ""},
                          {"check": "pack.file_hashes", "result": "PASS" if not bad else "FAIL",
                           "detail": f"{len(pm['files'])} files; bad {bad[:5]}"}]
        res["failed"] += [c["check"] for c in res["checks"][-2:] if c["result"] == "FAIL"]
        res["passed"] = len(res["checks"]) - len(res["failed"])
        res["result"] = "PASS" if not res["failed"] else "FAIL"
    else:
        args = [a for a in argv[1:] if not a.startswith("--")]
        if not args:
            print(__doc__)
            return 2
        res = verify(Source(args[0]), "--claims" in argv)
    if "--text" in argv:
        for c in res["claims"]:
            print(f"{'PASS' if c['evidence'] == 'VALID' else 'FAIL'}\n  claim: {c['claim']}\n  status: {c['status']} ({c['state']})\n"
                  f"  commit: {c['commit']}\n  artifacts: {'; '.join(c['artifacts']) or '-'}\n  signature: {c['signature']}\n"
                  f"  evidence: {c['evidence']}\n  result: {c['result']}\n")
        for c in res["checks"]:
            print(f"{c['result']:4s}  {c['check']}  {c['detail']}")
        print(f"\nRESULT: {res['result']}  ({res['passed']} passed, {len(res['failed'])} failed)")
    else:
        print(json.dumps(res, indent=2))
    return 0 if res["result"] == "PASS" else 1


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