#!/usr/bin/env python3
"""Serial, correct-path direction prediction. Standard library; stdout only.

predict(pc) sees no outcome. Exactly one identity-checked pending ticket per
instance. resolve(ticket, actual) trains the saved lookup, then moves history.
No branch targets, wrong-path execution, delayed feedback, or hardware timing.
"""
from dataclasses import dataclass, asdict
from itertools import product
import json
import random


def need(condition, message):
    if not condition:
        raise ValueError(message)


def integer(x, lo, hi, name):
    need(type(x) is int and lo <= x <= hi, name)


def pc_word(pc):
    integer(pc, 0, (1 << 32) - 1, '32-bit PC required')
    need(pc % 4 == 0, 'PC must be four-byte aligned')
    return pc >> 2


def sat(c, actual):
    return min(3, c + 1) if actual else max(0, c - 1)


def choose_next(c, b, g, actual):
    """High means prefer G. Equal suggestions provide no preference evidence."""
    return c if b == g else sat(c, int(g == actual))


@dataclass(frozen=True)
class Ticket:
    serial: int
    pc: int
    history: int
    indices: tuple
    counters: tuple
    predictions: tuple
    chosen: str
    prediction: int


class Direction:
    def __init__(self, m, h=0, initial=1):
        # The mathematical model allows m<=30. Bound the downloadable checker.
        integer(m, 0, 12, 'checker table bits must be 0..12')
        integer(h, 0, m, 'history bits must be 0..m')
        integer(initial, 0, 3, 'initial counter must be 0..3')
        self.m, self.h = m, h
        self.mask, self.hmask = (1 << m) - 1, (1 << h) - 1
        self.table = [initial] * (1 << m)
        self.history = 0
        self.serial = 0
        self.pending = None

    @property
    def logical_bits(self):
        return 2 * len(self.table) + self.h

    def state(self):
        # Diagnostic full copy, not a constant-time part of predict/resolve.
        return (tuple(self.table), self.history)

    def predict(self, pc):
        word = pc_word(pc)
        need(self.pending is None, 'resolve the pending prediction first')
        index = (word ^ (self.history << (self.m - self.h))) & self.mask
        c = self.table[index]
        token = Ticket(self.serial, pc, self.history, (index,), (c,),
                       (int(c >= 2),), 'D', int(c >= 2))
        self.pending = token
        return token

    def resolve(self, token, actual):
        need(token is self.pending and token is not None, 'not this pending ticket')
        integer(actual, 0, 1, 'actual direction must be integer 0 or 1')
        index = token.indices[0]
        self.table[index] = sat(self.table[index], actual)
        self.history = ((self.history << 1) | actual) & self.hmask
        self.serial += 1
        self.pending = None


class Bimodal(Direction):
    def __init__(self, m, initial=1):
        super().__init__(m, 0, initial)


class Gshare(Direction):
    pass


class Tournament:
    def __init__(self, mb, mg, h, mc, initial=1, chooser_initial=1):
        self.b = Bimodal(mb, initial)
        self.g = Gshare(mg, h, initial)
        integer(mc, 0, 12, 'checker chooser bits must be 0..12')
        integer(chooser_initial, 0, 3, 'initial chooser must be 0..3')
        self.mc = mc
        self.chooser = [chooser_initial] * (1 << mc)
        self.serial = 0
        self.pending = None
        self.components = None

    @property
    def logical_bits(self):
        return self.b.logical_bits + self.g.logical_bits + 2 * len(self.chooser)

    def state(self):
        return (self.b.state(), self.g.state(), tuple(self.chooser))

    def predict(self, pc):
        word = pc_word(pc)
        need(self.pending is None, 'resolve the pending prediction first')
        bt, gt = self.b.predict(pc), self.g.predict(pc)
        ci = word & (len(self.chooser) - 1)
        c = self.chooser[ci]
        chosen = 'G' if c >= 2 else 'B'
        p = gt.prediction if chosen == 'G' else bt.prediction
        t = Ticket(self.serial, pc, gt.history, (bt.indices[0], gt.indices[0], ci),
                   (bt.counters[0], gt.counters[0], c),
                   (bt.prediction, gt.prediction), chosen, p)
        self.pending, self.components = t, (bt, gt)
        return t

    def resolve(self, token, actual):
        need(token is self.pending and token is not None, 'not this pending ticket')
        integer(actual, 0, 1, 'actual direction must be integer 0 or 1')
        bt, gt = self.components
        ci = token.indices[2]
        self.chooser[ci] = choose_next(self.chooser[ci], *token.predictions, actual)
        # BOTH components train, whether or not selected this time.
        self.b.resolve(bt, actual)
        self.g.resolve(gt, actual)
        self.pending, self.components = None, None
        self.serial += 1


def run(model, trace, warmup=0, detail=False):
    integer(warmup, 0, len(trace), 'warmup exceeds trace')
    misses = after = 0
    rows = []
    for n, (pc, actual) in enumerate(trace):
        t = model.predict(pc)
        miss = int(t.prediction != actual)
        misses += miss
        after += miss if n >= warmup else 0
        model.resolve(t, actual)
        if detail:
            rows.append(dict(asdict(t), actual=actual, miss=miss, state_after=model.state()))
    ans = dict(events=len(trace), misses=misses, warmup=warmup,
               measured_events=len(trace)-warmup, measured_misses=after,
               logical_bits=model.logical_bits, final_state=model.state())
    if detail:
        ans['rows'] = rows
    return ans


def oracle_rows(trace, mb, mg, h, mc, initial, chooser_initial):
    """Independent recurrence: literal transition table + past-bit string.

    No Direction/Tournament/sat/choose_next calls. Each yielded observation is
    checked against the production ticket and all table cells after feedback.
    """
    nxt = ((0, 1), (0, 2), (1, 3), (2, 3))
    b = [initial] * (2 ** mb)
    g = [initial] * (2 ** mg)
    c = [chooser_initial] * (2 ** mc)
    past = ''
    for pc, actual in trace:
        bits = (('0' * h) + past)[-h:] if h else ''
        history = int(bits, 2) if bits else 0
        bi = (pc // 4) % len(b)
        address = format((pc // 4) % len(g), '0%db' % max(mg, 1))[-mg:] if mg else ''
        aligned = bits + '0' * (mg-h)
        xor = ''.join('0' if x == y else '1' for x, y in zip(address, aligned))
        gi = int(xor, 2) if xor else 0
        ci = (pc // 4) % len(c)
        bp, gp = int(b[bi] in (2, 3)), int(g[gi] in (2, 3))
        use_g = c[ci] in (2, 3)
        obs = (history, (bi, gi, ci), (b[bi], g[gi], c[ci]), (bp, gp),
               'G' if use_g else 'B', gp if use_g else bp)
        if bp != gp:
            c[ci] = nxt[c[ci]][int(gp == actual)]
        b[bi], g[gi] = nxt[b[bi]][actual], nxt[g[gi]][actual]
        past += str(actual)
        newbits = (('0'*h) + past)[-h:] if h else ''
        gh = int(newbits, 2) if newbits else 0
        state = ((tuple(b), 0), (tuple(g), gh), tuple(c))
        yield obs, state


def check_trace(trace, config):
    mb, mg, h, mc, init, ci = config
    m = Tournament(*config[:4], initial=init, chooser_initial=ci)
    db, dg = Bimodal(mb, init), Gshare(mg, h, init)
    for (pc, actual), (obs, state) in zip(trace, oracle_rows(trace, *config)):
        t = m.predict(pc)
        got = (t.history, t.indices, t.counters, t.predictions, t.chosen, t.prediction)
        need(got == obs, 'prediction differs from independent recurrence')
        bt, gt = db.predict(pc), dg.predict(pc)
        need(t.predictions == (bt.prediction, gt.prediction), 'component trained selectively')
        db.resolve(bt, actual); dg.resolve(gt, actual); m.resolve(t, actual)
        need(m.state() == state, 'state differs from independent recurrence')
        need((db.state(), dg.state()) == state[:2], 'component trajectory diverged')
    return len(trace)


def rejection_checks():
    count = 0
    for m in (Bimodal(2), Gshare(2, 2), Tournament(2, 2, 2, 2)):
        def reject(fn):
            nonlocal count
            state, pending, serial = m.state(), m.pending, m.serial
            try:
                fn()
            except ValueError:
                need((m.state(), m.pending, m.serial) == (state, pending, serial), 'rejection mutated state')
                count += 1
            else:
                raise ValueError('invalid action accepted')
        for pc in (-4, 2, 1 << 32, True, 3.0):
            reject(lambda pc=pc: m.predict(pc))
        reject(lambda: m.resolve(None, 0))
        t = m.predict(0x100)
        reject(lambda: m.predict(0x104))
        reject(lambda: m.resolve(Ticket(**asdict(t)), 0))
        other = Bimodal(2).predict(0x100)
        reject(lambda: m.resolve(other, 1))
        for actual in (-1, 2, True, 0.0):
            reject(lambda actual=actual: m.resolve(t, actual))
        m.resolve(t, 1)
        reject(lambda: m.resolve(t, 1))
    for args in [(-1, 0), (13, 0), (2, 3), (True, 0), (2, -1)]:
        try:
            Gshare(*args)
        except ValueError:
            count += 1
        else:
            raise ValueError('invalid configuration accepted')
    return count


def witnesses():
    # Real update-order mutants; small and explicit, not just labels.
    alternating = [(0x100, y) for y in [1, 0]*8]
    correct = run(Gshare(1, 1), alternating, detail=True)
    table, history, wrong_rows = [1, 1], 0, []
    for pc, actual in alternating:
        j = ((pc >> 2) ^ history) & 1
        before = tuple(table)
        p = int(table[j] >= 2)
        history = actual  # WRONG: forget the old lookup index.
        wrong = ((pc >> 2) ^ history) & 1
        table[wrong] = sat(table[wrong], actual)
        wrong_rows.append(dict(old_table=before, predicted_index=j, trained_index=wrong, prediction=p, actual=actual))
    wrong_misses = sum(r['prediction'] != r['actual'] for r in wrong_rows)
    need(wrong_misses != correct['misses'], 'wrong-history witness not distinguishing')
    # Future-outcome leakage: train the bimodal counter before asking for output.
    c, leaked = 1, []
    for _, actual in alternating:
        c = sat(c, actual)
        leaked.append(int(c >= 2))
    need(sum(p != y for p, (_, y) in zip(leaked, alternating)) == 0, 'leak witness')
    # Using updated predictions hides a candidate's just-observed error.
    c0, b0, g0, actual = 1, 1, 2, 1
    good = choose_next(c0, int(b0 >= 2), int(g0 >= 2), actual)
    bad = choose_next(c0, int(sat(b0, actual) >= 2), int(sat(g0, actual) >= 2), actual)
    need((good, bad) == (2, 1), 'updated suggestion witness')
    # Only selected B is trained: G remains 1 instead of recording this outcome.
    selected = Tournament(1, 1, 1, 1)
    t = selected.predict(0x100)
    need(t.chosen == 'B', 'selective witness chooses B')
    selected.b.resolve(selected.components[0], 1)
    selective_g = selected.g.table[t.indices[1]]
    correct_model = Tournament(1, 1, 1, 1)
    ct = correct_model.predict(0x100); correct_model.resolve(ct, 1)
    need((selective_g, correct_model.g.table[ct.indices[1]]) == (1, 2), 'unselected G must learn')
    # Selector-only linear-regret counterexample, constant candidates B=0/G=1.
    c, chooser_rows = 1, []
    for y in [1, 0]*8:
        old = c; p = int(c >= 2); c = choose_next(c, 0, 1, y)
        chooser_rows.append(dict(before=old, chosen_prediction=p, actual=y, after=c))
    need(all(r['chosen_prediction'] != r['actual'] for r in chooser_rows), 'chooser linear-regret witness')
    return dict(wrong_history=dict(correct_misses=correct['misses'], mutant_misses=wrong_misses, rows=wrong_rows[:4]),
                future_leak=dict(causal_bimodal_misses=run(Bimodal(1), alternating)['misses'], mutant_misses=0),
                updated_suggestions=dict(before=(c0, b0, g0), actual=actual, correct_chooser=good, mutant_chooser=bad),
                selective_training=dict(correct_g=2, mutant_g=selective_g),
                selector_regret=dict(events=16, selector_misses=16, best_fixed_misses=8, rows=chooser_rows[:4]))


def main():
    examples = {}
    loop = [(0x100, y) for y in [1, 1, 1, 0]*4]
    alternate = [(0x100, y) for y in [1, 0]*8]
    alias = [(pc, y) for _ in range(8) for pc, y in [(0x100, 1), (0x110, 0)]]
    # A alternates; the immediately following B copies A's actual direction.
    correlated = [(pc, y) for k in range(16) for pc, y in [(0x100, k % 2), (0x108, k % 2)]]
    phase = loop + alternate + [(0x100, 1)] * 16
    for name, trace in [('loop', loop), ('alternating', alternate), ('alias', alias), ('correlated', correlated), ('phase', phase)]:
        examples[name] = dict(trace=trace,
            bimodal=run(Bimodal(2), trace, warmup=min(4, len(trace)), detail=True),
            gshare=run(Gshare(2, 2), trace, warmup=min(4, len(trace)), detail=True),
            tournament=run(Tournament(2, 2, 2, 2), trace, warmup=min(4, len(trace)), detail=True))
    examples['alias_larger'] = run(Bimodal(3), alias, detail=True)
    moved = [(0x104 if pc == 0x110 else pc, y) for pc, y in alias]
    examples['alias_relayout'] = run(Bimodal(2), moved, detail=True)
    examples['alternating_one_bit'] = run(Gshare(1, 1), alternate, detail=True)
    examples['correlated_short_history'] = run(Gshare(2, 1), correlated, detail=True)
    budget_comparison = {}
    for name, trace in [('loop', loop), ('alternating', alternate), ('correlated', correlated), ('phase', phase)]:
        budget_comparison[name] = {label: run(model, trace) for label, model in (
            ('bimodal_m3', Bimodal(3)), ('gshare_m3_h2', Gshare(3, 2)),
            ('tournament_m2', Tournament(2, 2, 2, 2)))}
    migrations = []
    for m in (0, 1, 2, 3):
        for h in range(m + 1):
            r = run(Gshare(m, h), correlated)
            migrations.append(dict(m=m, h=h, misses=r['misses'], logical_bits=r['logical_bits']))
    # Exact degeneracy and constant-outcome bounds, not statistical estimates.
    degeneracies = 0
    rng = random.Random(22022)
    for m in range(5):
        for initial in range(4):
            trace = [(4*rng.randrange(64), rng.randrange(2)) for _ in range(80)]
            b, g = Bimodal(m, initial), Gshare(m, 0, initial)
            need(run(b, trace, detail=True) == run(g, trace, detail=True), 'h=0 differs')
            degeneracies += 1
    for initial, y in product(range(4), range(2)):
        need(run(Bimodal(0, initial), [(0, y)]*10)['misses'] <= 2, 'constant direction bound')
    configs = [(m, m, h, m, init, ci) for m in (0, 1, 2) for h in range(m+1) for init, ci in ((0, 3), (1, 1), (2, 2), (3, 0))]
    alphabet = tuple(product((0, 4, 16), (0, 1)))
    exhaustive = events = 0
    for n in range(5):
        for trace in product(alphabet, repeat=n):
            for cfg in configs:
                events += check_trace(trace, cfg)
                exhaustive += 1
    random_programs = random_events = 0
    for _ in range(1600):
        mb, mg, mc = (rng.randrange(5) for _ in range(3))
        cfg = (mb, mg, rng.randrange(mg+1), mc, rng.randrange(4), rng.randrange(4))
        trace = [(4*rng.randrange(256), rng.randrange(2)) for _ in range(rng.randrange(65))]
        random_events += check_trace(trace, cfg)
        random_programs += 1
    empty = run(Tournament(0, 0, 0, 0), [])
    result = dict(contract='serial correct-path directions; predict before feedback; no targets or pipeline simulation',
                  examples=examples, budget_cap_bits=26, budget_comparison=budget_comparison,
                  history_migrations=migrations, mutants=witnesses(),
                  tests=dict(exhaustive_configurations=exhaustive, exhaustive_events=events,
                    random_traces=random_programs, random_events=random_events,
                    zero_history_equivalences=degeneracies, rejected_actions=rejection_checks(), empty=empty))
    print(json.dumps(result, ensure_ascii=False, indent=2))


if __name__ == '__main__':
    main()
