"""Recompute every vector in data/ independently of the reference node, and check the results.

Run from the bundle root:

    pip install -r env/requirements.txt
    python code/check_vectors.py            # compare with the declared results/ and exit 1 on any difference
    python code/check_vectors.py --write    # regenerate results/ instead

Exits non-zero if any vector mismatches, any invalid case is accepted, or (without --write)
any recomputed result differs from the declared one in results/.
"""

import base64
import hashlib
import json
import math
import pathlib
import re
import sys
import unicodedata

from cryptography.exceptions import InvalidSignature
from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PrivateKey, Ed25519PublicKey
from cryptography.hazmat.primitives.asymmetric.mldsa import MLDSA44PrivateKey, MLDSA44PublicKey

ROOT = pathlib.Path(__file__).resolve().parent.parent
DATA = ROOT / "data"
RESULTS = ROOT / "results"


# --- RFC 8785 canonical JSON -------------------------------------------------------------


def es_number(x):
    """Format a number the way ECMAScript's Number.prototype.toString does (RFC 8785, 3.2.2.3)."""
    x = float(x)
    if math.isnan(x) or math.isinf(x):
        raise ValueError("RFC 8785 has no representation for NaN or infinity")
    if x == 0:
        return "0"
    if x < 0:
        return "-" + es_number(-x)
    mantissa, _, exponent = repr(x).partition("e")  # repr gives the shortest round-trip digits
    whole, _, fraction = mantissa.partition(".")
    raw = whole + fraction
    digits = raw.lstrip("0")
    # The value is 0.<digits> times 10**n.
    n = len(whole) + int(exponent or 0) - (len(raw) - len(digits))
    digits = digits.rstrip("0")
    k = len(digits)
    if k <= n <= 21:
        return digits + "0" * (n - k)
    if 0 < n <= 21:
        return digits[:n] + "." + digits[n:]
    if -6 < n <= 0:
        return "0." + "0" * -n + digits
    e = n - 1
    return digits[0] + ("." + digits[1:] if k > 1 else "") + ("e+" if e >= 0 else "e-") + str(abs(e))


def canonical(value):
    """RFC 8785: sorted keys (by UTF-16 code units), no whitespace, ECMAScript numbers."""
    if value is None or isinstance(value, bool):
        return json.dumps(value)
    if isinstance(value, (int, float)):
        return es_number(value)
    if isinstance(value, str):
        return json.dumps(value, ensure_ascii=False)
    if isinstance(value, list):
        return "[" + ",".join(canonical(item) for item in value) + "]"
    if isinstance(value, dict):
        keys = sorted(value, key=lambda key: key.encode("utf-16-be"))
        return "{" + ",".join(json.dumps(k, ensure_ascii=False) + ":" + canonical(value[k]) for k in keys) + "}"
    raise TypeError(f"No JSON representation for {type(value).__name__}")


def naive(value):
    """What many Python agents would write: correct except for how numbers are formatted."""
    return json.dumps(value, sort_keys=True, separators=(",", ":"), ensure_ascii=False)


class Invalid(Exception):
    pass


def strict_loads(text):
    """json.loads, rejecting what I-JSON (RFC 7493) forbids."""

    def no_duplicates(pairs):
        names = [name for name, _ in pairs]
        if len(set(names)) != len(names):
            raise Invalid("duplicate property name")
        return dict(pairs)

    def check(value):
        if isinstance(value, float) and not math.isfinite(value):
            raise Invalid("number outside binary64")
        if isinstance(value, int) and not isinstance(value, bool):
            try:
                float(value)
            except OverflowError:
                raise Invalid("number outside binary64") from None
        if isinstance(value, str):
            try:
                value.encode("utf-8")
            except UnicodeEncodeError:
                raise Invalid("lone surrogate") from None
        if isinstance(value, list):
            for item in value:
                check(item)
        if isinstance(value, dict):
            for name, item in value.items():
                check(name)
                check(item)

    value = json.loads(text, object_pairs_hook=no_duplicates)
    check(value)
    return value


def sha256_hex(data):
    return hashlib.sha256(data if isinstance(data, bytes) else data.encode("utf-8")).hexdigest()


# --- Claim IDs ----------------------------------------------------------------------------


class Unresolved(Exception):
    """A claim names a result the bundle doesn't declare, so it has no ID."""


def results_named(claim):
    """Each result the claim's evidence names, once, in the order they first appear."""
    names = []
    for item in claim["evidence"]:
        if "result" in item and item["result"] not in names:
            names.append(item["result"])
    return names


def claim_ids(claims, verification_inputs, results, serialize=canonical):
    """A claim with evidence binds the verification inputs and, when its evidence names
    results, a results object mapping each name to its declared value."""
    by_local_id = {claim["local_id"]: claim for claim in claims}
    ids = {}

    def resolve(claim):
        if claim["local_id"] not in ids:
            depends_on = sorted(
                {dep if dep.startswith("claim:") else resolve(by_local_id[dep]) for dep in claim["depends_on"]}
            )
            content = {key: value for key, value in claim.items() if key != "local_id"}
            content["depends_on"] = depends_on
            if claim["evidence"]:
                content["verification_inputs"] = verification_inputs
                names = results_named(claim)
                if names:
                    missing = [name for name in names if name not in results]
                    if missing:
                        raise Unresolved(f"no declared result {missing[0]}")
                    content["results"] = {name: results[name] for name in names}
            ids[claim["local_id"]] = "claim:" + sha256_hex(serialize(content))
        return ids[claim["local_id"]]

    for claim in claims:
        resolve(claim)
    return ids


CLAIM_TYPES = {"empirical", "theoretical", "methodological", "replication", "negative_result", "resource"}
LOCAL_ID = re.compile(r"^[A-Za-z][A-Za-z0-9_-]*$")
GLOBAL_ID = re.compile(r"^claim:[0-9a-f]{64}$")
REQUIRED = {"local_id", "type", "core", "statement", "evidence", "depends_on", "confidence"}
# A result name is a file under results/ (no dots or slashes) and at least one key.
RESULT_NAME = re.compile(r"[^./]+(\.[^.]+)+")
PROOF_CHECKERS = {"lean4", "rocq"}
# The kinds of evidence: the fields each must have, the ones it may have, and the directory
# its file must sit under.
EVIDENCE_KINDS = [
    ({"result", "produced_by"}, {"tolerance"}, ("produced_by", "code/")),
    ({"result", "measured"}, {"tolerance"}, ("measured", "data/")),
    ({"proof", "theorem", "checker"}, set(), ("proof", "proofs/")),
]


def is_number(value):
    return isinstance(value, (int, float)) and not isinstance(value, bool)


def is_text(value):
    return isinstance(value, str) and value != ""


def validate_claims(claims):
    """The claims.json rules from /llms.txt; raises Invalid on the first problem."""
    if not isinstance(claims, list) or not 1 <= len(claims) <= 30:
        raise Invalid("claims.json must be a list of 1 to 30 claims")
    for claim in claims:
        if not isinstance(claim, dict) or not REQUIRED <= set(claim) <= REQUIRED | {"falsified_if"}:
            raise Invalid("missing or undefined fields")
        if not (isinstance(claim["local_id"], str) and LOCAL_ID.fullmatch(claim["local_id"])):
            raise Invalid("bad local_id")
        if claim["type"] not in CLAIM_TYPES or not isinstance(claim["core"], bool):
            raise Invalid("bad type or core")
        if not is_text(claim["statement"]):
            raise Invalid("empty statement")
        if not isinstance(claim["evidence"], list) or not isinstance(claim["depends_on"], list):
            raise Invalid("evidence and depends_on must be lists")
        if not (is_number(claim["confidence"]) and 0 <= claim["confidence"] <= 1):
            raise Invalid("confidence outside 0 to 1")
        if "falsified_if" in claim and not is_text(claim["falsified_if"]):
            raise Invalid("empty falsified_if")
        for item in claim["evidence"]:
            validate_evidence(item)
        if claim["type"] == "empirical" and not claim["evidence"]:
            raise Invalid("empirical claim without evidence")
        for dep in claim["depends_on"]:
            if not (isinstance(dep, str) and (GLOBAL_ID.fullmatch(dep) or LOCAL_ID.fullmatch(dep))):
                raise Invalid("malformed dependency")
    local_ids = [claim["local_id"] for claim in claims]
    if len(set(local_ids)) != len(local_ids):
        raise Invalid("duplicate local_id")
    graph = {claim["local_id"]: [d for d in claim["depends_on"] if not d.startswith("claim:")] for claim in claims}
    state = {}

    def visit(node):
        if state.get(node) == "done":
            return
        if state.get(node) == "visiting":
            raise Invalid("dependency cycle")
        state[node] = "visiting"
        for dep in graph[node]:
            if dep not in graph:
                raise Invalid("unknown local dependency")
            visit(dep)
        state[node] = "done"

    for node in graph:
        visit(node)


def utf16_length(text):
    return len(text.encode("utf-16-le")) // 2


def validate_evidence(item):
    """One evidence item: exactly one of the three kinds, with its file in the right place."""
    if not isinstance(item, dict):
        raise Invalid("evidence items must be objects")
    kinds = [kind for kind in EVIDENCE_KINDS if kind[0] <= set(item) <= kind[0] | kind[1]]
    if len(kinds) != 1:
        raise Invalid("an evidence item must be a computation, a measurement, or a proof")
    _, _, (field, directory) = kinds[0]
    location = item[field]
    if not (isinstance(location, str) and location.startswith(directory) and len(location) > len(directory)):
        raise Invalid(f"{field} must name a file under {directory}")
    if "result" in item and not (isinstance(item["result"], str) and RESULT_NAME.fullmatch(item["result"])):
        raise Invalid("a result is named like R3.loss_delta")
    if "tolerance" in item and not (is_number(item["tolerance"]) and item["tolerance"] >= 0):
        raise Invalid("bad tolerance")
    if "checker" in item:
        if item["checker"] not in PROOF_CHECKERS:
            raise Invalid("unknown proof checker")
        if not (is_text(item["theorem"]) and utf16_length(item["theorem"]) <= 500):
            raise Invalid("a theorem is named by 1 to 500 characters")


def check_claim_ids():
    vectors = json.loads((DATA / "claim-id-vectors.json").read_text())
    mismatches = naive_mismatches = total = 0
    for case in vectors["cases"]:
        claims = strict_loads(case["claims_json"])
        validate_claims(claims)
        ids = claim_ids(claims, case["verification_inputs"], case["results"])
        naive_ids = claim_ids(claims, case["verification_inputs"], case["results"], serialize=naive)
        mismatches += set(ids) != set(case["expected"])
        for local_id, expected in case["expected"].items():
            total += 1
            mismatches += ids.get(local_id) != expected
            naive_mismatches += naive_ids.get(local_id) != expected
    accepted = 0
    for case in vectors["invalid"]:
        try:
            validate_claims(strict_loads(case["claims_json"]))
            accepted += 1
        except (Invalid, ValueError):
            pass
    unresolved_accepted = 0
    for case in vectors["unresolved"]:
        claims = strict_loads(case["claims_json"])
        validate_claims(claims)  # valid claims; only a named result is missing
        try:
            claim_ids(claims, case["verification_inputs"], case["results"])
            unresolved_accepted += 1
        except Unresolved:
            pass
    return {
        "cases": len(vectors["cases"]),
        "claims": total,
        "mismatches": mismatches,
        "invalid_cases": len(vectors["invalid"]),
        "invalid_accepted": accepted,
        "unresolved_cases": len(vectors["unresolved"]),
        "unresolved_accepted": unresolved_accepted,
    }, naive_mismatches


# --- Bundles ------------------------------------------------------------------------------

BUNDLE_DIRECTORIES = ("code/", "env/", "data/", "results/", "proofs/")
# Declared results are not verification inputs: each claim binds the values it names.
INPUT_DIRECTORIES = ("code/", "env/", "data/", "proofs/")
TOP_LEVEL_FILES = {"manifest.json", "paper.md", "claims.json", "references.json", "embeddings.json", "provenance.json", "signature"}


def check_paths(paths):
    """The bundle path rules from /llms.txt; raises Invalid on the first problem."""
    folded = {}
    for path in paths:
        segments = path.split("/")
        if any(segment in ("", ".", "..") for segment in segments):
            raise Invalid("not a normalized relative path")
        if "\\" in path or any(ord(char) < 0x20 or ord(char) == 0x7F for char in path):
            raise Invalid("backslash or control character")
        if not unicodedata.is_normalized("NFC", path):
            raise Invalid("not NFC")
        if not (path in TOP_LEVEL_FILES if len(segments) == 1 else path.startswith(BUNDLE_DIRECTORIES)):
            raise Invalid("outside the layout")
        if path.lower() in folded:
            raise Invalid("differ only in case")
        folded[path.lower()] = path
    for path in folded:
        parts = path.split("/")
        if any("/".join(parts[:depth]) in folded for depth in range(1, len(parts))):
            raise Invalid("both a file and a directory")


def check_bundles():
    cases = json.loads((DATA / "bundle-vectors.json").read_text())["cases"]
    mismatches = 0
    for case in cases:
        check_paths(case["files"])
        files = {path: base64.b64decode(content) for path, content in case["files"].items()}
        digests = {path: "sha256:" + sha256_hex(content) for path, content in files.items()}
        bundle = {path: digest for path, digest in digests.items() if path != "signature"}
        inputs = {path: digest for path, digest in digests.items() if path.startswith(INPUT_DIRECTORIES)}
        expected = case["expected"]
        mismatches += digests != expected["files"]
        mismatches += "sha256:" + sha256_hex(canonical(bundle)) != expected["bundle"]
        mismatches += "sha256:" + sha256_hex(canonical(inputs)) != expected["verification_inputs"]
    invalid = json.loads((DATA / "bundle-vectors.json").read_text())["invalid"]
    accepted = 0
    for case in invalid:
        try:
            check_paths(case["paths"])
            accepted += 1
        except Invalid:
            pass
    return {"cases": len(cases), "mismatches": mismatches, "invalid_cases": len(invalid), "invalid_accepted": accepted}


# --- The log (RFC 9162) -------------------------------------------------------------------


def leaf_hash(data):
    return hashlib.sha256(b"\x00" + data).digest()


def node_hash(left, right):
    return hashlib.sha256(b"\x01" + left + right).digest()


def split(n):
    k = 1
    while k * 2 < n:
        k *= 2
    return k


EMPTY_ROOT = hashlib.sha256(b"").digest()


def mth(leaves):
    if not leaves:
        return EMPTY_ROOT
    if len(leaves) == 1:
        return leaves[0]
    k = split(len(leaves))
    return node_hash(mth(leaves[:k]), mth(leaves[k:]))


def path(m, leaves):
    if len(leaves) == 1:
        return []
    k = split(len(leaves))
    if m < k:
        return path(m, leaves[:k]) + [mth(leaves[k:])]
    return path(m - k, leaves[k:]) + [mth(leaves[:k])]


def subproof(m, leaves, complete):
    if m == len(leaves):
        return [] if complete else [mth(leaves)]
    k = split(len(leaves))
    if m <= k:
        return subproof(m, leaves[:k], complete) + [mth(leaves[k:])]
    return subproof(m - k, leaves[k:], False) + [mth(leaves[:k])]


def is_position(n):
    return isinstance(n, int) and not isinstance(n, bool) and n >= 0


def is_hash(value):
    return isinstance(value, bytes) and len(value) == 32


def verify_inclusion(index, size, leaf, proof, root):
    if not (is_position(index) and is_position(size) and index < size):
        return False
    if not all(is_hash(h) for h in [leaf, root, *proof]):
        return False
    fn, sn, r = index, size - 1, leaf
    for p in proof:
        if sn == 0:
            return False
        if fn % 2 == 1 or fn == sn:
            r = node_hash(p, r)
            while fn % 2 == 0 and fn != 0:
                fn, sn = fn // 2, sn // 2
        else:
            r = node_hash(r, p)
        fn, sn = fn // 2, sn // 2
    return sn == 0 and r == root


def verify_consistency(first, second, first_root, second_root, proof):
    if not (is_position(first) and is_position(second) and first <= second):
        return False
    if not all(is_hash(h) for h in [first_root, second_root, *proof]):
        return False
    if second == 0:
        return not proof and first_root == second_root == EMPTY_ROOT
    if first == second:
        return not proof and first_root == second_root
    if first == 0:
        return not proof and first_root == EMPTY_ROOT
    if not proof:
        return False
    path_ = ([first_root] if first & (first - 1) == 0 else []) + list(proof)
    fn, sn = first - 1, second - 1
    while fn % 2 == 1:
        fn, sn = fn // 2, sn // 2
    fr = sr = path_[0]
    for c in path_[1:]:
        if sn == 0:
            return False
        if fn % 2 == 1 or fn == sn:
            fr, sr = node_hash(c, fr), node_hash(c, sr)
            while fn % 2 == 0 and fn != 0:
                fn, sn = fn // 2, sn // 2
        else:
            sr = node_hash(sr, c)
        fn, sn = fn // 2, sn // 2
    return sn == 0 and fr == first_root and sr == second_root


def check_log():
    vectors = json.loads((DATA / "log-vectors.json").read_text())
    leaves = [leaf_hash(bytes.fromhex(leaf)) for leaf in vectors["leaves"]]
    roots = {tree["size"]: bytes.fromhex(tree["root"]) for tree in vectors["trees"]}
    mismatches = sum(mth(leaves[:size]) != root for size, root in roots.items())
    for item in vectors["inclusion"]:
        index, size = item["index"], item["size"]
        proof = [bytes.fromhex(h) for h in item["proof"]]
        mismatches += proof != path(index, leaves[:size])
        mismatches += not verify_inclusion(index, size, leaves[index], proof, roots[size])
    for item in vectors["consistency"]:
        first, second = item["first"], item["second"]
        proof = [bytes.fromhex(h) for h in item["proof"]]
        expected = [] if first == second else subproof(first, leaves[:second], True)
        mismatches += proof != expected
        mismatches += not verify_consistency(first, second, roots[first], roots[second], proof)
    # Each example leaf holds its entry as signed, with every signature replaced by its digest,
    # and each signed entry verifies against the key the log names it by: an operator's key
    # entry is signed by its own key, and later entries by the key that entry registered. A key
    # entry's leaf names the operator by the ID that key makes.
    keys = {}
    for example in vectors["leaf_examples"]:
        signed, leaf = example["signed_entry"], example["leaf"]
        mismatches += canonical(leaf["entry"]) != canonical(detach_signatures(signed))
        mismatches += leaf_hash(canonical(leaf).encode("utf-8")).hex() != example["leaf_hash"]
        if signed["type"] == "key":
            mismatches += leaf["operator"] != operator_id(signed["key"])
            keys[leaf["operator"]] = signed["key"]
        signer = keys.get(leaf["operator"])
        mismatches += signer is None or not verify_hybrid(signed["sig"], signing_payload(signed), signer)
    mismatches += mth([]).hex() != vectors["empty_root"]
    unhex = lambda hashes: [bytes.fromhex(h) for h in hashes]
    accepted = sum(
        verify_inclusion(c["index"], c["size"], bytes.fromhex(c["leaf_hash"]), unhex(c["proof"]), bytes.fromhex(c["root"]))
        for c in vectors["invalid_inclusion"]
    ) + sum(
        verify_consistency(
            c["first"], c["second"], bytes.fromhex(c["first_root"]), bytes.fromhex(c["second_root"]), unhex(c["proof"])
        )
        for c in vectors["invalid_consistency"]
    )
    return {
        "tree_sizes": len(roots),
        "inclusion_proofs": len(vectors["inclusion"]),
        "consistency_proofs": len(vectors["consistency"]),
        "leaf_examples": len(vectors["leaf_examples"]),
        "mismatches": mismatches,
        "invalid_proofs": len(vectors["invalid_inclusion"]) + len(vectors["invalid_consistency"]),
        "invalid_accepted": accepted,
    }


# --- Signatures ---------------------------------------------------------------------------

# A key or signature is "ed25519-ml-dsa-44:" and lowercase hex: the Ed25519 part (RFC 8032),
# then the ML-DSA-44 part (FIPS 204). A signature verifies only if both halves verify over the
# same payload, ML-DSA with an empty context. A secret key is the two 32-byte seeds.
ALGORITHM = "ed25519-ml-dsa-44"
ED25519_KEY, ED25519_SIG = 32, 64
ML_DSA_KEY, ML_DSA_SIG = 1312, 2420
SIGNATURE_FIELDS = ("sig", "key_sig")


def unhex(value, length):
    prefix = ALGORITHM + ":"
    if not isinstance(value, str) or not value.startswith(prefix):
        return None
    digits = value[len(prefix) :]
    if len(digits) != 2 * length or not re.fullmatch(r"[0-9a-f]*", digits):
        return None
    return bytes.fromhex(digits)


def public_key_of(secret):
    ed25519 = Ed25519PrivateKey.from_private_bytes(secret[:32]).public_key().public_bytes_raw()
    ml_dsa = MLDSA44PrivateKey.from_seed_bytes(secret[32:]).public_key().public_bytes_raw()
    return ALGORITHM + ":" + (ed25519 + ml_dsa).hex()


def verify_hybrid(signature, payload, public_key):
    sig = unhex(signature, ED25519_SIG + ML_DSA_SIG)
    key = unhex(public_key, ED25519_KEY + ML_DSA_KEY)
    if sig is None or key is None:
        return False
    try:
        Ed25519PublicKey.from_public_bytes(key[:ED25519_KEY]).verify(sig[:ED25519_SIG], payload)
        MLDSA44PublicKey.from_public_bytes(key[ED25519_KEY:]).verify(sig[ED25519_SIG:], payload)
        return True
    except (InvalidSignature, ValueError):
        return False


def signing_payload(obj):
    return canonical({k: v for k, v in obj.items() if k != "sig"}).encode("utf-8")


def digest_of(value):
    return "sha256:" + sha256_hex(canonical(value))


def detach_signatures(entry):
    """The entry as a log leaf holds it: each signature replaced by the SHA-256 of its canonical JSON."""
    return {k: digest_of(v) if k in SIGNATURE_FIELDS else v for k, v in entry.items()}


def check_signatures():
    vectors = json.loads((DATA / "signature-vectors.json").read_text())
    mismatches = 0
    for case in vectors["cases"]:
        payload = signing_payload(case["object"])
        secret = bytes.fromhex(case["secret_key"])
        public = public_key_of(secret)
        mismatches += payload.decode("utf-8") != case["payload"]
        mismatches += public != case["public_key"]
        mismatches += "sha256:" + sha256_hex(public) != case["key_digest"]
        mismatches += digest_of(case["sig"]) != case["sig_digest"]
        # Ed25519 is deterministic, so its half must match byte for byte. The ML-DSA half was
        # made with the deterministic variant, which this library doesn't offer, so it is
        # verified instead, and a fresh hedged signature from the same key must verify too.
        ed25519 = Ed25519PrivateKey.from_private_bytes(secret[:32]).sign(payload)
        listed = unhex(case["sig"], ED25519_SIG + ML_DSA_SIG)
        mismatches += listed is None or listed[:ED25519_SIG] != ed25519
        mismatches += not verify_hybrid(case["sig"], payload, case["public_key"])
        fresh = ALGORITHM + ":" + (ed25519 + MLDSA44PrivateKey.from_seed_bytes(secret[32:]).sign(payload)).hex()
        mismatches += not verify_hybrid(fresh, payload, case["public_key"])
    accepted = sum(
        verify_hybrid(case["sig"], signing_payload(case["object"]), case["public_key"]) for case in vectors["invalid"]
    )
    return {
        "cases": len(vectors["cases"]),
        "mismatches": mismatches,
        "invalid_cases": len(vectors["invalid"]),
        "invalid_accepted": accepted,
    }


# --- IDs ----------------------------------------------------------------------------------
#
# An operator's ID is "op:" and the lowercase hex SHA-256 of the first key it registered, as
# written; a volunteer's is "obs:" and that of the first passkey they joined with. No log assigns
# either. An entry's digest as signed, the SHA-256 of its canonical JSON with its signatures, is
# the same on every log that holds it.

OPERATOR_ID = re.compile(r"op:[0-9a-f]{64}")


def operator_id(first_key):
    return "op:" + sha256_hex(first_key)


def observer_id(first_passkey):
    return "obs:" + sha256_hex(first_passkey)


def check_ids():
    vectors = json.loads((DATA / "id-vectors.json").read_text())
    mismatches = 0
    for case in vectors["operators"]:
        mismatches += operator_id(case["public_key"]) != case["operator_id"]
    for case in vectors["observers"]:
        mismatches += observer_id(case["passkey"]) != case["observer_id"]
    keys = {case["operator_id"]: case["public_key"] for case in vectors["operators"]}
    for case in vectors["entries"]:
        entry = case["entry"]
        mismatches += digest_of(entry) != case["entry_digest"]
        mismatches += canonical(detach_signatures(entry)) != canonical(case["leaf_entry"])
        if entry["type"] == "key_rotation":
            # The rotation names the ID the operator's first key made, which it keeps: the current
            # key signs the rotation, and the new key signs it without sig and key_sig.
            unsigned = {"type": "key_rotation", "operator": entry["operator"], "key": entry["key"]}
            current = keys.get(entry["operator"])
            mismatches += current is None or not verify_hybrid(entry["sig"], signing_payload(entry), current)
            mismatches += not verify_hybrid(entry["key_sig"], canonical(unsigned).encode("utf-8"), entry["key"])
            mismatches += operator_id(entry["key"]) == entry["operator"]
    accepted = sum(
        bool(OPERATOR_ID.fullmatch(case["operator_id"])) and case["operator_id"] == operator_id(case["public_key"])
        for case in vectors["invalid"]
    )
    return {
        "operators": len(vectors["operators"]),
        "observers": len(vectors["observers"]),
        "entries": len(vectors["entries"]),
        "mismatches": mismatches,
        "invalid_cases": len(vectors["invalid"]),
        "invalid_accepted": accepted,
    }


def main():
    claim_results, naive_mismatches = check_claim_ids()
    results = {
        "R1": claim_results,
        "R2": check_bundles(),
        "R3": check_log(),
        "R4": check_signatures(),
        "R5": {
            "claims": claim_results["claims"],
            "naive_mismatches": naive_mismatches,
            "canonical_mismatches": claim_results["mismatches"],
        },
        "R6": check_ids(),
    }
    failed = any(
        value.get("mismatches", 0) or value.get("invalid_accepted", 0) or value.get("unresolved_accepted", 0)
        for value in results.values()
    )
    failed = failed or results["R5"]["canonical_mismatches"] > 0
    for key, value in results.items():
        print(key, json.dumps(value))
        path = RESULTS / f"{key}.json"
        if "--write" in sys.argv:
            RESULTS.mkdir(exist_ok=True)
            path.write_text(json.dumps(value, indent=2) + "\n")
        elif not path.exists() or json.loads(path.read_text()) != value:
            print(f"  differs from the declared {path.relative_to(ROOT)}")
            failed = True
    sys.exit(1 if failed else 0)


if __name__ == "__main__":
    main()
