"""Finite register-integrity teaching model. Python 3, standard library only.

One boot attempt: reset@0, capture@2, stored upset AFTER update@3,
requests@4/6, horizon 0..7. One event at one coded register can affect
several named bits. Source@2 and grant@4 are SEPARATE expanded models.
This is not RTL simulation, formal proof, physical injection or a probability.
"""
import argparse
from itertools import combinations
import json
from pathlib import Path

VARIANTS = ("parity", "complement", "ecc_continue", "ecc_reject")


def masks(width, weight):
    return [sum(1 << bit for bit in bits)
            for bits in combinations(range(width), weight)]


def encode(kind, data):
    if kind == "complement":
        if data not in (0, 1):
            raise ValueError("complement has one logical data bit")
        return 2 if data else 1
    if not 0 <= data < 16:
        raise ValueError("four-bit payload required")
    if kind == "parity":
        return data | ((bin(data).count("1") & 1) << 4)
    if not kind.startswith("ecc_"):
        raise ValueError("unknown codec")
    word = 0
    for bit, position in enumerate((3, 5, 6, 7)):
        word |= ((data >> bit) & 1) << (position - 1)
    for parity in (1, 2, 4):
        value = 0
        for position in range(1, 8):
            if position & parity:
                value ^= (word >> (position - 1)) & 1
        word |= value << (parity - 1)
    return word | ((bin(word).count("1") & 1) << 7)


def decode(kind, word):
    syndrome = 0
    total = bin(word).count("1") & 1
    corrected = word
    ce = ue = False
    if kind == "parity":
        valid = total == 0
        data = word & 15
        ue = not valid
    elif kind == "complement":
        valid = word in (1, 2)
        data = (word >> 1) & 1
        ue = not valid
    else:
        for position in range(1, 8):
            if word & (1 << (position - 1)):
                syndrome ^= position
        valid = syndrome == 0 and total == 0
        ce = bool(total)
        ue = bool(syndrome and not total)
        if ce:
            corrected ^= 1 << ((syndrome - 1) if syndrome else 7)
        # Reject policy uses the raw payload. Correction is not a security input.
        selected = corrected if kind == "ecc_continue" else word
        data = sum(((selected >> (p - 1)) & 1) << i
                   for i, p in enumerate((3, 5, 6, 7)))
    bad = ue or (kind == "ecc_reject" and ce)
    return dict(raw_valid=valid, data=data, syndrome=syndrome,
                total_parity=total, ce=bool(ce), ue=bool(ue),
                corrected=corrected, local_bad=bool(bad))


def width(kind):
    return 2 if kind == "complement" else (5 if kind == "parity" else 8)


def run(kind, authorized=False, mask=0, source=False, grant=False,
        direct_block=True):
    if kind not in VARIANTS or not 0 <= mask < (1 << width(kind)):
        raise ValueError("invalid variant or stored-bit mask")
    if sum((bool(mask), source, grant)) > 1:
        raise ValueError("one event, one target; expanded models are separate")
    checked = ref_complete = sticky = 0
    word = encode(kind, 0)
    trace = []
    for edge in range(8):
        dec = decode(kind, word)
        permit = bool(checked and dec["data"] == 1 and not sticky
                      and (not dec["local_bad"] or not direct_block))
        actual_grant = permit ^ bool(grant and edge == 4)
        accepted = bool(edge in (4, 6) and actual_grant)
        unauthorized = bool(accepted and not (ref_complete and authorized))
        row = dict(edge=edge, raw_pre=word, checked_pre=checked,
                   sticky_pre=sticky, ref_complete_pre=ref_complete,
                   ref_pass=int(authorized), grant=actual_grant,
                   accepted=accepted, unauthorized=unauthorized, **dec)
        if edge == 0:
            checked = ref_complete = sticky = 0
            word = encode(kind, 0)
        else:
            sticky |= int(dec["local_bad"])
            if edge == 2 and not checked:
                word = encode(kind, int(authorized) ^ int(source))
                checked = ref_complete = 1
            if edge == 3:
                word ^= mask
        row.update(raw_post=word, sticky_post=sticky,
                   ref_complete_post=ref_complete)
        trace.append(row)
    bad_edges = [r["edge"] for r in trace if r["unauthorized"]]
    return dict(variant=kind, authorized=authorized, mask=mask,
                mask_bits=bin(mask).count("1"), source=source, grant_fault=grant,
                direct_block=direct_block, events=int(bool(mask) or source or grant),
                target="stored_code" if mask else ("source" if source else
                       ("final_grant" if grant else "none")),
                observation_edges=[0, 7],
                accepted_edges=[r["edge"] for r in trace if r["accepted"]],
                unauthorized_edges=bad_edges,
                error_flag_edges=[r["edge"] for r in trace if r["ce"] or r["ue"]],
                outcome="UNAUTHORIZED_COMMIT" if bad_edges else
                        "NO_UNAUTHORIZED_COMMIT_IN_WINDOW", trace=trace)


def codec_checks():
    """Use an independent codebook and distance oracle, not just round-trips."""
    counts = {}
    for kind in ("parity", "complement", "ecc_continue"):
        payloads = range(2 if kind == "complement" else 16)
        book = {data: encode(kind, data) for data in payloads}
        minimum = min(bin(a ^ b).count("1") for a, b in combinations(book.values(), 2))
        assert minimum == (4 if kind.startswith("ecc") else 2)
        checked_count = 0
        for data, origin in book.items():
            for weight in range((2 if kind == "complement" else 4) + 1):
                for mask in masks(width(kind), weight):
                    received = origin ^ mask
                    dec = decode(kind, received)
                    assert dec["raw_valid"] == (received in book.values())
                    if kind == "parity":
                        assert dec["local_bad"] == bool(weight % 2)
                    elif kind == "complement":
                        assert dec["local_bad"] == (weight == 1)
                    else:
                        if weight <= 1:
                            nearest = [d for d, code in book.items()
                                       if bin(code ^ received).count("1") <= 1]
                            assert nearest == [data]
                            assert dec["data"] == data and not dec["ue"]
                        if weight == 2:
                            assert dec["ue"] and not dec["ce"]
                        if weight == 3:
                            assert dec["ce"] and not dec["ue"]
                            assert dec["corrected"] in book.values()
                            assert dec["data"] != data
                        rejected = decode("ecc_reject", received)
                        if weight in (1, 2, 3):
                            assert rejected["local_bad"]
                    checked_count += 1
        counts[kind] = dict(cases=checked_count, minimum_distance=minimum)
    assert sum(v["cases"] for v in counts.values()) == 3112
    return counts


def exercise():
    code_checks = codec_checks()
    assert encode("ecc_continue", 1) == 0x87
    controls = [run(kind, authorized=authorized) for kind in VARIANTS
                for authorized in (False, True)]
    for control in controls:
        assert control["accepted_edges"] == ([4, 6] if control["authorized"] else [])
    campaign, summary = [], []
    for kind in VARIANTS:
        for weight in range(1, (2 if kind in ("parity", "complement") else 4) + 1):
            cases = [run(kind, mask=mask) for mask in masks(width(kind), weight)]
            campaign.extend(cases)
            summary.append(dict(variant=kind, weight=weight, cases=len(cases),
                                unauthorized=sum(bool(c["unauthorized_edges"]) for c in cases),
                                error_flag=sum(bool(c["error_flag_edges"]) for c in cases)))
    assert len(campaign) == 342
    assert sum(bool(c["unauthorized_edges"]) for c in campaign) == 8
    availability = [run(kind, authorized=True, mask=mask)
                    for kind in VARIANTS for mask in masks(width(kind), 1)]
    assert len(availability) == 23
    for case in availability:
        assert case["accepted_edges"] == ([4, 6] if case["variant"] == "ecc_continue" else [])
    source_cases = [run(kind, source=True) for kind in VARIANTS]
    grant_cases = [run(kind, grant=True) for kind in VARIANTS]
    assert all(c["unauthorized_edges"] == [4, 6] for c in source_cases)
    assert all(c["unauthorized_edges"] == [4] for c in grant_cases)
    named = {
        "parity_single": run("parity", mask=1),
        "parity_paired": run("parity", mask=0x11),
        "complement_both": run("complement", mask=3),
        "ecc_triple_continue": run("ecc_continue", mask=7),
        "ecc_triple_reject": run("ecc_reject", mask=7),
        "ecc_four_reject": run("ecc_reject", mask=0x87),
        "late_alert_only": run("parity", mask=1, direct_block=False),
        "same_edge_block": run("parity", mask=1),
    }
    assert named["late_alert_only"]["unauthorized_edges"] == [4]
    assert named["same_edge_block"]["unauthorized_edges"] == []
    assert named["parity_paired"]["unauthorized_edges"] == [4, 6]
    assert named["complement_both"]["unauthorized_edges"] == [4, 6]
    assert named["ecc_triple_continue"]["unauthorized_edges"] == [4, 6]
    assert named["ecc_triple_reject"]["unauthorized_edges"] == []
    assert named["ecc_four_reject"]["unauthorized_edges"] == [4, 6]
    late = named["late_alert_only"]["trace"][4]
    assert late["local_bad"] and late["sticky_pre"] == 0 and late["sticky_post"] == 1
    triple = named["ecc_triple_continue"]["trace"][4]
    assert triple["raw_pre"] == 7 and triple["corrected"] == 0x87
    assert triple["syndrome"] == 0 and triple["ce"] and not triple["ue"]
    return dict(model="lesson03-v1", codec_checks=code_checks,
                campaign_summary=summary, controls=controls, campaign=campaign,
                availability=availability, source_cases=source_cases,
                grant_cases=grant_cases, named=named)


if __name__ == "__main__":
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--output-dir", default="lesson03-results")
    args = parser.parse_args()
    result = exercise()
    folder = Path(args.output_dir)
    folder.mkdir(parents=True, exist_ok=True)
    (folder / "lesson03-results.json").write_text(json.dumps(result, indent=2), encoding="utf-8")
    small = {key: result[key] for key in ("model", "codec_checks", "campaign_summary", "named")}
    small.update(controls=8, availability_cases=23, source_cases=4, grant_cases=4,
                 stored_fault_cases=342, stored_unauthorized_cases=8)
    (folder / "lesson03-summary.json").write_text(json.dumps(small, indent=2), encoding="utf-8")
    print("PASS: 3112 codec checks; 342 stored-fault traces (8 unauthorized); "
          "8 controls, 23 availability cases, 4 source cases, 4 grant cases; same-edge blocking.")
    print("Finite two-state model only. Counts are not physical attack probabilities.")
