#!/usr/bin/env python3
"""Clean-room offline verifier for CAIN governance contracts (cain.contract.v1).

Written from the spec in platform-gateway/cain_governance_contract.py's docstring and the evidence bundle's
CONTRACT_SPEC.md. It imports NOTHING from CAIN: standard library + `cryptography` only. It re-implements the
canonical form, digests, signatures, structural rules, risk scoring and the decision function independently, so a
bug in the runtime shows up as a disagreement here (tests/test_governance_contract_cleanroom.py diffs the two).

    python3 scripts/contract_verify_offline.py export.json --issuer-key <base64 Ed25519> [--tenant T] [--at-ms N]

Input: the JSON produced by GovernanceContractFabric.export(contract_id). Exit 0 if every check passes.
Without --issuer-key the issuer key inside the export is used and the result says SELF_ASSERTED_KEY: that proves only
internal consistency, not who issued it.
"""
from __future__ import annotations

import argparse
import base64
import hashlib
import json
import os
import re
import sys
import unicodedata

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

SECTION_NAMES = ["header", "subject", "purpose", "capability", "authority", "policy", "risk", "action",
                 "execution_path", "environment", "evidence_policy", "conditions", "validity", "annotations"]
US = b"\x1f"
LIMIT = 1 << 53


# -------------------------------------------------------------------------------------------- canonical form
class Bad(Exception):
    pass


def _walk(v, where="$"):
    if v is None or v is True or v is False:
        return
    t = type(v)
    if t is int:
        if v <= -LIMIT or v >= LIMIT:
            raise Bad(f"{where}: integer out of range")
    elif t is float:
        raise Bad(f"{where}: float")
    elif t is str:
        if unicodedata.normalize("NFC", v) != v or "\x00" in v:
            raise Bad(f"{where}: non-NFC or NUL string")
    elif t is list:
        for n, x in enumerate(v):
            _walk(x, f"{where}[{n}]")
    elif t is dict:
        for k in v:
            if type(k) is not str or not re.fullmatch(r"[a-z0-9_]{1,64}", k):
                raise Bad(f"{where}: bad key {k!r}")
            _walk(v[k], f"{where}.{k}")
    else:
        raise Bad(f"{where}: type {t.__name__}")


def canon(v) -> bytes:
    _walk(v)
    return json.dumps(v, sort_keys=True, separators=(",", ":"), ensure_ascii=False, allow_nan=False).encode()


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


def fp_of(body) -> str:
    digests = {}
    for name in SECTION_NAMES:
        digests[name] = h(b"CAIN/contract/v1/section" + US + name.encode() + US + canon(body[name]))
    return h(b"CAIN/contract/v1/fingerprint" + US + canon(digests))


def req_hash(req) -> str:
    return h(b"CAIN/contract/v1/request" + US + canon(req))


def grant_of(snapshot) -> str:
    return h(b"CAIN/contract/v1/grant" + US + canon(snapshot))


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


# -------------------------------------------------------------------------------------------- spec rules
KINDS = {"openshell": ["exec"], "memvault": ["memory_write", "memory_read"], "scitt": ["evidence_register"]}
ENGINE = {"openshell": "cain_openshell_gateway", "memvault": "cain_memvault", "scitt": "cain_scitt_clearinghouse"}
BASE = {"exec": 40, "memory_write": 25, "memory_read": 10, "evidence_register": 5}
CLASSES = ["low", "medium", "high", "critical"]
TRUST = {"low": 0, "medium": 1, "high": 2}


def score(body) -> int:
    s = BASE.get(body["action"]["kind"], 50)
    s += 0 if body["action"]["request_hash"] else 20
    s += min(body["capability"]["max_uses"], 10)
    v = body["validity"]
    s += 10 if v["expires_at_ms"] - v["not_before_ms"] > 900_000 else 0
    s += 0 if body["conditions"] else 10
    s += 0 if body["evidence_policy"]["scitt_required"] else 10
    return min(s, 100)


def klass(sc: int) -> str:
    for threshold, name in ((75, "critical"), (50, "high"), (25, "medium")):
        if sc >= threshold:
            return name
    return "low"


def structural(body) -> list:
    """Independent subset of the validator: the rules an auditor most needs. Returns a list of problems."""
    p = []
    if sorted(body) != sorted(SECTION_NAMES):
        return ["sections differ from the 14 required"]
    hd = body["header"]
    if hd.get("schema") != "cain.contract.v1":
        p.append("schema")
    if not re.fullmatch(r"ctr_[0-9a-f]{32}", str(hd.get("contract_id"))):
        p.append("contract_id")
    ver = hd.get("version")
    if type(ver) is not int or ver < 1:
        p.append("version")
    elif (ver == 1) != (hd.get("supersedes") == ""):
        p.append("supersedes/version mismatch")
    prod, kind = body["policy"].get("product"), body["action"].get("kind")
    if kind not in KINDS.get(prod, []):
        p.append("action kind does not belong to product (confusion)")
    if body["capability"].get("action") != f"{prod}.{kind}":
        p.append("capability.action != product.kind")
    ep = body["execution_path"]
    if ep.get("product") != prod or ep.get("engine") != ENGINE.get(prod) or ep.get("adapter_version") != "1":
        p.append("execution_path mismatch or downgrade")
    if body["authority"].get("source") != "rig":
        p.append("authority source is not rig")
    ch = body["authority"].get("agent_chain") or []
    if not ch or ch[-1] != body["subject"].get("agent_id") or len(set(ch)) != len(ch):
        p.append("agent_chain does not end at the subject")
    if body["environment"].get("model_runtime_digest") != body["subject"].get("runtime_digest"):
        p.append("model_runtime_digest != subject.runtime_digest")
    v = body["validity"]
    span = v["expires_at_ms"] - v["not_before_ms"]
    if not 1000 <= span <= 86_400_000:
        p.append("validity window out of bounds")
    try:
        sc = score(body)
        declared = body["risk"]["risk_class"]
        if declared not in CLASSES or CLASSES.index(declared) < CLASSES.index(klass(sc)):
            p.append(f"risk downgrade: declared {declared}, computed {klass(sc)} ({sc})")
        if sc > body["risk"]["max_risk_score"]:
            p.append("risk score above declared ceiling")
        if declared == "critical" and (not body["action"]["request_hash"] or body["capability"]["max_uses"] != 1
                                       or span > 300_000):
            p.append("critical contract not pinned / single-use / short")
        if declared == "high" and not body["action"]["request_hash"] and not body["conditions"]:
            p.append("high contract neither pinned nor conditioned")
    except (KeyError, TypeError):
        p.append("risk section malformed")
    return p


# -------------------------------------------------------------------------------------------- decision (re-derived)
def _under(path, root):
    path = os.path.normpath(path)
    root = root.rstrip("/") or "/"
    return root == "/" or path == root or path.startswith(root + "/")


def _scope_has(res, sc):
    if sc == "*":
        return True
    if sc[-1:] == "*":
        return res.startswith(sc[:-1])
    return res == sc or res.startswith(sc.rstrip("/") + "/")


def _request_shape_ok(kind, r) -> bool:
    if type(r) is not dict:
        return False
    try:
        if len(canon(r)) > 64 * 1024:
            return False
    except Bad:
        return False
    s = lambda x: type(x) is str and x != ""  # noqa: E731
    if kind == "exec":
        if set(r) != {"argv", "cwd", "timeout_s"}:
            return False
        av = r["argv"]
        if type(av) is not list or not 1 <= len(av) <= 64 or any(type(a) is not str or len(a) > 4096 for a in av):
            return False
        cwd = r["cwd"]
        if type(cwd) is not str or not cwd.startswith("/") or len(cwd) > 512 or ".." in cwd.split("/") or \
                any(ch in cwd for ch in "\0\n\r"):
            return False
        return type(r["timeout_s"]) is int and 1 <= r["timeout_s"] <= 60
    if kind == "memory_write":
        return set(r) == {"namespace", "memory_key", "content", "source_type", "source_id"} and \
            all(s(r[k]) for k in ("namespace", "memory_key", "content", "source_type")) and type(r["source_id"]) is str
    if kind == "memory_read":
        return set(r) == {"namespace", "query", "top_k", "min_trust"} and s(r["query"]) and s(r["namespace"]) and \
            type(r["top_k"]) is int and 1 <= r["top_k"] <= 50 and type(r["min_trust"]) is str and r["min_trust"] in TRUST
    if kind == "evidence_register":
        return set(r) == {"payload_type", "subject", "body"} and type(r["body"]) is dict and bool(r["body"]) and \
            all(type(r[k]) is str and 0 < len(r[k]) <= 256 for k in ("payload_type", "subject"))
    return False


def _condition_holds(op, arg, kind, r, at_ms) -> bool:
    if op == "argv0_in":
        a0 = r["argv"][0]
        return "/" not in a0 and a0 in arg
    if op == "max_argc":
        return len(r["argv"]) <= arg
    if op == "cwd_under":
        return any(_under(r["cwd"], x) for x in arg)
    if op == "max_timeout_s":
        return r["timeout_s"] <= arg
    if op == "namespace_in":
        return r["namespace"] in arg
    if op == "source_type_in":
        return r["source_type"] in arg
    if op == "min_trust":
        return TRUST[r["min_trust"]] >= TRUST[arg]
    if op == "max_content_bytes":
        text = r["content"] if kind == "memory_write" else canon(r["body"]).decode()
        return len(text.encode()) <= arg
    if op == "max_top_k":
        return r["top_k"] <= arg
    if op == "payload_type_in":
        return r["payload_type"] in arg
    if op == "utc_hour_window":
        return arg[0] <= (at_ms // 3_600_000) % 24 < arg[1]
    return False


def decision(rp) -> tuple:
    """Final decision: decide -> (Evolution 3) forecast gate -> quorum, the order the runtime uses."""
    if "forecast_policy" not in rp:
        return _decide(rp)
    base = _decide(rp, quorum=False)
    d, why, _v = forecast_gate(base, rp.get("forecast"), rp["forecast_policy"], rp.get("budget_before"),
                               rp.get("forecast_error", ""))
    return _decide(rp) if d == "ALLOW" else (d, why)


def _decide(rp, quorum=True) -> tuple:
    """Independent derivation of the runtime's decide(); check order follows the spec."""
    b, at = rp["contract"], rp["decided_at_ms"]
    if rp["tenant"] != b["header"]["tenant"]:
        return "DENY", "TENANT_MISMATCH"
    if rp["state_before"] != "ACTIVE":
        return "DENY", "STATE_" + rp["state_before"]
    if at < b["validity"]["not_before_ms"]:
        return "DENY", "NOT_YET_VALID"
    if at >= b["validity"]["expires_at_ms"]:
        return "DENY", "EXPIRED"
    if rp["uses_before"] >= b["capability"]["max_uses"]:
        return "DENY", "USES_EXHAUSTED"
    if rp.get("epoch_now") is not None and b["validity"]["epoch"] != rp["epoch_now"]:
        return "DENY", "STALE_EPOCH"
    snap = rp["snapshot"]
    if snap.get("unavailable"):
        return "DENY", "AUTHORITY_UNAVAILABLE"
    if not snap.get("complete"):
        return "DENY", "AUTHORITY_UNKNOWN_AGENT"
    ag = snap["agents"]
    if not ag or ag[-1]["agent_id"] != b["subject"]["agent_id"]:
        return "DENY", "AUTHORITY_SUBJECT_MISMATCH"
    if [x["agent_id"] for x in ag] != b["authority"]["agent_chain"]:
        return "DENY", "AUTHORITY_CHAIN_MISMATCH"
    if any(x["status"] != "ACTIVE" for x in ag):
        return "DENY", "AUTHORITY_REVOKED"
    if ag[-1]["public_key_b64"] != b["subject"]["public_key_b64"]:
        return "DENY", "IDENTITY_SUBSTITUTION"
    if ag[-1]["runtime_digest"] != b["subject"]["runtime_digest"]:
        return "DENY", "MODEL_SUBSTITUTION"
    for x in ag:
        if "*" not in x["actions"] and b["capability"]["action"] not in x["actions"]:
            return "DENY", "AUTHORITY_AMPLIFICATION"
        if not any(_scope_has(b["capability"]["resource"], sc) for sc in x["scopes"]):
            return "DENY", "AUTHORITY_AMPLIFICATION"
    if (rp["runtime_digest"] or "") != b["subject"]["runtime_digest"]:
        return "DENY", "MODEL_SUBSTITUTION"
    if rp["environment_digest"] != b["environment"]["environment_digest"]:
        return "DENY", "ENVIRONMENT_DRIFT"
    if rp["policy_digest"] == "UNAVAILABLE":
        return "DENY", "POLICY_UNAVAILABLE"
    if rp["policy_digest"] != b["policy"]["policy_digest"]:
        return "DENY", "POLICY_DRIFT"
    kind, r = b["action"]["kind"], rp["request"]
    if not _request_shape_ok(kind, r):
        return "DENY", "MALFORMED_REQUEST"
    rh = req_hash(r)
    if b["action"]["request_hash"] and b["action"]["request_hash"] != rh:
        return "DENY", "REQUEST_NOT_PINNED"
    inv = rp["invocation"]
    if type(inv) is not dict or sorted(inv) != ["nonce", "request_hash", "signature_b64", "ts_ms"]:
        return "DENY", "POP_MISSING"
    if inv["request_hash"] != rh:
        return "DENY", "POP_REQUEST_MISMATCH"
    if type(inv["ts_ms"]) is not int or abs(at - inv["ts_ms"]) > 60_000:
        return "DENY", "POP_STALE"
    if type(inv["nonce"]) is not str or not re.fullmatch(r"[0-9a-f]{32}", inv["nonce"]):
        return "DENY", "POP_MALFORMED"
    msg = b"CAIN/contract/v1/invoke" + US + canon({"fp": fp_of(b), "nonce": inv["nonce"], "request_hash": rh,
                                                  "ts_ms": inv["ts_ms"]})
    if not sig_ok(b["subject"]["public_key_b64"], msg, inv["signature_b64"]):
        return "DENY", "POP_INVALID"
    if rp.get("nonce_used"):
        return "DENY", "POP_REPLAY"
    res = b["capability"]["resource"]
    inside = {"exec": lambda: _under(r["cwd"], res),
              "memory_write": lambda: res == "memvault/ns/" + r["namespace"],
              "memory_read": lambda: res == "memvault/ns/" + r["namespace"],
              "evidence_register": lambda: res == "scitt/" + r["payload_type"]}[kind]()
    if not inside:
        return "DENY", "RESOURCE_OUTSIDE_CONTRACT"
    conds = sorted(b["conditions"], key=lambda c: c["op"])
    for c in conds:
        arg = c["arg"]
        if c["op"] == "cwd_under":
            arg = [os.path.normpath(x) for x in arg]
        if not _condition_holds(c["op"], arg, kind, r, at):
            return "DENY", "CONDITION_" + c["op"].upper()
    qs = quorum_status_of(rp) if quorum else "NOT_REQUIRED"
    if qs not in ("NOT_REQUIRED", "VALID"):
        return "DENY", "QUORUM_" + qs
    return "ALLOW", "OK"


# -------------------------------------------------------------------------------------------- forecasts (Evolution 3)
# Re-implemented from the spec, not imported. Verifies STRUCTURE, ARITHMETIC, BINDING and the gate -- never whether a
# forecast was right about the world.
F_BP = 10_000
F_DIMS = ("security", "privacy", "financial", "operational", "reputational", "availability", "data", "identity",
          "authority", "compliance", "physical", "network", "infrastructure", "model", "agent", "organizational",
          "economic", "supply_chain", "social")
F_UNC = ("aleatoric", "epistemic", "model", "data", "environment", "distribution_shift", "sensor", "dependency",
         "execution")
F_REV = ("REVERSIBLE", "PARTIAL", "IRREVERSIBLE", "UNKNOWN")
F_TTL = 30_000


def f_digest(fc) -> str:
    return h(b"CAIN/contract/v1/forecast" + US + canon(fc))


def f_policy_digest(p) -> str:
    return h(b"CAIN/contract/v1/forecast-policy" + US + canon(p))


def _f_eff(d):
    return d.get("residual_impact_bp", d["impact_bp"])


def _f_exp(d):
    return _f_eff(d) * d["probability_bp"] // F_BP


def _f_hor(d):
    e, out = _f_exp(d), {}
    for t in (0, 300, 86_400):
        out[str(t)] = 0 if t < d["time_to_effect_s"] else (e // 2 if d["reversibility"] == "REVERSIBLE" and
                                                           t >= 86_400 else e)
    return out


def f_summary(fc):
    m = {k: d for k, d in fc["consequence_vector"].items() if d.get("status") == "MODELED"}
    tot = sum(_f_exp(d) for d in m.values())
    return {"expected_total_bp": tot, "max_impact_bp": max([_f_eff(d) for d in m.values()] or [0]),
            "adverse_probability_bp": max([d["probability_bp"] for d in m.values()] or [0]),
            "worst_dimension": max(m, key=lambda k: (_f_exp(m[k]), k)) if m else "",
            "irreversible_dimensions": sorted(k for k, d in m.items() if d["reversibility"] in ("IRREVERSIBLE",
                                                                                               "UNKNOWN")),
            "budget_cost": max(1, tot // 100),
            "unknown_dimensions": sorted(k for k, d in fc["consequence_vector"].items() if d.get("status") != "MODELED")}


def f_problems(fc) -> list:
    try:
        canon(fc)
    except Bad:
        return ["not canonical"]
    if type(fc) is not dict or fc.get("schema") != "cain.contract.forecast.v1" or fc.get("grants_authority") is not False:
        return ["schema/grants_authority"]
    cv = fc.get("consequence_vector")
    if type(cv) is not dict or set(cv) != set(F_DIMS):
        return ["consequence vector incomplete"]
    p = []
    for k, d in cv.items():
        st = d.get("status") if type(d) is dict else None
        if st not in ("MODELED", "NOT_MODELED", "NOT_APPLICABLE"):
            p.append(k)
            continue
        if st != "MODELED":
            continue
        if any(type(d.get(f)) is not int or d[f] < 0 or (f != "time_to_effect_s" and d[f] > F_BP)
               for f in ("impact_bp", "probability_bp", "time_to_effect_s", "expected_bp")) or \
                d.get("reversibility") not in F_REV:
            p.append(k)
            continue
        if "residual_impact_bp" in d and (type(d["residual_impact_bp"]) is not int or
                                          not 0 <= d["residual_impact_bp"] <= d["impact_bp"]):
            p.append(k)
            continue
        if d["expected_bp"] != _f_exp(d) or d.get("horizon_bp") != _f_hor(d):
            p.append(k + " arithmetic")
    u = fc.get("uncertainty")
    if type(u) is not dict or set(u) != set(F_UNC) or any(
            type(v) is not dict or not (("bp" in v and type(v["bp"]) is int and 0 <= v["bp"] <= F_BP) or
                                        v.get("status") in ("UNQUANTIFIED", "NOT_APPLICABLE")) for v in u.values()):
        p.append("uncertainty")
    v = fc.get("validity") or {}
    if not (type(v.get("issued_at_ms")) is int and type(v.get("valid_until_ms")) is int and
            0 < v["valid_until_ms"] - v["issued_at_ms"] <= F_TTL):
        p.append("validity")
    if not p and fc.get("summary") != f_summary(fc):
        p.append("summary")
    return p


def forecast_gate(base, fc, pol, budget, err=""):
    if base[0] != "ALLOW":
        return base[0], base[1], "NOT_EVALUATED"
    v = "PASS"
    if err:
        v = "FORECAST_UNAVAILABLE"
    elif fc is None or f_problems(fc):
        v = "FORECAST_INVALID"
    else:
        cv, s = fc["consequence_vector"], fc["summary"]
        for dim in sorted(pol["max_expected_bp"]):
            d = cv.get(dim, {})
            if d.get("status") == "MODELED" and d["expected_bp"] >= pol["max_expected_bp"][dim]:
                v = "FORECAST_IMPACT_" + dim.upper()
                break
        else:
            if fc["uncertainty"]["epistemic"].get("status") == "UNQUANTIFIED" and \
                    s["max_impact_bp"] >= pol["unquantified_impact_bp"]:
                v = "FORECAST_UNQUANTIFIED_HIGH_IMPACT"
            elif fc["uncertainty"]["model"].get("bp", F_BP) >= 8000 and s["max_impact_bp"] >= pol["unquantified_impact_bp"]:
                v = "FORECAST_MODEL_CANNOT_PREDICT"
            elif fc["drift"].get("status") == "MATERIAL" and s["max_impact_bp"] >= pol["drift_deny_impact_bp"]:
                v = "FORECAST_DRIFT"
            elif budget is None:
                v = "AUTONOMY_BUDGET_UNAVAILABLE"
            elif budget["used"] + s["budget_cost"] > budget["limit"]:
                v = "AUTONOMY_BUDGET_EXHAUSTED"
    if v == "PASS":
        return "ALLOW", "OK", "PASS"
    if pol.get("mode") == "observe":
        return "ALLOW", "OK", "OBSERVE_WOULD_DENY:" + v
    return "DENY", v, "DENY:" + v


def forecast_bound(rc, rp) -> list:
    """Every forecast input must be signed by the receipt, and the recorded verdict must be re-derivable."""
    p = []
    fc = rp.get("forecast")
    try:
        if (f_digest(fc) if fc is not None else "") != rc.get("forecast_digest"):
            p.append("forecast digest")
        if f_policy_digest(rp["forecast_policy"]) != rc.get("forecast_policy_digest"):
            p.append("forecast policy digest")
    except (Bad, KeyError, TypeError):
        return ["forecast record malformed"]
    if rp.get("budget_before") != rc.get("budget_before"):
        p.append("budget")
    if fc is not None and f_problems(fc):
        p.append("forecast structure: " + ",".join(f_problems(fc))[:120])
    try:
        _d, _w, verdict = forecast_gate(_decide(rp, quorum=False), fc, rp["forecast_policy"], rp.get("budget_before"),
                                        rp.get("forecast_error", ""))
    except Exception as exc:  # noqa: BLE001 -- malformed records fail closed
        return p + [f"gate raised {type(exc).__name__}"]
    if verdict != rc.get("forecast_verdict"):
        p.append(f"verdict {rc.get('forecast_verdict')} not re-derivable (got {verdict})")
    bd = rp.get("execution_binding")
    if rc.get("decision") == "ALLOW" and fc is not None and (not bd or bd.get("forecast_digest") != f_digest(fc)):
        p.append("ALLOW binding does not carry the forecast")
    return p


# -------------------------------------------------------------------------------------------- quorum (re-derived)
def quorum_params(n: int):
    f = (n - 1) // 3
    return f, -(-(n + f + 1) // 2)


def membership_digest(replicas: dict) -> str:
    return h(b"CAIN/contract/v1/membership" + US + canon(dict(sorted(replicas.items()))))


def qcert_ok(cert, replicas: dict, kind: str, expect: dict) -> bool:
    """>= q valid signatures from distinct members over the exact fields, for this membership."""
    if type(cert) is not dict or cert.get("schema") != "cain.contract.qcert.v1" or cert.get("kind") != kind:
        return False
    if cert.get("membership_digest") != membership_digest(replicas):
        return False
    try:
        if kind == "use":
            fields = {"fp": cert["fp"], "use_index": cert["use_index"], "nonce": cert["nonce"],
                      "request_hash": cert["request_hash"]}
            tag = b"CAIN/contract/v1/q-use"
        elif kind == "authorize":
            fields = {"fp": cert["fp"], "grant_digest": cert["grant_digest"], "epoch": cert["epoch"]}
            tag = b"CAIN/contract/v1/q-authorize"
        else:
            fields = {"fp": cert["fp"], "revoked": True}
            tag = b"CAIN/contract/v1/q-revoke"
    except KeyError:
        return False
    if any(fields.get(k) != v for k, v in expect.items()):
        return False
    msg = tag + US + canon(fields)
    good = {rid for rid, sig in (cert.get("signatures") or {}).items()
            if rid in replicas and type(sig) is str and sig_ok(replicas[rid], msg, sig)}
    return len(good) >= quorum_params(len(replicas))[1]


def quorum_status_of(rp) -> str:
    """Never trust a recorded VALID: re-check the certificate. If the export names a membership, a decision without a
    valid certificate cannot be an ALLOW."""
    st = rp.get("quorum_status", "NOT_REQUIRED")
    mem = rp.get("quorum_membership")
    if mem is None:
        return "NOT_REQUIRED" if st == "NOT_REQUIRED" else st
    if st != "VALID":
        return st if st != "NOT_REQUIRED" else "MISSING"
    inv = rp.get("invocation") if type(rp.get("invocation")) is dict else {}
    ok = qcert_ok(rp.get("quorum_cert"), dict(mem.get("replicas") or {}), "use",
                  {"fp": fp_of(rp["contract"]), "nonce": inv.get("nonce"), "request_hash": inv.get("request_hash")})
    return "VALID" if ok else "INVALID"


# -------------------------------------------------------------------------------------------- execution binding
def binding_digest_of(binding) -> str:
    return h(b"CAIN/contract/v1/execution-binding" + US + canon(binding))


def binding_consistent(rp) -> list:
    """The decision-time binding must restate exactly the recorded inputs."""
    bd, b = rp.get("execution_binding"), rp["contract"]
    if not bd:
        return []
    p = []
    inv = rp.get("invocation") or {}
    exp = {"contract_digest": fp_of(b), "authority_digest": grant_of(rp["snapshot"]),
           "policy_digest": rp["policy_digest"], "environment_digest": rp["environment_digest"],
           "action_digest": req_hash(rp["request"]) if rp.get("request") is not None else None,
           "nonce": inv.get("nonce"), "tenant_digest": h(rp["tenant"].encode()),
           "governance_epoch": rp.get("epoch_now"), "expiration": b["validity"]["expires_at_ms"],
           "model_runtime_digest": h((rp.get("runtime_digest") or "").encode()), "revocation_state": "NOT_REVOKED"}
    for k, v in exp.items():
        if bd.get(k) != v:
            p.append(k)
    return p


# -------------------------------------------------------------------------------------------- key trust chain
def _keycert_msg(c) -> bytes:
    body = {k: c[k] for k in ("schema", "key_id", "public_key_b64", "purpose", "not_before_ms", "not_after_ms",
                              "serial", "root_key_id", "kind")}
    return b"CAIN/contract/v1/keycert" + US + canon(body)


def _keyrevoke_msg(r) -> bytes:
    body = {k: r[k] for k in ("schema", "key_id", "revoked_at_ms", "reason", "compromised", "root_key_id")}
    return b"CAIN/contract/v1/keyrevoke" + US + canon(body)


def key_for(trust: dict, root_pub: str, key_id: str, at_ms: int):
    """(public_key or None, reason). Independent re-derivation of cain_contract_keys.check_key_at."""
    try:
        rkid = hashlib.sha256(base64.b64decode(root_pub)).hexdigest()[:16]
    except (ValueError, TypeError):
        return None, "BAD_ROOT"
    issue, cap = None, None
    for c in trust.get("certs", []):
        try:
            if c["key_id"] != key_id or c["root_key_id"] != rkid or c["purpose"] != "contract-authority":
                continue
            if hashlib.sha256(base64.b64decode(c["public_key_b64"])).hexdigest()[:16] != key_id:
                continue
            if not sig_ok(root_pub, _keycert_msg(c), c["signature_b64"]):
                continue
        except (KeyError, TypeError, ValueError):
            continue
        if c["kind"] == "issue":
            issue = c
        elif c["kind"] == "retire":
            cap = c["not_after_ms"] if cap is None else min(cap, c["not_after_ms"])
    if issue is None:
        return None, "KEY_NOT_CERTIFIED"
    for r in trust.get("revocations", []):
        try:
            if r["key_id"] != key_id or not sig_ok(root_pub, _keyrevoke_msg(r), r["signature_b64"]):
                continue
        except (KeyError, TypeError, ValueError):
            continue
        if r["compromised"]:
            return None, "KEY_COMPROMISED"
        cap = r["revoked_at_ms"] if cap is None else min(cap, r["revoked_at_ms"])
    end = issue["not_after_ms"] if cap is None else min(issue["not_after_ms"], cap)
    if not issue["not_before_ms"] <= at_ms < end:
        return None, "KEY_OUTSIDE_VALIDITY"
    return issue["public_key_b64"], "OK"


# -------------------------------------------------------------------------------------------- the verifier
def verify(export: dict, issuer_key_b64: str = "", tenant: str = "", at_ms: int = -1, root_key_b64: str = "") -> dict:
    """With `root_key_b64` (and an export whose trust mode is root-chain), every signature's key is resolved through
    the root-certified chain at that artifact's signed time; otherwise `issuer_key_b64` (or, unpinned, the export's own
    key) is used for every signature."""
    checks, problems = {}, []

    def mark(name, ok, why=""):
        checks[name] = bool(ok)
        if not ok:
            problems.append(f"{name}: {why}" if why else name)

    body = export.get("contract")
    try:
        raw = canon(body)
        mark("canonical_encoding", json.loads(raw) == body)
        fp = fp_of(body)
        mark("fingerprint", fp == export.get("fingerprint"), "recomputed fingerprint differs")
    except (Bad, KeyError, TypeError) as exc:
        mark("canonical_encoding", False, str(exc))
        return {"ok": False, "checks": checks, "problems": problems}
    st = structural(body)
    mark("structure_and_risk", not st, "; ".join(st))

    trust = export.get("trust") if type(export.get("trust")) is dict else {}
    chain_mode = bool(root_key_b64) and trust.get("mode") == "root-chain"
    if root_key_b64 and not chain_mode:
        # Found by the attack engine: flipping trust.mode made a root-pinned check fall back to the export's own key.
        mark("trust_downgrade", False, "a root was pinned but the export is not root-chain signed")
    if root_key_b64 and trust.get("mode") == "root-chain" and trust.get("root_public_key_b64") not in (None, "", root_key_b64):
        mark("root_matches_pin", False, "export's root differs from the pinned root")

    def resolve(key_id, at):
        if chain_mode:
            pub, why = key_for(trust, root_key_b64, key_id or "", int(at or 0))
            return pub or "", why
        return (issuer_key_b64 or export.get("issuer_public_key_b64", "")), "OK"

    key, why = resolve(export.get("issuer_key_id"), export.get("authorized_at_ms"))
    checks["issuer_key_pinned"] = bool(issuer_key_b64) or chain_mode
    if issuer_key_b64 and not chain_mode and export.get("issuer_public_key_b64") not in ("", None, issuer_key_b64):
        mark("issuer_key_matches_pin", False, "export names a different issuer key than the pinned one")
    mark("issuer_signature", bool(key) and sig_ok(key, b"CAIN/contract/v1/issuer" + US + fp.encode() + US +
                                                  str(export.get("authority_digest", "")).encode(),
                                                  export.get("issuer_sig", "")),
         f"issuer signature does not verify ({why})")

    subj_key = body["subject"]["public_key_b64"]
    if export.get("consent_mode") not in ("direct", "carried", "", None):
        mark("consent_mode", False, "unknown consent mode")
    if export.get("consent_mode") == "carried":
        pred = export.get("predecessor") or {}
        pb = pred.get("body")
        ok = False
        why = "predecessor missing"
        if pb:
            pfp = fp_of(pb)
            only_annotations = all(pb[s] == body[s] for s in SECTION_NAMES if s not in ("header", "validity",
                                                                                         "annotations"))
            hdr_ok = (pb["header"]["lineage_id"] == body["header"]["lineage_id"] and
                      pb["header"]["tenant"] == body["header"]["tenant"] and
                      pb["header"]["schema"] == body["header"]["schema"] and
                      body["header"]["version"] == pb["header"]["version"] + 1 and
                      body["header"]["supersedes"] == pfp)
            val_ok = all(pb["validity"][k] == body["validity"][k] for k in ("not_before_ms", "expires_at_ms"))
            sig = sig_ok(subj_key, b"CAIN/contract/v1/subject" + US + pfp.encode(), export.get("subject_sig", ""))
            ok = only_annotations and hdr_ok and val_ok and sig and export.get("consent_from") == pfp
            why = "carried consent is valid only for annotation-only amendments with the predecessor's signature"
        mark("subject_consent", ok, why)
    else:
        mark("subject_consent", sig_ok(subj_key, b"CAIN/contract/v1/subject" + US + fp.encode(),
                                       export.get("subject_sig", "")), "subject consent signature does not verify")

    if tenant:
        mark("tenant", body["header"]["tenant"] == tenant, "contract tenant differs from expected")
    if at_ms >= 0:
        v = body["validity"]
        mark("validity_at_time", v["not_before_ms"] <= at_ms < v["expires_at_ms"], "outside validity window")

    rv = export.get("revocation")
    if rv is not None and type(rv) is not dict:
        mark("revocation_statement", False, "malformed revocation")
        rv = None
    if rv is not None:
        unsigned = {k: x for k, x in rv.items() if k not in ("signature_b64", "quorum_certified")}
        rkey, _ = resolve(rv.get("issuer_key_id"), rv.get("revoked_at_ms"))
        good = bool(rkey) and sig_ok(rkey, b"CAIN/contract/v1/revocation" + US + canon(unsigned),
                                     rv.get("signature_b64", "")) and rv.get("fingerprint") == fp
        mark("revocation_statement", good, "revocation statement invalid")
        checks["revoked"] = True
        if export.get("state") != "REVOKED":
            mark("revocation_state_consistent", False, "signed revocation exists but state is not REVOKED")

    mem = export.get("quorum_membership")
    for j, qc in enumerate(export.get("quorum_certs") or []):
        if type(qc) is not dict:
            mark(f"quorum_cert[{j}]", False, "malformed")
            continue
        reps = dict((mem or {}).get("replicas") or {}) if type(mem) is dict else {}
        exp = {"fp": fp}
        if qc.get("kind") == "authorize":
            exp.update(grant_digest=export.get("authority_digest"), epoch=body["validity"]["epoch"])
        mark(f"quorum_cert[{j}].{qc.get('kind')}", bool(mem) and qcert_ok(qc.get("cert"), reps, qc.get("kind"), exp),
             "quorum certificate invalid")
    if mem is not None and export.get("state") not in ("DRAFT", "PROPOSED", "VALIDATED", "INVALID", "REJECTED",
                                                       "ABORTED"):
        mark("quorum_authorized", any(q.get("kind") == "authorize" for q in export.get("quorum_certs") or []),
             "quorum-governed contract without an authorization certificate")
    receipts = export.get("receipts") or []
    head = export.get("head")
    if head is not None or export.get("schema") == "cain.contract.export.v2":
        try:
            unsigned_h = {k: x for k, x in head.items() if k != "signature_b64"}
            hkey, hwhy = resolve(head.get("issuer_key_id"), head.get("exported_at_ms"))
            ok_h = (bool(hkey) and sig_ok(hkey, b"CAIN/contract/v1/export-head" + US + canon(unsigned_h),
                                          head.get("signature_b64", ""))
                    and head.get("fingerprint") == fp and head.get("receipts") == len(receipts)
                    and head.get("last_receipt_hash") == (receipts[-1]["receipt"].get("receipt_hash")
                                                          if receipts else "0" * 64)
                    and head.get("state") == export.get("state"))
            mark("export_head", ok_h, "signed export head does not match the receipts/state (truncation or edit)")
        except (AttributeError, TypeError, KeyError, Bad):
            mark("export_head", False, "missing or malformed export head")
    prev_h = "0" * 64
    for i, item in enumerate(receipts):
        rc = item.get("receipt") if type(item) is dict else None
        if type(rc) is not dict:
            mark(f"receipt[{i}].shape", False, "malformed receipt entry")
            continue
        if "contract_seq" in rc:
            mark(f"receipt[{i}].contract_chain", rc.get("contract_seq") == i and
                 rc.get("prev_contract_receipt_hash") == prev_h, "receipt omitted, duplicated or reordered")
        prev_h = rc.get("receipt_hash")
        if rc.get("kind") == "outcome" and item.get("replay") is not None:
            mark(f"receipt[{i}].outcome_replay", False, "an outcome receipt carries no replay record")
        rp = item.get("replay")
        if rc.get("kind") == "decision" and type(rp) is dict and rp.get("quorum_membership") is not None:
            mark(f"receipt[{i}].membership", rp["quorum_membership"].get("replicas") ==
                 (mem or {}).get("replicas"), "replay names a different quorum membership than the export")
            if rc.get("decision") == "ALLOW":
                qc = rp.get("quorum_cert")
                mark(f"receipt[{i}].use_cert_exported", any(
                    q.get("kind") == "use" and q.get("cert") == qc for q in export.get("quorum_certs") or []),
                    "allowed use's quorum certificate missing from the export")
    replays, allowed_decisions = [], set()
    for i, item in enumerate(receipts):
        try:
            _verify_receipt(i, item, mark, resolve, fp, body, allowed_decisions, replays)
        except Exception as exc:  # malformed input is a failed check, never a crash or a pass
            mark(f"receipt[{i}].malformed", False, type(exc).__name__)
    return {"ok": not problems, "checks": checks, "problems": problems, "replays": replays,
            "key_status": "ROOT_CHAIN" if chain_mode else ("PINNED" if issuer_key_b64 else "SELF_ASSERTED_KEY"),
            "fingerprint": fp}


def inputs_signed(rc, rp) -> bool:
    """Every replay input with a signed counterpart in the receipt must equal it."""
    try:
        rh = req_hash(rp["request"]) if rp.get("request") is not None else ""
    except Bad:
        rh = ""
    ok = (rp.get("state_before") == rc.get("state_before") and rp.get("uses_before") == rc.get("uses_before")
          and rp.get("policy_digest") == rc.get("policy_digest")
          and rp.get("environment_digest") == rc.get("environment_digest")
          and (rp.get("runtime_digest") or "") == (rc.get("runtime_digest") or "")
          and rp.get("tenant") == rc.get("tenant") and rh == rc.get("request_hash"))
    if "quorum_status" in rc:
        ok = ok and rp.get("quorum_status", "NOT_REQUIRED") == rc["quorum_status"]
    if "quorum_cert_digest" in rc:
        qc = rp.get("quorum_cert")
        try:
            ok = ok and (h(canon(qc)) if qc else "") == rc["quorum_cert_digest"]
        except Bad:
            ok = False
    return ok


def _verify_receipt(i, item, mark, resolve, fp, body, allowed_decisions, replays):
    if True:  # noqa: SIM108 (kept for diff locality)
        rc, rp = item["receipt"], item["replay"]
        unsigned = {k: x for k, x in rc.items() if k not in ("receipt_hash", "signature_b64")}
        try:
            dig = h(canon(unsigned))
        except Bad as exc:
            mark(f"receipt[{i}].hash", False, str(exc))
            return
        mark(f"receipt[{i}].hash", dig == rc.get("receipt_hash"), "receipt hash mismatch")
        rkey, kwhy = resolve(rc.get("issuer_key_id"), rc.get("decided_at_ms", rc.get("completed_at_ms")))
        mark(f"receipt[{i}].signature", bool(rkey) and sig_ok(rkey, b"CAIN/contract/v1/receipt" + US + dig.encode(),
                                                             rc.get("signature_b64", "")), f"receipt signature ({kwhy})")
        mark(f"receipt[{i}].binding", rc.get("contract_fingerprint") == fp and rc.get("tenant") ==
             body["header"]["tenant"] and rc.get("contract_version") == body["header"]["version"],
             "receipt bound to a different contract/tenant/version")
        if rc.get("kind") == "outcome":
            # an outcome must follow an ALLOW decision of this contract, once
            link = rc.get("decision_receipt_hash")
            mark(f"receipt[{i}].outcome_link", link in allowed_decisions, "outcome without a preceding ALLOW decision")
            allowed_decisions.discard(link)
            pe = rc.get("prediction_error")
            if pe is not None:   # Evolution 3: structure only -- the observation itself is not in the export
                mark(f"receipt[{i}].prediction_error", type(pe) is dict and pe.get("class") in (
                    "true_positive", "false_positive", "true_negative", "false_negative", "underprediction",
                    "overprediction", "unknown") and type(pe.get("divergence")) is bool and
                    re.fullmatch(r"[0-9a-f]{64}", str(pe.get("observation_digest", ""))) is not None,
                     "malformed prediction_error")
            e8 = rc.get("e8")
            if e8 is not None and type(e8) is not dict:
                mark(f"receipt[{i}].e8", False, "malformed e8 record")
            elif e8 is not None:  # Evolution 2: an executed outcome must say E8 committed it; a refusal must not execute
                committed = e8.get("e8_status") == "COMMITTED"
                mark(f"receipt[{i}].e8", committed or rc.get("execution_failed") is True,
                     "outcome executed without an E8 commit")
            return
        if rc.get("kind") != "decision" or rp is None:
            mark(f"receipt[{i}].kind", False, "unknown receipt kind or missing replay record")
            return
        if rc.get("decision") == "ALLOW":
            allowed_decisions.add(rc.get("receipt_hash"))
        bound = (fp_of(rp["contract"]) == rc["contract_fingerprint"] and grant_of(rp["snapshot"]) ==
                 rc["authority_snapshot_digest"] and rp["decided_at_ms"] == rc["decided_at_ms"] and
                 bool(rp.get("nonce_used")) == bool(rc.get("nonce_used")) and
                 rp.get("epoch_now") == rc.get("governance_epoch") and inputs_signed(rc, rp))
        if "execution_binding_digest" in rc and rc["execution_binding_digest"] == "" and \
                rp.get("execution_binding") is not None:
            mark(f"receipt[{i}].execution_binding", False, "binding present but the receipt signed none")
        if rp.get("execution_binding") or rc.get("execution_binding_digest"):
            bad = binding_consistent(rp)
            ok_b = bool(rp.get("execution_binding")) and not bad and \
                binding_digest_of(rp["execution_binding"]) == rc.get("execution_binding_digest")
            mark(f"receipt[{i}].execution_binding", ok_b, f"binding inconsistent: {bad}")
        if rc.get("decision") == "ALLOW" and not rp.get("execution_binding") and "governance_epoch" in rc:
            mark(f"receipt[{i}].execution_binding", False, "an Evolution 2 ALLOW must carry an execution binding")
        if "forecast_digest" in rc or "forecast_policy" in rp:
            fb = forecast_bound(rc, rp) if "forecast_digest" in rc else ["forecast record without a signed digest"]
            mark(f"receipt[{i}].forecast", not fb, "; ".join(fb))
        got = decision(rp)
        same = list(got) == [rc["decision"], rc["reason"]]
        mark(f"receipt[{i}].replay", bound and same,
             f"replay gives {list(got)}, receipt says {[rc['decision'], rc['reason']]}, inputs_bound={bound}")
        replays.append({"recorded": [rc["decision"], rc["reason"]], "replayed": list(got), "bound": bound})


def main(argv=None) -> int:
    ap = argparse.ArgumentParser(description=__doc__.split("\n")[0])
    ap.add_argument("export")
    ap.add_argument("--issuer-key", default="")
    ap.add_argument("--tenant", default="")
    ap.add_argument("--at-ms", type=int, default=-1)
    ap.add_argument("--root-key", default="", help="pinned contract trust root (base64 Ed25519)")
    a = ap.parse_args(argv)
    with open(a.export, "rb") as f:
        exp = json.loads(f.read())
    res = verify(exp, a.issuer_key, a.tenant, a.at_ms, a.root_key)
    print(json.dumps(res, indent=1, sort_keys=True))
    return 0 if res["ok"] else 1


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