"""CAIN 37.0 Phase 4/10 -- Independent Evidence Chain Verifier (zero-import).

Same discipline as cain_35_independent_verifier.py (CAIN 35, Phase 10): this
module NEVER imports cain_evidence_fabric_37 (the production module that
mints and chains evidence). Every canonicalization rule, hash computation,
chain-link recomputation, and Merkle root/inclusion-proof check below is
reimplemented independently from the schema documented in
cain_evidence_fabric_37.py's docstrings -- so a bug in the production chain
code cannot be silently inherited by whatever is supposed to be checking it.

Verified by `test_independent_verifier_shares_no_classes_with_production_module`
in tests/distributed/test_cain37_evidence_fabric.py (ast-parses this file's
own imports).

Usage:
    python3 cain_37_independent_verifier.py verify-chain evidence_chain_37.jsonl
    python3 cain_37_independent_verifier.py verify-checkpoint evidence_chain_37.jsonl evidence_checkpoints_37.jsonl
"""

from __future__ import annotations

import hashlib
import json
import sys
from typing import Any, Dict, List, Tuple


class Verdict:
    VALID = "VALID"
    CHAIN_INTEGRITY_FAILURE = "CHAIN_INTEGRITY_FAILURE"
    CONSISTENCY_FAILURE = "CONSISTENCY_FAILURE"
    VERIFICATION_FAILURE = "VERIFICATION_FAILURE"
    EMPTY = "EMPTY_CHAIN"


def _canonical_bytes(obj: Any) -> bytes:
    """Independent reimplementation of the canonicalization rule (sorted
    keys, no whitespace, ensure_ascii=False) -- written from the spec in
    cain_evidence_fabric_37.py's docstring, not imported from it."""
    return json.dumps(obj, sort_keys=True, separators=(",", ":"), ensure_ascii=False, default=str).encode("utf-8")


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


def _content_hash(evidence: Dict[str, Any]) -> str:
    """Recomputes the same content hash the production module signs/chains,
    by stripping the same three signature-related fields and canonicalizing
    the rest -- independently re-derived, not copy-pasted from a shared
    helper."""
    stripped = {k: v for k, v in evidence.items() if k not in ("signature", "algorithm", "signer_node_id")}
    return _sha256_hex(_canonical_bytes(stripped))


def _merkle_root(leaf_hashes: List[str]) -> str:
    if not leaf_hashes:
        return _sha256_hex(b"")
    level = list(leaf_hashes)
    while len(level) > 1:
        if len(level) % 2 == 1:
            level.append(level[-1])
        nxt = []
        for i in range(0, len(level), 2):
            nxt.append(_sha256_hex(bytes.fromhex(level[i]) + bytes.fromhex(level[i + 1])))
        level = nxt
    return level[0]


def verify_chain(records: List[Dict[str, Any]]) -> Dict[str, Any]:
    """Independently walks the whole chain, recomputing:
    1. each entry's content_hash from its own evidence payload
    2. each entry's chain_hash from prev_chain_hash + recomputed content_hash
    3. that chain_index and prev_chain_hash linkage are strictly sequential

    Any tamper (field mutation, reordering, deletion, rewrite of an old
    entry) is expected to surface here as a hash mismatch.
    """
    if not records:
        return {"verdict": Verdict.EMPTY, "checked": 0}

    genesis = "0" * 64
    prev_chain_hash = genesis
    for i, rec in enumerate(records):
        recomputed_content_hash = _content_hash(rec["evidence"])
        if recomputed_content_hash != rec["content_hash"]:
            return {
                "verdict": Verdict.VERIFICATION_FAILURE,
                "reason": f"content_hash mismatch at index {i}: evidence payload does not hash to its recorded content_hash (tampered field)",
                "index": i,
                "expected": rec["content_hash"],
                "recomputed": recomputed_content_hash,
            }
        if rec["prev_chain_hash"] != prev_chain_hash:
            return {
                "verdict": Verdict.CHAIN_INTEGRITY_FAILURE,
                "reason": f"prev_chain_hash mismatch at index {i}: this entry's recorded predecessor does not match the actual prior entry's chain_hash (reordering, deletion, or insertion detected)",
                "index": i,
                "expected_prev": prev_chain_hash,
                "recorded_prev": rec["prev_chain_hash"],
            }
        recomputed_chain_hash = _sha256_hex(bytes.fromhex(rec["prev_chain_hash"]) + bytes.fromhex(recomputed_content_hash))
        if recomputed_chain_hash != rec["chain_hash"]:
            return {
                "verdict": Verdict.CHAIN_INTEGRITY_FAILURE,
                "reason": f"chain_hash mismatch at index {i}: recomputed hash chain diverges from the recorded one",
                "index": i,
                "expected": rec["chain_hash"],
                "recomputed": recomputed_chain_hash,
            }
        if rec.get("chain_index") != i:
            return {
                "verdict": Verdict.CHAIN_INTEGRITY_FAILURE,
                "reason": f"chain_index out of sequence at position {i} (recorded {rec.get('chain_index')}) -- entries missing or reordered",
                "index": i,
            }
        prev_chain_hash = recomputed_chain_hash

    return {
        "verdict": Verdict.VALID,
        "checked": len(records),
        "final_chain_hash": prev_chain_hash,
    }


def verify_checkpoint(records: List[Dict[str, Any]], checkpoint: Dict[str, Any]) -> Dict[str, Any]:
    """Independently recomputes the Merkle root over the exact slice of
    records the checkpoint claims to cover, and compares against the
    checkpoint's recorded root -- a missing/reordered/tampered record in
    that slice changes the recomputed root."""
    since = checkpoint["since_index"]
    through = checkpoint["through_index"]
    batch = records[since:through + 1]
    if len(batch) != checkpoint["leaf_count"]:
        return {
            "verdict": Verdict.CONSISTENCY_FAILURE,
            "reason": f"expected {checkpoint['leaf_count']} entries in range [{since},{through}], found {len(batch)} -- an entry was removed",
        }
    leaf_hashes = [_content_hash(e["evidence"]) for e in batch]
    recomputed_root = _merkle_root(leaf_hashes)
    if recomputed_root != checkpoint["merkle_root"]:
        return {
            "verdict": Verdict.CONSISTENCY_FAILURE,
            "reason": "recomputed Merkle root does not match the sealed checkpoint's root -- an entry in this range was tampered, removed, or reordered after sealing",
            "expected": checkpoint["merkle_root"],
            "recomputed": recomputed_root,
        }
    return {"verdict": Verdict.VALID, "checkpoint_id": checkpoint.get("checkpoint_id"), "leaf_count": len(batch), "merkle_root": recomputed_root}


def verify_inclusion_proof(leaf_hash: str, path: List[Dict[str, str]], expected_root: str) -> Dict[str, Any]:
    cur = leaf_hash
    for step in path:
        sib = step["hash"]
        if step["side"] == "left":
            cur = _sha256_hex(bytes.fromhex(sib) + bytes.fromhex(cur))
        else:
            cur = _sha256_hex(bytes.fromhex(cur) + bytes.fromhex(sib))
    if cur != expected_root:
        return {"verdict": Verdict.VERIFICATION_FAILURE, "reason": "recomputed root from inclusion proof does not match expected root", "recomputed": cur, "expected": expected_root}
    return {"verdict": Verdict.VALID, "recomputed_root": cur}


def _load_jsonl(path: str) -> List[Dict[str, Any]]:
    out = []
    with open(path) as f:
        for line in f:
            line = line.strip()
            if line:
                out.append(json.loads(line))
    return out


def main(argv: List[str]) -> int:
    if len(argv) < 2:
        print(__doc__)
        return 2
    cmd = argv[0]
    if cmd == "verify-chain":
        records = _load_jsonl(argv[1])
        result = verify_chain(records)
        print(json.dumps(result, indent=2))
        return 0 if result["verdict"] == Verdict.VALID else 1
    if cmd == "verify-checkpoint":
        records = _load_jsonl(argv[1])
        checkpoints = _load_jsonl(argv[2])
        results = [verify_checkpoint(records, ckpt) for ckpt in checkpoints]
        print(json.dumps(results, indent=2))
        return 0 if all(r["verdict"] == Verdict.VALID for r in results) else 1
    print(f"unknown command: {cmd}")
    return 2


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