#!/usr/bin/env python3
"""Independent verifier for the Modelometer provenance chain.

Runs with the STANDARD LIBRARY ONLY — no `mom` package, no third-party deps — so a stranger
can clone the repo (or just point at the public JSON) and reproduce every hash. It
re-implements the frozen canonical-JSON spec and the chain rule:

    chain_hash(run) = sha256( prev_hash  ++  canonical_json(payload) )

where payload = {"canonical_spec": "1", "run": {...}, "content_digest": <sha256>} and
content_digest = sha256(canonical_json(sorted [{id, digest}])) with digest = sha256 of each
published record. Genesis prev_hash = 64 zeros.

Usage:
    python scripts/verify_chain.py --api ./public/api/v0
    python scripts/verify_chain.py --api https://modelometer.fly.dev/api/v0
"""

from __future__ import annotations

import argparse
import hashlib
import json
import urllib.request
from typing import Any

GENESIS_PREV_HASH = "0" * 64


# ── frozen canonicalization (must match src/mom/canonical.py exactly) ────────────────────
def canonical_json(obj: Any) -> str:
    return json.dumps(obj, sort_keys=True, separators=(",", ":"), ensure_ascii=False, allow_nan=False)


def sha256_hex(data: bytes) -> str:
    return hashlib.sha256(data).hexdigest()


def hash_obj(obj: Any) -> str:
    return sha256_hex(canonical_json(obj).encode("utf-8"))


# ── fetch (file or URL) ─────────────────────────────────────────────────────────────────
def _load(api: str, name: str) -> Any:
    if api.startswith("http://") or api.startswith("https://"):
        with urllib.request.urlopen(api.rstrip("/") + "/" + name, timeout=30) as r:  # noqa: S310
            return json.loads(r.read().decode("utf-8"))
    with open(f"{api.rstrip('/')}/{name}", encoding="utf-8") as f:
        return json.load(f)


# ── verification ────────────────────────────────────────────────────────────────────────
def verify(api: str, check_records: bool = True) -> tuple[bool, list[str]]:
    msgs: list[str] = []
    chain = _load(api, "chain.json")
    runs = chain.get("runs", [])
    if not runs:
        return False, ["chain.json has no runs"]

    ok = True
    prev = GENESIS_PREV_HASH
    for entry in runs:
        run_id = entry["id"]
        # linkage
        if entry.get("prev_hash") != prev:
            ok = False
            msgs.append(f"[{run_id}] prev_hash linkage broken (expected {prev[:12]}…, got {str(entry.get('prev_hash'))[:12]}…)")

        # pull the published run doc for content + payload
        try:
            doc = _load(api, f"runs/{run_id}.json")
        except Exception as exc:  # noqa: BLE001
            ok = False
            msgs.append(f"[{run_id}] cannot load published run doc: {type(exc).__name__}")
            prev = entry.get("chain_hash") or prev
            continue

        # recompute content_digest from the published records
        if check_records:
            manifest = sorted(({"id": r["id"], "digest": hash_obj(r)} for r in doc.get("records", [])),
                              key=lambda x: x["id"])
            cdigest = hash_obj(manifest)
            if cdigest != doc.get("content_digest"):
                ok = False
                msgs.append(f"[{run_id}] content_digest mismatch (records tampered?)")
        else:
            cdigest = doc.get("content_digest")

        # recompute the chain hash
        payload = {"canonical_spec": doc.get("canonical_spec", "1"), "run": doc["run"], "content_digest": cdigest}
        recomputed = sha256_hex((prev + canonical_json(payload)).encode("utf-8"))
        published = doc.get("chain_hash")
        if recomputed != published or recomputed != entry.get("chain_hash"):
            ok = False
            msgs.append(f"[{run_id}] chain_hash mismatch (recomputed {recomputed[:12]}…, published {str(published)[:12]}…)")
        else:
            msgs.append(f"[{run_id}] OK  seq={entry.get('seq')}  {recomputed[:16]}…")
        prev = published or recomputed

    return ok, msgs


def main(argv: list[str] | None = None) -> int:
    ap = argparse.ArgumentParser(description="Verify the Modelometer hash chain from public JSON")
    ap.add_argument("--api", required=True, help="path or URL to /api/v0")
    ap.add_argument("--no-records", action="store_true", help="skip per-record content verification")
    args = ap.parse_args(argv)

    ok, msgs = verify(args.api, check_records=not args.no_records)
    for m in msgs:
        print(m)
    print("-" * 48)
    print("CHAIN VERIFICATION:", "PASS ✓" if ok else "FAIL ✗")
    return 0 if ok else 1


if __name__ == "__main__":
    raise SystemExit(main())
