#!/usr/bin/env python3
"""Offline verifier for CAIN SCITT receipts and log consistency. Imports nothing from CAIN.

Needs only Python 3.9+ and the `cryptography` package. Written independently of cain_scitt_clearinghouse so an
auditor does not have to trust CAIN's own code: it re-implements RFC 9162 inclusion/consistency verification,
Ed25519 checkpoint verification and SD-JWT disclosure checks.

Usage:
  scitt_verify_offline.py keys.json receipt.json [--body body.json]
      keys.json: {"keys": {key_id: base64 raw Ed25519 public key}} (GET /fabric/scitt/keys)
      receipt.json: a receipt from POST /fabric/scitt/register-statement or GET /fabric/scitt/receipts/{id}
  scitt_verify_offline.py keys.json --consistency proof.json
      proof.json: GET /fabric/scitt/consistency?first=N[&second=M] (checks both checkpoints' signatures too)

Exit status 0 = verified, 1 = not verified, 2 = usage error.
"""
import base64
import hashlib
import json
import sys

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


def cj(obj):
    return json.dumps(obj, sort_keys=True, separators=(",", ":"), ensure_ascii=False).encode("utf-8")


def node(a, b):
    return hashlib.sha256(b"\x01" + a + b).digest()


def sig_ok(keys, key_id, msg, sig_b64):
    pub = keys.get(key_id)
    if not pub:
        return False
    try:
        Ed25519PublicKey.from_public_bytes(base64.b64decode(pub)).verify(base64.b64decode(sig_b64), msg)
        return True
    except (InvalidSignature, ValueError):
        return False


def checkpoint_ok(keys, cp):
    msg = cj({k: cp[k] for k in ("log_id", "tree_size", "root_hash", "timestamp", "key_id")})
    return sig_ok(keys, cp["key_id"], msg, cp["signature_b64"])


def inclusion_ok(leaf, index, size, path, root):
    if not 0 <= index < size:
        return False
    fn, sn, r = index, size - 1, leaf
    for p in path:
        if sn == 0:
            return False
        if fn & 1 or fn == sn:
            r = node(p, r)
            while not fn & 1 and fn != 0:
                fn >>= 1
                sn >>= 1
        else:
            r = node(r, p)
        fn >>= 1
        sn >>= 1
    return sn == 0 and r == root


def consistency_ok(first, second, r1, r2, proof):
    if first < 1 or first > second:
        return False
    if first == second:
        return not proof and r1 == r2
    if not proof:
        return False
    if first & (first - 1) == 0:
        proof = [r1] + proof
    fn, sn = first - 1, second - 1
    while fn & 1:
        fn >>= 1
        sn >>= 1
    fr = sr = proof[0]
    for c in proof[1:]:
        if sn == 0:
            return False
        if fn & 1 or fn == sn:
            fr, sr = node(c, fr), node(c, sr)
            while not fn & 1 and fn != 0:
                fn >>= 1
                sn >>= 1
        else:
            sr = node(sr, c)
        fn >>= 1
        sn >>= 1
    return fr == r1 and sr == r2 and sn == 0


def verify_receipt(keys, rc, body=None):
    cp, entry = rc["checkpoint"], rc["entry"]
    if not checkpoint_ok(keys, cp):
        return False, "checkpoint signature invalid or key not pinned"
    if entry["log_id"] != cp["log_id"] or int(entry["index"]) != int(rc["leaf_index"]):
        return False, "entry does not match the checkpoint's log or the proven index"
    leaf = hashlib.sha256(b"\x00" + cj(entry)).digest()
    if not inclusion_ok(leaf, int(rc["leaf_index"]), int(cp["tree_size"]),
                        [bytes.fromhex(p) for p in rc["inclusion_proof"]], bytes.fromhex(cp["root_hash"])):
        return False, "inclusion proof does not reach the signed root"
    if body is not None and hashlib.sha256(cj(body)).hexdigest() != entry["statement_digest"]:
        return False, "body does not match the logged digest"
    return True, f"entry {entry['statement_id']} is in log {cp['log_id']} at tree size {cp['tree_size']}"


def verify_consistency_doc(keys, doc):
    for name in ("first_checkpoint", "second_checkpoint"):
        cp = doc.get(name)
        if not cp or not checkpoint_ok(keys, cp):
            return False, f"{name} missing or its signature is invalid"
    a, b = doc["first_checkpoint"], doc["second_checkpoint"]
    if a["log_id"] != b["log_id"] or (a["tree_size"], b["tree_size"]) != (doc["first"], doc["second"]):
        return False, "checkpoints do not match the proof's log and sizes"
    if not consistency_ok(doc["first"], doc["second"], bytes.fromhex(a["root_hash"]), bytes.fromhex(b["root_hash"]),
                          [bytes.fromhex(p) for p in doc["proof"]]):
        return False, "NOT consistent: the log was rewritten or forked between the two checkpoints"
    return True, f"log {a['log_id']} only appended between sizes {doc['first']} and {doc['second']}"


def main(argv):
    if len(argv) < 3:
        print(__doc__)
        return 2
    keys = json.load(open(argv[1]))
    keys = keys.get("keys", keys)
    if argv[2] == "--consistency":
        ok, msg = verify_consistency_doc(keys, json.load(open(argv[3])))
    else:
        body = json.load(open(argv[argv.index("--body") + 1])) if "--body" in argv else None
        ok, msg = verify_receipt(keys, json.load(open(argv[2])), body)
    print(("VERIFIED: " if ok else "NOT VERIFIED: ") + msg)
    return 0 if ok else 1


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