"""Replay synthetic RTL security worksheets (Python 3.8+, standard library only).

No hardware access, fault injector, RTL simulator, network, or credentials.
Prints deterministic JSON; --output saves it to a user-selected file.
Existing HTML definitions are used for lesson12/13/15; new toy models are
explicitly separated from the original HTML fixtures.
"""
from itertools import product
from math import comb, log2, sqrt, ceil
from pathlib import Path
import argparse
import hashlib
import json
VERSION = 'rtl-university-workbook-v1'
TOY_SBOX = [6, 13, 0, 11, 7, 2, 15, 8, 5, 10, 4, 1, 12, 9, 3, 14]

def authorization_counts():
    counts = {'both': 0, 'fetchOnly': 0, 'keyOnly': 0, 'neither': 0}
    fetch_only = []
    for bits in product((0, 1), repeat=7):
        s, v, d, lf, lk, tf, tk = bits
        fetch = bool(s and v and d and lf and tf)
        key = bool(s and v and d and lk and tk)
        name = 'both' if fetch and key else 'fetchOnly' if fetch else 'keyOnly' if key else 'neither'
        counts[name] += 1
        if name == 'fetchOnly':
            fetch_only.append(list(bits))
    assert counts == {'both': 1, 'fetchOnly': 3, 'keyOnly': 3, 'neither': 121}
    return {
        'fieldOrder': ['S', 'V', 'D', 'L_f', 'L_k', 'T_f', 'T_k'],
        'combinations': 128,
        'counts': counts,
        'fetchOnlyRows': fetch_only,
    }

def authorization_slot_counts():
    counts = {'both': 0, 'fetchOnly': 0, 'keyOnly': 0, 'neither': 0}
    for s, v, d, lf, lk, tf, tk, ks in product((0, 1), repeat=8):
        fetch = bool(s and v and d and lf and tf)
        key = bool(s and v and d and lk and tk and ks)
        counts['both' if fetch and key else 'fetchOnly' if fetch else 'keyOnly' if key else 'neither'] += 1
    assert counts == {'both': 1, 'fetchOnly': 7, 'keyOnly': 3, 'neither': 245}
    return {
        'combinations': 256,
        'counts': counts,
        'scope': 'new eighth key-slot input, not implemented in original HTML',
    }

def noisy_samples(channel=0, separated=True, noise=1, n=8, holdout=25):
    a, b = ([0, 1, 2][channel], [3, 2, 1][channel]) if separated else ([0, 0, 1][channel], [0, 0, 1][channel])
    seed = 13
    xa = []
    xb = []

    def draw():
        nonlocal seed
        seed = seed * 1664525 + 1013904223 & 4294967295
        return seed / 2147483648 - 1
    for _ in range(n):
        xa.append(a + noise * draw())
        xb.append(b + noise * draw())
    # Positive inputs: JS Math.round uses this rule, not Python round-to-even.
    k = max(1, int(n * holdout / 100 + 0.5))
    tr = n - k
    ma = sum(xa[:tr]) / tr
    mb = sum(xb[:tr]) / tr
    tests = []
    for actual, vals in [('A', xa[tr:]), ('B', xb[tr:])]:
        for value in vals:
            da, db = (abs(value - ma), abs(value - mb))
            predicted = 'A' if da <= db else 'B'
            tests.append(dict(actual=actual, value=value, distanceA=da, distanceB=db, predicted=predicted, correct=predicted == actual))
    return dict(
        channel=channel,
        separated=separated,
        noise=noise,
        seed=13,
        trainA=xa[:tr],
        trainB=xb[:tr],
        meanA=ma,
        meanB=mb,
        tests=tests,
        correct=sum((t['correct'] for t in tests)),
        testCount=2 * k,
    )

def toy_dfa():
    assert sorted(TOY_SBOX) == list(range(16))
    inv = [0] * 16
    for x, y in enumerate(TOY_SBOX):
        inv[y] = x
    key = 10
    error = 1
    survivors = set(range(16))
    pairs = []
    for x in (3, 6, 9):
        c = TOY_SBOX[x] ^ key
        cf = TOY_SBOX[x ^ error] ^ key
        candidates = [k for k in range(16) if inv[c ^ k] ^ inv[cf ^ k] == error]
        survivors.intersection_update(candidates)
        pairs.append(dict(teacherInternalX=x, C=c, Cfault=cf, error=error, candidates=candidates, intersection=sorted(survivors)))
    assert key in survivors
    checks = [dict(k=k, inverseC=inv[pairs[0]['C'] ^ k], inverseFaultC=inv[pairs[0]['Cfault'] ^ k], difference=inv[pairs[0]['C'] ^ k] ^ inv[pairs[0]['Cfault'] ^ k], survives=k in pairs[0]['candidates']) for k in range(16)]
    return dict(
        kind='custom four-bit single substitution plus output XOR; not AES or a standard cipher',
        sbox=TOY_SBOX,
        inverse=inv,
        teacherKey=key,
        error=error,
        pairs=pairs,
        firstPairAll16Guesses=checks,
        finalCandidates=sorted(survivors),
    )

def confusion(threshold):
    pairs = [(1, 1), (4, 1), (5, 1), (2, 1), (1, 0), (3, 0), (5, 0), (2, 0)]
    tp = fp = fn = tn = 0
    for score, label in pairs:
        positive = score >= threshold
        tp += int(positive and label == 1)
        fp += int(positive and label == 0)
        fn += int(not positive and label == 1)
        tn += int(not positive and label == 0)
    return dict(
        threshold=threshold,
        TP=tp,
        FP=fp,
        FN=fn,
        TN=tn,
        cost10FNplusFP=10 * fn + fp,
        cost2FNplusFP=2 * fn + fp,
        precision=None if tp + fp == 0 else tp / (tp + fp),
        recall=tp / (tp + fn),
        accuracy=(tp + tn) / 8,
    )

def wilson(successes, n, z=1.959963984540054):
    p = successes / n
    den = 1 + z * z / n
    center = (p + z * z / (2 * n)) / den
    half = z * sqrt(p * (1 - p) / n + z * z / (4 * n * n)) / den
    return [center - half, center + half]
FIRST_FOUR_FIXTURE_ROWS = [
    {
        'attemptId': 'try-001',
        'targetClass': 'ROM check_done',
        'timeBin': 1,
        'requiredEvents': 1,
        'triggered': True,
        'toolObservable': True,
        'unauthorizedCommit': True,
        'availabilityFailure': False,
        'blockedInTime': False,
        'lateAlert': True,
        'firstBadCommitEdge': 4,
    },
    {
        'attemptId': 'try-002',
        'targetClass': 'ROM check_done',
        'timeBin': 1,
        'requiredEvents': 2,
        'triggered': True,
        'toolObservable': True,
        'unauthorizedCommit': False,
        'availabilityFailure': False,
        'blockedInTime': True,
        'lateAlert': False,
        'firstBadCommitEdge': None,
    },
    {
        'attemptId': 'try-003',
        'targetClass': 'debug permission',
        'timeBin': 3,
        'requiredEvents': 1,
        'triggered': True,
        'toolObservable': True,
        'unauthorizedCommit': False,
        'availabilityFailure': True,
        'blockedInTime': False,
        'lateAlert': False,
        'firstBadCommitEdge': None,
    },
    {
        'attemptId': 'try-004',
        'targetClass': 'ROM check_done',
        'timeBin': 1,
        'requiredEvents': 1,
        'triggered': False,
        'toolObservable': True,
        'unauthorizedCommit': False,
        'availabilityFailure': False,
        'blockedInTime': False,
        'lateAlert': False,
        'firstBadCommitEdge': None,
    },
]

def ideal_held_request(request_rise_ns, destination_period_ns=40):
    first_capture = ceil(request_rise_ns / destination_period_ns) * destination_period_ns
    second_observation = first_capture + destination_period_ns
    return dict(
        requestRiseNs=request_rise_ns,
        firstStageCapturesNs=first_capture,
        secondStageObservesNs=second_observation,
        observationDelayNs=second_observation - request_rise_ns,
        scope='ideal held request, destination phase 0, no metastability or edge-aperture effects',
    )

def campaign_first_four(budget):
    rows = FIRST_FOUR_FIXTURE_ROWS
    records = []
    for row in rows:
        activated = row['triggered'] and row['requiredEvents'] <= budget
        records.append(dict(row, activated=activated, eventsApplied=row['requiredEvents'] if activated else 0))
    selected = {(r['targetClass'], r['timeBin']) for r in records}
    observed = {(r['targetClass'], r['timeBin']) for r in records if r['activated'] and r['toolObservable']}
    summary = dict(
        attemptsExecuted=len(records),
        eventsApplied=sum((r['eventsApplied'] for r in records)),
        selectedUniqueTargetBins=len(selected),
        activatedObservableUniqueBins=len(observed),
        targetTimeBinDenominator=12,
        untouchedCount=12 - len(observed),
        notActivated=sum((not r['activated'] for r in records)),
    )
    for outcome in ('unauthorizedCommit', 'availabilityFailure', 'blockedInTime', 'lateAlert'):
        summary[outcome] = sum((r['activated'] and r[outcome] for r in records))
    return dict(
        fixtureVersion='rtl14-fixture-v3',
        scope='first four synthetic rows; no simulator or physical injection',
        attempts=4,
        eventBudgetPerAttempt=budget,
        records=records,
        summary=summary,
    )



def four_phase_handshake():
    state = dict(req=0, req1=0, req2=0, ack=0, ack1=0, ack2=0, sender='idle')
    changes = []
    captured = None
    released = None
    completed = None
    for time_ns in range(261):
        old = dict(state)
        new = dict(old)
        if time_ns == 6:
            new['req'] = 1
            new['sender'] = 'wait_ack_high'
        if time_ns >= 1 and (time_ns-1) % 5 == 0:
            new['ack1'] = old['ack']
            new['ack2'] = old['ack1']
            if old['sender'] == 'wait_ack_high' and old['ack2']:
                new['req'] = 0
                new['sender'] = 'wait_ack_low'
                released = time_ns
            elif old['sender'] == 'wait_ack_low' and not old['ack2']:
                new['sender'] = 'idle'
                completed = time_ns
        if time_ns % 40 == 0:
            new['req1'] = old['req']
            new['req2'] = old['req1']
            if old['req2'] and not old['ack']:
                new['ack'] = 1
                captured = time_ns
            elif not old['req2'] and old['ack']:
                new['ack'] = 0
        state = new
        if state != old:
            changes.append(dict(timeNs=time_ns, **state))
    assert (captured, released, completed) == (120, 131, 251)
    return dict(sourcePeriodNs=5, sourcePhaseNs=1, destinationPeriodNs=40,
                destinationPhaseNs=0, requestRiseNs=6, capturedDataNs=captured,
                dataMayReleaseNs=released, nextTransactionMayStartNs=completed,
                assumptions='ideal registers read old values at each edge; two sync stages each direction; registered controllers; no metastability or reset',
                eventTrace=changes)

def power_loss_contract():
    cases = [
        ('before_B_write',5,5,True),
        ('B_incomplete',5,5,True),
        ('B_verified_pending',5,6,True),
        ('B_confirmed_before_floor',5,6,True),
        ('atomic_floor_write_old_value',5,6,True),
        ('atomic_floor_write_new_value',6,6,True),
        ('floor_6_durable',6,6,True),
    ]
    rows = []
    for point, floor, selected_version, verified in cases:
        allowed = verified and selected_version >= floor
        assert allowed
        rows.append(dict(interruptionPoint=point,floor=floor,
                         selectedVersion=selected_version,imageVerified=verified,
                         authorized=allowed, A5AllowedUnderCurrentFloor=(5 >= floor),
                         R6AllowedUnderCurrentFloor=(6 >= floor)))
    unsafe_update = dict(
        floor=6,
        A=dict(version=5, verified=True),
        B=dict(version=6, verified=False),
        recoveryAvailable=False,
    )
    assert not (unsafe_update['A']['verified'] and unsafe_update['A']['version'] >= 6)
    assert not unsafe_update['B']['verified']
    return dict(scope='separate illustrative dual-slot contract; no storage/physical testing',
                assumptions=['verified images and metadata','atomic durable floor and confirmed record','available independently authorized R6'],
                interruptions=rows,earlyFloorUpdateCounterexample=unsafe_update,
                counterexampleResult='safe rejection possible, legitimate service has no recovery path')


def build():
    scan = [confusion(t) for t in range(1, 7)]
    assert [[r[k] for k in ('TP', 'FP', 'FN', 'TN')] for r in scan] == [[4, 4, 0, 0], [3, 3, 1, 1], [2, 2, 2, 2], [2, 1, 2, 3], [1, 1, 3, 3], [0, 0, 4, 4]]
    assert [r['cost10FNplusFP'] for r in scan] == [4, 13, 22, 21, 31, 40]
    p0 = 50 / 60
    p1 = 10 / 60
    separated = noisy_samples(noise=0)
    overlap = noisy_samples(separated=False, noise=0)
    assert (separated['correct'], overlap['correct']) == (4, 2)
    output = dict(
        schemaVersion=1,
        version=VERSION,
        scope='Synthetic teaching calculations only; no RTL/formal/netlist/physical validation',
        lesson11=dict(
            sourcePeriodNs=5,
            destinationPeriodNs=40,
            idealUniformPhaseSamplingFraction=5 / 40,
            phaseExamples=[dict(destinationPhaseNs=p, samplingEdgesNs=[p, p + 40, p + 80], pulseIntervalNs=[0, 5], observed=0 <= p < 5) for p in (2, 10, 30)],
            heldRequest=dict(
                requestRiseNs=6,
                destinationEdgesNs=[40, 80],
                firstStageCapturesNs=40,
                secondStageObservesNs=80,
                assumptions='ideal edge model, no metastability, old first-stage value used by second stage',
            ),
        ),
        lesson12=authorization_counts(),
        lesson12NewSlot=authorization_slot_counts(),
        lesson13=dict(
            toyDfa=toy_dfa(),
            conditionedSynthetic=dict(
                prior=[0.5, 0.5],
                probCorrectGivenX=[1, 0.2],
                retainedCounts=[50, 10],
                posterior=[p0, p1],
                entropyBits=-p0 * log2(p0) - p1 * log2(p1),
            ),
            defaultNoiseSamples=noisy_samples(),
            zeroNoiseSeparated=separated,
            zeroNoiseOverlap=overlap,
        ),
        lesson14=dict(
            singleBitLocations=1000,
            timeEdges=500,
            singleEvents=500000,
            unorderedDistinctTwoEventPairs=comb(500000, 2),
            zeroObservedBernoulliUpper95={str(n): 1 - 0.05 ** (1 / n) for n in (30, 3000)},
            assumptions='iid Bernoulli draws from a specified fixed sampling distribution; no guarantee over every fault-space point',
        ),
        lesson15=dict(
            syntheticScoreLabelPairs=[[1, 1], [4, 1], [5, 1], [2, 1], [1, 0], [3, 0], [5, 0], [2, 0]],
            thresholdScan=scan,
            recall2of4Wilson95=wilson(2, 4),
            unknownEffectsOutsideMatrix=['multi-bit upset', 'power-rail coupling'],
        ),
        lesson16=dict(
            existingProjectionBins=12,
            existingObservableProjectionBins=2,
            newHypotheticalDimensions=[4, 2, 3, 3],
            newHypotheticalDenominator=72,
            projectionDoesNotEstablishFullDimensionCoverage=True,
            newHypotheticalRecords=[['ROM', 'skip', 't1', 'PROD'], ['debug', 'delay', 't3', 'PROD']],
            newCoveredBins=2,
            newUncoveredBins=70,
            repeatedSameBin=dict(extraAttempts=3, extraObservableActivatedRecords=3, extraUniqueCoveredBins=0),
            teachingClaimPacket=[
                dict(
                    claimId='DBG-01',
                    boundary='debug_accept',
                    evidenceLevel='Boolean and ideal timing teaching model',
                    ownerRole='CDC reviewer',
                    date='2026-10-08',
                    status='PRODUCT_CDC_RESET_PENDING',
                ),
                dict(
                    claimId='BOOT-01',
                    boundary='first_fetch',
                    evidenceLevel='Boolean teaching model',
                    ownerRole='Boot designer',
                    date='2026-10-08',
                    status='PRODUCT_LOAD_PATH_PENDING',
                ),
                dict(
                    claimId='CAM-01',
                    boundary='commit',
                    evidenceLevel='synthetic fixed fixture',
                    ownerRole='Campaign reviewer',
                    date='2026-10-08',
                    status='COUNTEREXAMPLE_TRY001',
                ),
                dict(
                    claimId='PHY-01',
                    boundary='effect mapping',
                    evidenceLevel='synthetic confusion matrix',
                    ownerRole='Measurement owner',
                    date='2026-10-08',
                    status='PHYSICAL_EVIDENCE_MISSING',
                ),
            ],
            riskAcceptance='NOT_SIGNED; illustrative responsibility roles only, not a real product decision',
        ),
    )
    output['lesson11']['pulseFractionsByWidthNs'] = {'5': 5 / 40, '10': 10 / 40}
    output['lesson11']['idealHeldRequestCalculations'] = [ideal_held_request(6), ideal_held_request(41)]
    output['lesson14']['firstFourCampaigns'] = [campaign_first_four(2), campaign_first_four(1)]
    output['lesson11']['fourPhaseHandshake'] = four_phase_handshake()
    output['lesson12']['powerLossContract'] = power_loss_contract()
    output['sourceSha256'] = hashlib.sha256(Path(__file__).read_bytes()).hexdigest()
    return output
if __name__ == '__main__':
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument('--output', type=Path)
    args = parser.parse_args()
    text = json.dumps(build(), ensure_ascii=False, indent=2) + '\n'
    if args.output:
        args.output.write_text(text, encoding='utf-8')
    else:
        print(text, end='')
