#!/usr/bin/env python3
"""Check that a stored CAIN hosted decision record is exactly what the gateway signed.

No CAIN imports; needs Python 3.8+ and `cryptography`.

  python3 verify_decision_record.py record.json --key https://mcpgate.online/fabric/decision-signing-key
  python3 verify_decision_record.py record.json --key signing-key.json --self-test

record.json is one full row of the hosted Fabric's decision table, exactly as stored.

  DIGEST     the record's SHA-256 digest is recomputed HERE from its own fields (same field order and
             separator as the gateway) and must equal the stored digest
  KEY        the key id in the record equals sha256(public key)[:16] of the key given with --key,
             which should be fetched from a DIFFERENT site than the one serving the record
  SIGNATURE  the Ed25519 signature verifies over  "cain.fabric.decision-record.v1" 0x1F <digest hex>

--self-test also runs two negative controls on copies of the record and requires both to FAIL:
  1. the verdict is changed and the digest is left alone      -> DIGEST fails
  2. the verdict is changed and the digest is recomputed too  -> SIGNATURE fails
(the second is what a database writer without the signing key would do).

Not covered: someone holding the gateway host's root account holds both the key and the database.
"""
from __future__ import annotations

import base64
import copy
import hashlib
import json
import sys
import urllib.request
from pathlib import Path

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

DOMAIN = "cain.fabric.decision-record.v1"
SCHEMA_VERSION = 3


def load(src: str):
    if src.startswith("https://") or src.startswith("http://"):
        req = urllib.request.Request(src, headers={"User-Agent": "cain42-verify-decision-record"})
        with urllib.request.urlopen(req, timeout=60) as r:
            return json.loads(r.read())
    return json.loads(Path(src).read_text())


def digest(r) -> str:
    payload = "\x1f".join([
        r["decision_id"], r["tenant"], r.get("principal_id") or "", r.get("agent_id") or "",
        r.get("service") or "", r.get("path") or "", r["verdict"],
        "1" if r["blocked"] else "0", "1" if r["enforcing"] else "0",
        r["stages"], r["created_at"], r.get("chain_id") or "", r.get("parent_decision_id") or "",
        str(SCHEMA_VERSION), r.get("outcome") or "",
    ])
    return hashlib.sha256(payload.encode("utf-8")).hexdigest()


def check(r, key):
    raw = base64.b64decode(key["public_key_b64"])
    kid = hashlib.sha256(raw).hexdigest()[:16]
    out = []
    d = digest(r)
    out.append(("DIGEST", d == r.get("digest"), f"recomputed {d[:16]}, stored {str(r.get('digest'))[:16]}"))
    out.append(("KEY", kid == r.get("record_key_id"), f"record key id {r.get('record_key_id')}, --key id {kid}"))
    try:
        Ed25519PublicKey.from_public_bytes(raw).verify(base64.b64decode(r.get("record_signature") or ""),
                                                        f"{DOMAIN}\x1f{r.get('digest')}".encode("utf-8"))
        ok = True
    except Exception:  # noqa: BLE001
        ok = False
    out.append(("SIGNATURE", ok, f"Ed25519 over {DOMAIN} 0x1F <stored digest>"))
    return out


def main(argv) -> int:
    if len(argv) < 3 or argv[1] != "--key":
        print(__doc__); return 2
    r, key = load(argv[0]), load(argv[2])
    if int(r.get("schema_version") or 0) != SCHEMA_VERSION:
        print(f"record schema_version {r.get('schema_version')} is not {SCHEMA_VERSION}; this verifier does not cover it")
        return 2
    print(f"decision {r['decision_id']}  verdict {r['verdict']}  created {r['created_at']}")
    res = check(r, key)
    for name, ok, detail in res:
        print(f"[{'PASS' if ok else 'FAIL'}] {name:9s} {detail}")
    good = all(ok for _, ok, _ in res)
    if "--self-test" in argv:
        other = "ALLOW" if r["verdict"] != "ALLOW" else "DENY"
        t1 = copy.deepcopy(r); t1["verdict"] = other
        t2 = copy.deepcopy(t1); t2["digest"] = digest(t2)
        f1 = {n: ok for n, ok, _ in check(t1, key)}
        f2 = {n: ok for n, ok, _ in check(t2, key)}
        n1, n2 = not f1["DIGEST"], f2["DIGEST"] and not f2["SIGNATURE"]
        print(f"[{'PASS' if n1 else 'FAIL'}] NEGATIVE1 verdict -> {other}, digest untouched: DIGEST {'fails' if n1 else 'DID NOT FAIL'}")
        print(f"[{'PASS' if n2 else 'FAIL'}] NEGATIVE2 verdict -> {other}, digest recomputed: SIGNATURE {'fails' if n2 else 'DID NOT FAIL'}")
        good = good and n1 and n2
    print("\nVERIFIED: this record is exactly what the gateway signed" if good else "\nINVALID")
    return 0 if good else 1


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