#!/usr/bin/env python3
"""grip-verify — the reference verifier for GRIP/1.0-draft.2 §4.6.

Runs the eight-rule audit any third party can perform from a trust bundle
plus an object set. The output is a PER-RULE VERDICT, never a scalar score.

Verdicts
  pass                  the rule held over every object it applies to
  fail                  at least one violation
  unevaluable-critical  a branch a conformant gate refused for an unknown
                        critical extension — correctly-refused authority is
                        not chain failure (§4.6)
  incomplete            the chain could not be shown to extend its latest
                        head attestation (§4.6 Anchoring)
  n/a                   the rule has no subject in this object set

Usage
  grip_verify.py --bundle trust-bundle.json --objects DIR [--json] [--level GRIP-3]
"""
from __future__ import annotations

import argparse
import base64
import glob
import json
import os
import sys
from typing import Any

sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
from gripcore import (  # noqa: E402
    canonicalize, object_id, preimage, verify_detached, jws_kid, digest_bytes,
)
from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PublicKey  # noqa: E402

CLASS_SEVERITY = {"observe": 1, "egress": 2, "settle": 3}
RULES = {
    1: "Signatures", 2: "Chain integrity", 3: "Authority", 4: "Budget",
    5: "Stops", 6: "Settlement and closure", 7: "Context", 8: "Anchoring and time",
}


# ---------------------------------------------------------------- utilities

def resource_matches(pattern: str, resource: str) -> bool:
    """v1 wildcards: a single trailing '*' matches any completion of the
    prefix; exact strings match only themselves. No other forms (§4.3.2)."""
    if pattern.endswith("*"):
        return resource.startswith(pattern[:-1])
    return pattern == resource


def contained(child: str, parent: str) -> bool:
    """Child resource pattern contained in parent's."""
    if parent.endswith("*"):
        prefix = parent[:-1]
        return child.startswith(prefix)
    return child == parent


def governing_class(allow: list[dict], verb: str, resource: str) -> str | None:
    """Most consequence-severe matching class governs (§4.3.2)."""
    best = None
    for e in allow:
        if e.get("verb") == verb and resource_matches(e.get("resource", ""), resource):
            c = e.get("class")
            if c and (best is None or CLASS_SEVERITY[c] > CLASS_SEVERITY[best]):
                best = c
    return best


def budget_le(child: dict, parent: dict) -> list[str]:
    """Law 2: child carries every axis present in parent, each <= parent's.
    Absent-in-parent means unbounded. Returns list of violations."""
    bad = []
    for axis, pv in (parent or {}).items():
        if axis not in (child or {}):
            bad.append(f"child omits axis '{axis}' present in parent")
        elif child[axis] > pv:
            bad.append(f"axis '{axis}': child {child[axis]} > parent {pv}")
    return bad


class Finding:
    def __init__(self, rule: int, obj: str, msg: str, kind: str = "fail"):
        self.rule, self.obj, self.msg, self.kind = rule, obj, msg, kind

    def __repr__(self):
        return f"[rule {self.rule}] {self.obj[:20]}… {self.msg}"


# ------------------------------------------------------------------ loading

def load_bundle(path: str) -> dict:
    raw = json.load(open(path))
    keys = {}
    for k in raw.get("keys", []):
        keys[k["kid"]] = {
            "key": Ed25519PublicKey.from_public_bytes(base64.b64decode(k["public_key"])),
            "nbf": k.get("valid_from", 0),
            "exp": k.get("valid_until", 1 << 62),
        }
    return {"keys": keys, "attestations": raw.get("head_attestations", [])}


def load_objects(d: str) -> list[dict]:
    objs = []
    for f in sorted(glob.glob(os.path.join(d, "*.json"))):
        # NB: never inject metadata into the object — the JCS preimage is
        # the object minus `id` and `sig`, and nothing else.
        objs.append(json.load(open(f)))
    return objs


# ----------------------------------------------------------------- verifier

class Verifier:
    def __init__(self, objects: list[dict], bundle: dict, external: dict | None = None):
        self.objs = objects
        self.bundle = bundle
        self.external = external or {}
        self.by_id = {o["id"]: o for o in objects if "id" in o}
        self.by_type = {}
        for o in objects:
            self.by_type.setdefault(o.get("type"), []).append(o)
        self.findings: list[Finding] = []
        self.verdicts: dict[int, str] = {}
        self.crit_refused: set[str] = set()
        self.notes: list[str] = []

    def add(self, rule, obj, msg, kind="fail"):
        self.findings.append(Finding(rule, obj, msg, kind))

    def _mandate(self):
        m = self.by_type.get("mandate") or []
        return m[0] if m else None

    def _receipts_ordered(self):
        rs = self.by_type.get("receipt", [])
        head = [r for r in rs if "prev" not in r]
        chain, seen = [], set()
        cur = head[0] if head else None
        while cur and cur["id"] not in seen:
            chain.append(cur)
            seen.add(cur["id"])
            cur = next((r for r in rs if r.get("prev") == cur["id"]), None)
        return chain, [r for r in rs if r["id"] not in seen]

    # -- rule 1 ------------------------------------------------------------
    def rule1(self):
        subjects = 0
        for o in self.objs:
            if "sig" not in o:
                continue
            subjects += 1
            kid = o.get("iss")
            entry = self.bundle["keys"].get(kid)
            if not entry:
                self.add(1, o["id"], f"no key in trust bundle for iss {kid}")
                continue
            if not (entry["nbf"] <= o.get("iat", 0) <= entry["exp"]):
                self.add(1, o["id"], f"iat {o.get('iat')} outside key validity window")
            if jws_kid(o["sig"]) != kid:
                self.add(1, o["id"], "JWS kid does not match iss")
            if not verify_detached(o["sig"], preimage(o), entry["key"]):
                self.add(1, o["id"], "signature does not verify")
            # GRIP-2+: actor signature over (grant id ‖ canonical act)
            asig = (o.get("act") or {}).get("actor_sig")
            if asig:
                g = self._grant_for(o)
                actor_entry = self.bundle["keys"].get(o.get("actor"))
                if not actor_entry:
                    self.add(1, o["id"], f"no key for actor {o.get('actor')}")
                else:
                    act_wo = {k: v for k, v in o["act"].items() if k != "actor_sig"}
                    pl = (g["id"] if g else "").encode() + canonicalize(act_wo)
                    if not verify_detached(asig, pl, actor_entry["key"]):
                        self.add(1, o["id"], "act.actor_sig does not verify against actor key")
        self.verdicts[1] = self._verdict(1, subjects)

    # -- rule 2 ------------------------------------------------------------
    def rule2(self):
        subjects = 0
        for o in self.objs:
            if "id" not in o:
                continue
            subjects += 1
            if object_id(o) != o["id"]:
                self.add(2, o["id"], "id does not match JCS digest of its own preimage")
            for p in o.get("parents", []):
                if p not in self.by_id:
                    self.add(2, o["id"], f"parent {p[:18]}… does not resolve")
        chain, orphans = self._receipts_ordered()
        for r in orphans:
            self.add(2, r["id"], "receipt not reachable on the prev-chain (broken or forked link)")
        # referenced digests match supplied content, where content was supplied
        for ref, content in self.external.items():
            if digest_bytes(content) != ref:
                self.add(2, ref, "supplied content does not match its digest")
        self.verdicts[2] = self._verdict(2, subjects)

    def _grant_for(self, receipt):
        for p in receipt.get("parents", []):
            o = self.by_id.get(p)
            if o and o.get("type") == "grant":
                return o
        return None

    # -- rule 3 ------------------------------------------------------------
    def rule3(self):
        m = self._mandate()
        subjects = 0
        for g in self.by_type.get("grant", []):
            subjects += 1
            parent = self.by_id.get(g.get("parents", [None])[0])
            if parent is None:
                self.add(3, g["id"], "grant parent does not resolve")
                continue
            if parent.get("type") == "mandate":
                if g["iss"] not in parent.get("issuers", []):
                    self.add(3, g["id"], f"root grant issuer {g['iss']} not in mandate issuers")
                scope = parent.get("scope")
                if scope is None:
                    self.notes.append(f"root hop of {g['id'][:18]}… reported ISSUER-ATTESTED (mandate carries no scope)")
                else:
                    self._dominance(g, scope, parent.get("budget"), parent.get("exp"), machine=True)
            else:  # delegated
                if not parent.get("delegable"):
                    self.add(3, g["id"], "parent grant is not delegable")
                allowed_issuers = [parent.get("actor")] + list(parent.get("delegation_issuers", []))
                if g["iss"] not in allowed_issuers:
                    self.add(3, g["id"], "delegated grant issuer holds no issuance authority (minting, not delegation)")
                for e in g.get("allow", []):
                    if e.get("resource", "").endswith("*"):
                        self.add(3, g["id"], "wildcard confinement: '*' may appear only in root grants")
                self._dominance(g, parent.get("allow", []), parent.get("budget"), parent.get("exp"), machine=True)

        for r in self.by_type.get("receipt", []):
            verb = (r.get("act") or {}).get("verb", "")
            if verb in ("stop.approve", "stop.reject"):
                continue  # authorised by the Stop itself (§4.3.4)
            subjects += 1
            g = self._grant_for(r)
            if g is None:
                if r.get("decision", {}).get("result") == "allowed":
                    self.add(3, r["id"], "allowed act with no grant on its parents")
                continue
            gc = governing_class(g.get("allow", []), verb, (r.get("act") or {}).get("resource", ""))
            if gc is None and r.get("decision", {}).get("result") == "allowed":
                self.add(3, r["id"], f"act {verb} on {(r.get('act') or {}).get('resource')} matches no allow entry")
            if not (g.get("iat", 0) <= r.get("iat", 0) < g.get("exp", 1 << 62)):
                if r.get("decision", {}).get("result") == "allowed":
                    self.add(3, r["id"], "receipt iat outside grant validity window")
            if m and r.get("iat", 0) >= m.get("exp", 1 << 62) and r.get("decision", {}).get("result") == "allowed":
                self.add(3, r["id"], "allowed act after mandate expiry")
        self.verdicts[3] = self._verdict(3, subjects)

    def _dominance(self, child, parent_allow, parent_budget, parent_exp, machine):
        for ce in child.get("allow", []):
            ok = False
            for pe in parent_allow:
                if pe.get("verb") != ce.get("verb"):
                    continue
                if not contained(ce.get("resource", ""), pe.get("resource", "")):
                    continue
                pc, cc = pe.get("class"), ce.get("class")
                if pc and cc and CLASS_SEVERITY[cc] < CLASS_SEVERITY[pc]:
                    continue  # child widened (less severe than parent) — not a restriction
                pb, cb = set(pe.get("boundary", [])), set(ce.get("boundary", []))
                if not pb.issubset(cb):
                    continue  # boundaries must be added, never dropped
                ok = True
                break
            if not ok:
                self.add(3, child["id"], f"allow entry {ce.get('verb')} {ce.get('resource')} not dominated by parent")
        for msg in budget_le(child.get("budget", {}), parent_budget or {}):
            self.add(3, child["id"], f"budget dominance: {msg}")
        if parent_exp is not None and child.get("exp", 0) > parent_exp:
            self.add(3, child["id"], f"expiry dominance: child exp {child.get('exp')} > parent {parent_exp}")

    # -- rule 4 ------------------------------------------------------------
    def rule4(self):
        m = self._mandate()
        chain, _ = self._receipts_ordered()
        # outcome actuals supersede their paired gate estimates
        superseded = {r["gate"] for r in chain if r.get("gate")}
        running: dict[str, int] = {}
        subjects = 0
        for r in chain:
            if r["id"] in superseded and r.get("decision", {}).get("result") == "allowed":
                continue  # estimate replaced by its outcome receipt
            subjects += 1
            for axis, v in (r.get("cost") or {}).items():
                running[axis] = running.get(axis, 0) + v
            for anc in self._ancestors_of(r):
                for axis, cap in (anc.get("budget") or {}).items():
                    if running.get(axis, 0) > cap:
                        self.add(4, r["id"], f"prefix total {axis}={running[axis]} exceeds {anc['type']} cap {cap}")
        self.recomputed = running
        if m:
            for axis, cap in (m.get("budget") or {}).items():
                if running.get(axis, 0) > cap:
                    self.add(4, m["id"], f"final total {axis}={running.get(axis,0)} exceeds mandate {cap}")
        self.verdicts[4] = self._verdict(4, subjects)

    def _ancestors_of(self, receipt):
        out, seen, stack = [], set(), list(receipt.get("parents", []))
        while stack:
            pid = stack.pop()
            if pid in seen:
                continue
            seen.add(pid)
            o = self.by_id.get(pid)
            if not o:
                continue
            if o.get("type") in ("grant", "mandate"):
                out.append(o)
                stack.extend(o.get("parents", []))
        return out

    # -- rule 5 ------------------------------------------------------------
    def rule5(self):
        consumed: dict[str, int] = {}
        approvals = {}
        for r in self.by_type.get("receipt", []):
            if (r.get("act") or {}).get("verb") == "stop.approve":
                approvals[(r.get("act") or {}).get("resource")] = r
        subjects = 0
        chain, _ = self._receipts_ordered()
        for r in chain:
            if r.get("decision", {}).get("result") != "allowed":
                continue
            g = self._grant_for(r)
            verb = (r.get("act") or {}).get("verb", "")
            if verb.startswith("stop."):
                continue
            cls = governing_class(g.get("allow", []), verb, (r.get("act") or {}).get("resource", "")) if g else None
            if cls != "settle":
                continue
            if not r.get("gate"):
                subjects += 1
            sid = r.get("decision", {}).get("stop")
            if not sid:
                self.add(5, r["id"], "allowed settle-class act references no Stop")
                continue
            stop = self.by_id.get(sid)
            if not stop:
                self.add(5, r["id"], f"referenced Stop {sid[:18]}… does not resolve")
                continue
            if stop.get("act_digest") != (r.get("act") or {}).get("input_digest"):
                self.add(5, r["id"], "Stop act_digest does not equal the act's input_digest")
            if r.get("iat", 0) > stop.get("exp", 0):
                self.add(5, r["id"], "Stop had expired at act time")
            ap = approvals.get(sid)
            if not ap:
                self.add(5, r["id"], "no stop.approve receipt for the referenced Stop")
            else:
                if ap.get("actor") != stop.get("approver"):
                    self.add(5, ap["id"], "approval not attributable to the Stop's named approver")
                if stop.get("mode") == "signature":
                    entry = self.bundle["keys"].get(stop.get("approver"))
                    payload = (stop["id"] + stop["act_digest"]).encode()
                    if not entry or not ap.get("approval") or not verify_detached(ap["approval"], payload, entry["key"]):
                        self.add(5, ap["id"], "signature-mode approval does not verify over (stop.id ‖ act_digest)")
                elif stop.get("mode") == "token":
                    tok = ap.get("token", "")
                    if digest_bytes(tok.encode()) != stop.get("token_digest"):
                        self.add(5, ap["id"], "token-mode approval does not match the Stop's token digest")
            # A paired outcome Receipt records the same act as its gate
            # Receipt (§4.3.3); only the gate consumes the Stop.
            if not r.get("gate"):
                allowance = (stop.get("envelope") or {}).get("count", 1)
                consumed[sid] = consumed.get(sid, 0) + 1
                if consumed[sid] > allowance:
                    self.add(5, r["id"], f"Stop reused: {consumed[sid]} acts reference a Stop allowing {allowance}")
        self.verdicts[5] = self._verdict(5, subjects)

    # -- rule 6 ------------------------------------------------------------
    def rule6(self):
        m = self._mandate()
        sets = self.by_type.get("settlement", [])
        if not m or not sets:
            self.verdicts[6] = "n/a"
            return
        subjects = 0
        for s in sets:
            subjects += 1
            addressed = {c["code"] for c in s.get("criteria", [])}
            for crit in m.get("acceptance", []):
                if crit["code"] not in addressed:
                    self.add(6, s["id"], f"acceptance criterion '{crit['code']}' not addressed")
            indep = {c["code"] for c in m.get("acceptance", []) if c.get("independent")}
            for c in s.get("criteria", []):
                if c["code"] in indep:
                    if not c.get("validator_sig"):
                        self.add(6, s["id"], f"independent criterion '{c['code']}' carries no validator signature")
                    else:
                        entry = self.bundle["keys"].get(c.get("validator"))
                        body = {k: v for k, v in c.items() if k != "validator_sig"}
                        if not entry or not verify_detached(c["validator_sig"], canonicalize(body), entry["key"]):
                            self.add(6, s["id"], f"validator signature for '{c['code']}' does not verify")
                    actors = {r.get("actor") for r in self.by_type.get("receipt", [])}
                    if c.get("validator") in actors:
                        self.add(6, s["id"], f"independence: validator of '{c['code']}' also acted in this chain")
                if c.get("verdict") == "rejected" and s.get("state") == "settled":
                    self.add(6, s["id"], f"criterion '{c['code']}' rejected but state is settled")
            for axis, v in (s.get("totals") or {}).items():
                if getattr(self, "recomputed", {}).get(axis, 0) != v:
                    self.add(6, s["id"], f"totals.{axis}={v} does not match recomputation {getattr(self,'recomputed',{}).get(axis,0)}")
            for r in self.by_type.get("receipt", []):
                if r.get("iat", 0) > s.get("iat", 0):
                    self.add(6, r["id"], "receipt postdates the terminal settlement")
        self.verdicts[6] = self._verdict(6, subjects)

    # -- rule 7 ------------------------------------------------------------
    def rule7(self):
        subjects = 0
        for c in self.by_type.get("context", []):
            subjects += 1
            if c.get("used_tokens", 0) > c.get("budget_tokens", 0):
                self.add(7, c["id"], "used_tokens exceeds budget_tokens")
            for it in c.get("items", []):
                if it.get("trust") not in ("principal", "governed", "external"):
                    self.add(7, c["id"], f"item '{it.get('ref')}' carries no valid trust class")
                srcs = it.get("derived_from") or []
                if srcs and not it.get("declassified_by"):
                    lowest = "principal"
                    for s in srcs:
                        for other in self.by_type.get("context", []):
                            for oi in other.get("items", []):
                                if oi.get("digest") == s:
                                    order = ["external", "governed", "principal"]
                                    if order.index(oi["trust"]) < order.index(lowest):
                                        lowest = oi["trust"]
                    if it.get("trust") != lowest:
                        self.add(7, c["id"], f"taint: item '{it.get('ref')}' is {it.get('trust')} but its sources bottom out at {lowest}")
        for r in self.by_type.get("receipt", []):
            if (r.get("act") or {}).get("verb", "").startswith("model.") and r.get("decision", {}).get("result") == "allowed":
                subjects += 1
                if not r.get("context"):
                    self.add(7, r["id"], "model-driven allowed act references no Context Envelope")
        self.notes.append("rendering constraints are conformance-clause territory and are NOT claimed as chain-verified (§4.6 rule 7)")
        self.verdicts[7] = self._verdict(7, subjects)

    # -- rule 8 ------------------------------------------------------------
    def rule8(self):
        subjects = 0
        for o in self.objs:
            if "id" not in o:
                continue
            for p in o.get("parents", []):
                po = self.by_id.get(p)
                if po and o.get("iat", 0) < po.get("iat", 0):
                    self.add(8, o["id"], f"iat precedes parent {p[:18]}…")
                    subjects += 1
        chain, _ = self._receipts_ordered()
        atts = self.bundle.get("attestations", [])
        if not atts:
            self.verdicts[8] = "incomplete" if not self.findings else self._verdict(8, subjects or 1)
            self.notes.append("no head attestation supplied: chain reported INCOMPLETE, not clean (§4.6 Anchoring)")
            return
        head = chain[-1]["id"] if chain else None
        latest = max(atts, key=lambda a: a.get("iat", 0))
        if latest.get("head") != head:
            ids = [r["id"] for r in chain]
            if latest.get("head") not in ids:
                self.add(8, latest.get("head", "?"), "latest head attestation is not on the disclosed chain")
            else:
                self.notes.append("chain extends its latest attestation")
        self.verdicts[8] = self._verdict(8, max(subjects, 1))

    # -- driver ------------------------------------------------------------
    def _verdict(self, rule, subjects):
        fails = [f for f in self.findings if f.rule == rule and f.kind == "fail"]
        if subjects == 0:
            return "n/a"
        if fails:
            return "fail"
        return "pass"

    def run(self):
        # branches a conformant gate refused for an unknown critical extension
        for r in self.by_type.get("receipt", []):
            if r.get("decision", {}).get("reason") == "GRIP_DENY_UNKNOWN_CRIT":
                self.crit_refused.add(r["id"])
        for n in range(1, 9):
            getattr(self, f"rule{n}")()
        return self

    def report(self) -> dict:
        return {
            "grip": "1.0-draft.2",
            "objects": len([o for o in self.objs if "id" in o]),
            "verdicts": {f"rule{n}": {"name": RULES[n], "verdict": self.verdicts.get(n, "n/a")} for n in range(1, 9)},
            "findings": [{"rule": f.rule, "object": f.obj, "message": f.msg} for f in self.findings],
            "notes": self.notes,
            "unevaluable_critical_branches": sorted(self.crit_refused),
            "recomputed_totals": getattr(self, "recomputed", {}),
            "clean": all(v in ("pass", "n/a") for v in self.verdicts.values()),
        }


def main():
    ap = argparse.ArgumentParser(description="GRIP/1.0-draft.2 chain verifier (§4.6)")
    ap.add_argument("--bundle", required=True)
    ap.add_argument("--objects", required=True)
    ap.add_argument("--json", action="store_true")
    args = ap.parse_args()

    v = Verifier(load_objects(args.objects), load_bundle(args.bundle)).run()
    rep = v.report()
    if args.json:
        print(json.dumps(rep, indent=2))
        return 0 if rep["clean"] else 1

    print(f"grip-verify · GRIP/{rep['grip']} · {rep['objects']} objects\n")
    for n in range(1, 9):
        r = rep["verdicts"][f"rule{n}"]
        mark = {"pass": "PASS", "fail": "FAIL", "n/a": " n/a", "incomplete": "INCM"}[r["verdict"]]
        print(f"  [{mark}] rule {n}  {r['name']}")
    if rep["findings"]:
        print("\nFindings:")
        for f in rep["findings"]:
            print(f"  rule {f['rule']}  {f['object'][:22]}…  {f['message']}")
    if rep["notes"]:
        print("\nNotes:")
        for n_ in rep["notes"]:
            print(f"  · {n_}")
    if rep["recomputed_totals"]:
        print(f"\nRecomputed totals: {rep['recomputed_totals']}")
    print("\nRESULT:", "clean" if rep["clean"] else "NOT clean")
    return 0 if rep["clean"] else 1


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