#!/usr/bin/env python3
"""S23: exact single-thread integer-slot memory windows; Python 3.10+, stdlib.
No hardware timing, ISA memory-model, cache, byte-overlap or security claim.
"""
import sys
sys.dont_write_bytecode = True
from dataclasses import dataclass
from itertools import product
import json
import random


def require(ok, message):
    if not ok:
        raise AssertionError(message)


def integer(x):
    return type(x) is int


def demand(ok, message):
    if not ok:
        raise ValueError(message)


@dataclass(frozen=True)
class Ref:
    load: int
    scale: int = 1
    offset: int = 0


@dataclass(frozen=True)
class Op:
    kind: str
    pc: int
    address: object
    data: object = None
    address_delay: int = 1
    data_delay: int = 1


def validate(program, memory):
    demand(type(program) in (list, tuple), 'program must be a list/tuple')
    demand(type(memory) is dict and all(integer(k) and integer(v)
           for k, v in memory.items()), 'integer-slot memory required')
    roles = {}
    for i, op in enumerate(program):
        demand(type(op) is Op and op.kind in ('L', 'S'), 'bad operation')
        demand(integer(op.pc) and op.pc >= 0 and op.pc % 4 == 0, 'bad PC')
        demand(roles.get(op.pc, op.kind) == op.kind, 'one PC has two roles')
        roles[op.pc] = op.kind
        demand(integer(op.address_delay) and op.address_delay > 0 and
               integer(op.data_delay) and op.data_delay > 0, 'positive delays')
        demand(op.data is None if op.kind == 'L' else op.data is not None,
               'load has no data operand; store needs one')
        for expr in ([op.address] if op.kind == 'L' else [op.address, op.data]):
            if integer(expr):
                continue
            demand(type(expr) is Ref and integer(expr.load) and
                   0 <= expr.load < i and program[expr.load].kind == 'L' and
                   integer(expr.scale) and integer(expr.offset),
                   'references must name earlier loads')


def select_source(stores, address, memory, conservative=True):
    """Pure query. stores = [(age, known_address_or_None, data_or_None), ...]."""
    demand(type(conservative) is bool and integer(address) and
           type(stores) in (list, tuple), 'bad query')
    demand(type(memory) is dict and all(integer(k) and integer(v)
           for k, v in memory.items()), 'bad memory')
    ages = []
    for row in stores:
        demand(type(row) in (tuple, list) and len(row) == 3, 'bad store row')
        i, a, d = row
        demand(integer(i) and i >= 0 and (a is None or integer(a)) and
               (d is None or integer(d)), 'bad store fields')
        ages.append(i)
    demand(len(set(ages)) == len(ages), 'duplicate store age')
    unknown = sorted(i for i, a, d in stores if a is None)
    if conservative and unknown:
        return {'status': 'WAIT', 'reason': 'unknown-address', 'stores': unknown}
    matches = [(i, d) for i, a, d in stores if a == address]
    if matches:
        i, d = max(matches)
        if d is None:
            return {'status': 'WAIT', 'reason': 'matching-data', 'stores': [i]}
        return {'status': 'VALUE', 'source': i, 'value': d}
    return {'status': 'VALUE', 'source': -1, 'value': memory.get(address, 0)}


class StoreSets:
    """Exact PC keys; eager whole-set merge, rather than tagless hardware tables."""
    def __init__(self):
        self.ssit = {}
        self.next_set = 0

    def learn(self, store_pc, load_pc):
        demand(all(integer(p) and p >= 0 and p % 4 == 0
                   for p in (store_pc, load_pc)), 'bad training PCs')
        a, b = self.ssit.get(store_pc), self.ssit.get(load_pc)
        if a is None and b is None:
            winner = self.next_set
            self.next_set += 1
        elif a is None or b is None:
            winner = b if a is None else a
        else:
            winner, loser = min(a, b), max(a, b)
            if winner != loser:
                for pc in self.ssit:
                    if self.ssit[pc] == loser:
                        self.ssit[pc] = winner
        self.ssit[store_pc] = self.ssit[load_pc] = winner
        return winner

    def clear(self):
        self.ssit.clear()
        self.next_set = 0

    def snapshot(self):
        return [[pc, self.ssit[pc]] for pc in sorted(self.ssit)]


@dataclass(frozen=True)
class Token:
    index: int
    generation: int


@dataclass
class Entry:
    token: Token
    born: int = 0
    address: object = None
    data: object = None
    result: object = None
    source: object = None
    done: bool = False
    predecessor: object = None


class MemoryWindow:
    def __init__(self, program, memory, mode='conservative', predictor=None,
                 log=True):
        validate(program, memory)
        demand(mode in ('conservative', 'speculate', 'predict'), 'bad mode')
        demand(predictor is None or type(predictor) is StoreSets, 'bad predictor')
        demand(type(log) is bool, 'bad log flag')
        self.program = tuple(program)
        self.memory = dict(memory)
        self.mode = mode
        self.predictor = predictor if predictor is not None else StoreSets()
        self.entries = [Entry(Token(i, 0)) for i in range(len(program))]
        self.head = 0
        self.committed_loads = {}
        self.lfst = {}
        self.events = []
        self.logging = log
        self.round = 0
        self.replays = 0
        self.cancelled = 0
        self.loads_executed = 0
        self._rebuild_predictions()

    def _emit(self, event_name, **kw):
        if self.logging:
            self.events.append({'round': self.round, 'event': event_name, **kw})

    def token(self, i):
        demand(integer(i) and self.head <= i < len(self.entries), 'inactive index')
        return self.entries[i].token

    def _entry(self, token, kind=None):
        demand(type(token) is Token and integer(token.index) and
               self.head <= token.index < len(self.entries), 'inactive token')
        i = token.index
        demand(token is self.entries[i].token, 'stale or foreign token')
        demand(kind is None or self.program[i].kind == kind, 'wrong operation kind')
        return i, self.entries[i]

    def _value(self, expr):
        if integer(expr):
            return expr
        value = self.entries[expr.load].result
        return None if value is None else expr.scale * value + expr.offset

    def _prediction_ready(self, entry):
        p = entry.predecessor
        return p is None or self.entries[p].done

    def _rebuild_predictions(self):
        self.lfst = {}
        for i in range(self.head, len(self.entries)):
            e, op = self.entries[i], self.program[i]
            e.predecessor = None
            if self.mode != 'predict' or e.done:
                continue
            sid = self.predictor.ssit.get(op.pc)
            if sid is not None:
                e.predecessor = self.lfst.get(sid)
                if op.kind == 'S':
                    self.lfst[sid] = i

    def _violations(self, store):
        a = self.entries[store].address
        return [i for i in range(store + 1, len(self.entries))
                if self.program[i].kind == 'L' and self.entries[i].done
                and self.entries[i].address == a
                and self.entries[i].source < store]

    def _reset_suffix(self, first):
        killed = []
        for i in range(first, len(self.entries)):
            old = self.entries[i].token
            killed.append([i, old.generation])
            self.entries[i] = Entry(Token(i, old.generation + 1), born=self.round)
        return killed

    def _recover(self, store, load):
        self.replays += 1
        killed = self._reset_suffix(load)
        self.cancelled += len(killed)
        if self.mode == 'predict':
            self.predictor.learn(self.program[store].pc, self.program[load].pc)
        self._rebuild_predictions()
        self._emit('replay', store=store, load=load, cancelled=killed,
                   ssit=self.predictor.snapshot())
        return load

    def address(self, token):
        i, e = self._entry(token)
        demand(e.address is None, 'address already known')
        a = self._value(self.program[i].address)
        if a is None:
            return {'status': 'WAIT', 'reason': 'operand'}
        e.address = a
        self._emit('address', index=i, generation=token.generation, address=a)
        bad = self._violations(i) if self.program[i].kind == 'S' else []
        if bad:
            j = self._recover(i, min(bad))
            return {'status': 'REPLAY', 'first': j}
        return {'status': 'READY', 'address': a}

    def data(self, token):
        i, e = self._entry(token, 'S')
        demand(e.data is None, 'data already known')
        d = self._value(self.program[i].data)
        if d is None:
            return {'status': 'WAIT', 'reason': 'operand'}
        e.data = d
        self._emit('data', index=i, generation=token.generation, value=d)
        return {'status': 'READY', 'value': d}

    def finish_store(self, token):
        i, e = self._entry(token, 'S')
        demand(not e.done, 'store already complete')
        if e.address is None or e.data is None:
            return {'status': 'WAIT', 'reason': 'store-operands'}
        if not self._prediction_ready(e):
            return {'status': 'WAIT', 'reason': 'predicted-store',
                    'store': e.predecessor}
        e.done = True
        sid = self.predictor.ssit.get(self.program[i].pc)
        if sid is not None and self.lfst.get(sid) == i:
            del self.lfst[sid]
        self._emit('store-ready', index=i, generation=token.generation)
        return {'status': 'READY'}

    def load(self, token):
        i, e = self._entry(token, 'L')
        demand(not e.done, 'load already executed')
        if e.address is None:
            return {'status': 'WAIT', 'reason': 'load-address'}
        if not self._prediction_ready(e):
            return {'status': 'WAIT', 'reason': 'predicted-store',
                    'store': e.predecessor}
        stores = [(j, self.entries[j].address, self.entries[j].data)
                  for j in range(self.head, i) if self.program[j].kind == 'S']
        answer = select_source(stores, e.address, self.memory,
                               self.mode == 'conservative')
        if answer['status'] == 'WAIT':
            return answer
        e.source, e.result, e.done = answer['source'], answer['value'], True
        self.loads_executed += 1
        self._emit('load', index=i, generation=token.generation,
                   address=e.address, source=e.source, value=e.result)
        return answer

    def retire(self):
        if self.head == len(self.entries) or not self.entries[self.head].done:
            return {'status': 'WAIT', 'reason': 'head'}
        i, e, op = self.head, self.entries[self.head], self.program[self.head]
        if op.kind == 'S':
            self.memory[e.address] = e.data
        else:
            self.committed_loads[i] = e.result
        self.head += 1
        self._emit('retire', index=i, kind=op.kind, address=e.address,
                   value=e.data if op.kind == 'S' else e.result)
        return {'status': 'RETIRED', 'index': i}

    def clear_history(self):
        self.predictor.clear()
        self._rebuild_predictions()
        self._emit('clear-history')

    def state(self):
        return {'head': self.head, 'memory': sorted(self.memory.items()),
                'committed_loads': sorted(self.committed_loads.items()),
                'entries': [dict(index=i, generation=e.token.generation,
                    born=e.born, address=e.address, data=e.data, result=e.result,
                    source=e.source, done=e.done, predecessor=e.predecessor)
                    for i, e in enumerate(self.entries)],
                'ssit': self.predictor.snapshot(), 'lfst': sorted(self.lfst.items()),
                'replays': self.replays, 'cancelled': self.cancelled,
                'loads_executed': self.loads_executed, 'round': self.round}


def sequential(program, memory):
    """Independent literal sequential semantics, retaining each commit prefix."""
    validate(program, memory)
    mem, values, prefixes = dict(memory), {}, []
    def ev(x):
        return x if type(x) is int else values[x.load] * x.scale + x.offset
    for i, op in enumerate(program):
        a = ev(op.address)
        if op.kind == 'S':
            mem[a] = ev(op.data)
        else:
            values[i] = mem.get(a, 0)
        prefixes.append((dict(mem), dict(values)))
    return mem, values, prefixes


def run(program, memory, mode='conservative', predictor=None, order='oldest',
        max_rounds=100000, log=True, cls=MemoryWindow, verify=True):
    demand(order in ('oldest', 'youngest') and integer(max_rounds) and
           max_rounds > 0 and type(verify) is bool, 'bad driver options')
    w = cls(program, memory, mode, predictor, log)
    prefixes = sequential(program, memory)[2] if verify else None
    stalls = {}
    if not program:
        return w, stalls
    for c in range(1, max_rounds + 1):
        w.round = c
        answer = w.retire()
        if answer['status'] == 'RETIRED' and verify:
            em, ev = prefixes[answer['index']]
            require(w.memory == em and w.committed_loads == ev,
                    ('retirement-prefix', mode, c, answer, w.state(), em, ev))
        if w.head == len(program):
            return w, stalls
        for i in range(w.head, len(program)):
            e, op = w.entries[i], program[i]
            if e.address is None and c >= e.born + op.address_delay:
                w.address(w.token(i))
        for i in range(w.head, len(program)):
            e, op = w.entries[i], program[i]
            if op.kind == 'S' and e.data is None and c >= e.born + op.data_delay:
                w.data(w.token(i))
        for i in range(w.head, len(program)):
            if program[i].kind == 'S' and not w.entries[i].done:
                w.finish_store(w.token(i))
        indices = range(w.head, len(program))
        if order == 'youngest':
            indices = reversed(indices)
        for i in indices:
            if program[i].kind == 'L' and not w.entries[i].done:
                answer = w.load(w.token(i))
                if answer['status'] == 'VALUE':
                    break
                why = answer['reason']
                stalls[why] = stalls.get(why, 0) + 1
    raise RuntimeError('round budget exhausted; no completed result returned')


def summary(w, stalls):
    return dict(rounds=w.round, replays=w.replays, cancelled=w.cancelled,
                loads_executed=w.loads_executed,
                memory=sorted(w.memory.items()),
                loads=sorted(w.committed_loads.items()), stalls=stalls)


def main():
    # The load's tentative value immediately reaches a younger store operand.
    p = [Op('S', 0x100, 4, 17, 6, 1), Op('L', 0x104, 4),
         Op('S', 0x108, 8, Ref(1, 2, 1)), Op('L', 0x10c, 8)]
    examples = {}
    for mode in ('conservative', 'speculate', 'predict'):
        w, stalls = run(p, {4: 3, 8: 0}, mode)
        examples[mode] = dict(summary=summary(w, stalls), events=w.events,
                              final=w.state())
    trained = StoreSets()
    cold, cs = run(p, {4: 3, 8: 0}, 'predict', trained)
    warm, ws = run(p, {4: 3, 8: 0}, 'predict', trained)
    changed = [Op('S', 0x100, 9, 17, 6, 1), *p[1:]]
    phase, ps = run(changed, {4: 3, 8: 0}, 'predict', trained)
    cleared = StoreSets()
    cleared.learn(0x100, 0x104)
    cleared.clear()
    reset, rs = run(p, {4: 3, 8: 0}, 'predict', cleared)
    examples['learning'] = {'cold': summary(cold, cs), 'warm': summary(warm, ws),
       'changed_address': summary(phase, ps), 'cleared': summary(reset, rs),
       'trained_ssit': trained.snapshot(), 'warm_events': warm.events}
    # A younger already-known matching store masks an older late store.
    shadow = [Op('S', 0x200, 4, 17, 7), Op('S', 0x204, 4, 23),
              Op('L', 0x208, 4)]
    w, st = run(shadow, {4: 3}, 'speculate')
    require(w.replays == 0 and w.committed_loads[2] == 23, 'masking')
    examples['shadow'] = dict(summary=summary(w, st), events=w.events)
    same = [Op('S', 0x200, 4, 3, 6), Op('L', 0x204, 4)]
    w, st = run(same, {4: 3}, 'speculate')
    require(w.replays == 1, 'equal-value provenance still replays')
    examples['same_value'] = summary(w, st)
    derived = [Op('S', 0x100, 4, 17, 6), Op('L', 0x104, 4),
               Op('S', 0x108, Ref(1), Ref(1, 2, 1)), Op('L', 0x10c, Ref(1))]
    w, st = run(derived, {3: 0, 4: 3, 17: 0}, 'speculate')
    require(w.memory == {3: 0, 4: 17, 17: 35}, 'derived address rollback')
    examples['derived_address'] = dict(summary=summary(w, st), events=w.events)
    examples['source_queries'] = [
        select_source([(0, None, 17), (2, 4, 23)], 4, {4: 3}, True),
        select_source([(0, None, 17), (2, 4, 23)], 4, {4: 3}, False),
        select_source([(0, 4, 17), (2, 4, None)], 4, {4: 3}, True),
        select_source([(0, 4, 17), (2, 4, 23)], 4, {4: 3}, True),
        select_source([(0, 9, None)], 4, {4: 3}, True)]
    # Pure forwarding interface: enumerate partially revealed stores.
    selections = 0
    for n in range(5):
        for rows in product(tuple(product((None, 0, 1), (None, 7, 9))), repeat=n):
            stores = [(i, *row) for i, row in enumerate(rows)]
            for a, conservative in product((0, 1), (False, True)):
                answer = select_source(stores, a, {0: 2, 1: 4}, conservative)
                unknown = any(x[1] is None for x in stores)
                matching = [s for s in stores if s[1] == a]
                if conservative and unknown:
                    require(answer['status'] == 'WAIT' and
                            answer['reason'] == 'unknown-address', 'unknown')
                elif matching and matching[-1][2] is None:
                    require(answer == {'status': 'WAIT', 'reason': 'matching-data',
                            'stores': [matching[-1][0]]}, 'youngest-data')
                else:
                    source, value = ((matching[-1][0], matching[-1][2])
                                     if matching else (-1, {0: 2, 1: 4}[a]))
                    require(answer == {'status': 'VALUE', 'source': source,
                                       'value': value}, 'source')
                selections += 1
    rng = random.Random(230010)
    random_runs = rounds = replays = 0
    for trial in range(1200):
        program, load_ids = [], []
        def expr():
            if load_ids and rng.random() < .38:
                return Ref(rng.choice(load_ids), rng.choice((-1, 0, 1, 2)),
                           rng.randrange(-2, 5))
            return rng.randrange(-2, 6)
        for i in range(rng.randrange(1, 15)):
            k = rng.choice(('L', 'S'))
            program.append(Op(k, 4 * (2 * rng.randrange(5) + (k == 'L')),
                              expr(), expr() if k == 'S' else None,
                              rng.randrange(1, 11), rng.randrange(1, 11)))
            if k == 'L': load_ids.append(i)
        memory = {i: rng.randrange(-3, 5) for i in range(-2, 6)}
        for mode, order in product(('conservative', 'speculate', 'predict'),
                                   ('oldest', 'youngest')):
            predictor = StoreSets()
            for s in program:
                for l in program:
                    if s.kind == 'S' and l.kind == 'L' and rng.random() < .06:
                        predictor.learn(s.pc, l.pc)
            w, _ = run(program, memory, mode, predictor, order, log=False)
            random_runs += 1; rounds += w.round; replays += w.replays
    # Exact table merge and live store chain; completing an older entry must
    # not erase the latest entry's LFST pointer.
    ss = StoreSets()
    ss.learn(0x300, 0x304); ss.learn(0x308, 0x30c)
    before_merge = ss.snapshot()
    ss.learn(0x300, 0x30c)
    q = [Op('S', 0x300, 0, 5), Op('S', 0x308, 1, 6), Op('L', 0x304, 2)]
    chain = MemoryWindow(q, {}, 'predict', ss)
    require([e.predecessor for e in chain.entries] == [None, 0, 1], 'chain')
    for i in (0, 1): chain.address(chain.token(i)); chain.data(chain.token(i))
    require(chain.finish_store(chain.token(1))['status'] == 'WAIT', 'store-chain')
    chain.finish_store(chain.token(0))
    require(list(chain.lfst.values()) == [1], 'conditional LFST clear')
    chain.finish_store(chain.token(1))
    require(chain.lfst == {}, 'last clear')
    chain.address(chain.token(2))
    chain_read = chain.load(chain.token(2))
    require(chain_read == {'status': 'VALUE', 'source': -1, 'value': 0},
            'predicted predecessor is not the actual value source')
    for _ in q: chain.retire()
    require(chain.memory == {0: 5, 1: 6} and chain.committed_loads == {2: 0},
            'chain retirement')
    examples['merge'] = dict(before=before_merge, after=ss.snapshot(),
        chain=[None, 0, 1], actual_read=chain_read, final=chain.state(),
        events=chain.events)
    # Clear genuinely trained history while a live load is waiting on it.
    history = StoreSets(); history.learn(0x100, 0x104)
    live = MemoryWindow(p, {4: 3, 8: 0}, 'predict', history)
    live.address(live.token(1))
    wait_before = live.load(live.token(1))
    require(wait_before['reason'] == 'predicted-store', 'trained live wait')
    before_clear = live.state()
    live.clear_history()
    after_clear = live.state()
    require(not live.predictor.ssit and not live.lfst and
            all(e.predecessor is None for e in live.entries), 'clear live chain')
    require(live.memory == {4: 3, 8: 0} and live.entries[1].address == 4,
            'clear preserves program state')
    require(live.load(live.token(1))['value'] == 3, 'clear allows speculation')
    require(live.address(live.token(0))['status'] == 'REPLAY', 'backstop after clear')
    live.data(live.token(0)); live.finish_store(live.token(0)); live.retire()
    live.address(live.token(1)); live.load(live.token(1)); live.retire()
    live.address(live.token(2)); live.data(live.token(2)); live.finish_store(live.token(2))
    live.address(live.token(3)); live.load(live.token(3)); live.retire(); live.retire()
    require((live.memory, live.committed_loads) == sequential(p, {4: 3, 8: 0})[:2],
            'live clear and retraining preserve result')
    examples['live_clear'] = dict(before=before_clear, after_clear=after_clear,
        wait_before=wait_before, final=live.state(), events=live.events)
    # Tokens reject old, foreign and forged replies without state/log changes.
    stale = MemoryWindow(p, {4: 3, 8: 0}, 'speculate')
    old = stale.token(2)
    stale.address(stale.token(1)); stale.load(stale.token(1))
    stale.data(stale.token(2))
    stale.address(stale.token(0))
    rejects = 0
    bad_calls = [lambda: stale.data(old),
                 lambda: stale.data(Token(2, old.generation + 1)),
                 lambda: stale.data(MemoryWindow(p, {}, 'speculate').token(2)),
                 lambda: stale.address(stale.token(0)),
                 lambda: stale.data(stale.token(1)),
                 lambda: stale.token(True)]
    for call in bad_calls:
        before = json.dumps([stale.state(), stale.events], sort_keys=True)
        try: call()
        except ValueError: rejects += 1
        else: raise AssertionError('accepted invalid action')
        require(before == json.dumps([stale.state(), stale.events], sort_keys=True),
                'rejection mutates state')
    bad_inputs = [
        lambda: MemoryWindow([Op('L', True, 0)], {}),
        lambda: MemoryWindow([Op('L', 2, 0)], {}),
        lambda: MemoryWindow([Op('L', 0, True)], {}),
        lambda: MemoryWindow([Op('L', 0, Ref(0))], {}),
        lambda: MemoryWindow([Op('S', 0, 0, 1), Op('L', 4, Ref(0))], {}),
        lambda: MemoryWindow([Op('L', 0, 0), Op('S', 0, 1, 2)], {}),
        lambda: MemoryWindow([Op('S', 0, 0)], {}),
        lambda: MemoryWindow([Op('L', 0, 0, 1)], {}),
        lambda: MemoryWindow([Op('L', 0, 0, address_delay=0)], {}),
        lambda: MemoryWindow([], {True: 0}),
        lambda: select_source([(0, 1, 2), (0, 1, 3)], 1, {}),
        lambda: select_source([(0, False, 2)], 1, {}),
        lambda: run([], {}, max_rounds=True)]
    for call in bad_inputs:
        try: call()
        except ValueError: rejects += 1
        else: raise AssertionError('accepted invalid configuration')
    empty, _ = run([], {})
    require(empty.round == 0 and not empty.events, 'empty program')
    # External mutants of the control protocol, not hidden checker switches.
    class NoReplay(MemoryWindow):
        def _violations(self, store): return []
    class OnlyLoad(MemoryWindow):
        def _reset_suffix(self, first):
            old = self.entries[first].token
            self.entries[first] = Entry(Token(first, old.generation + 1),
                                        born=self.round)
            return [[first, old.generation]]
    mutants = []
    for candidate in (NoReplay, OnlyLoad):
        wrong, st = run(p, {4: 3, 8: 0}, 'speculate', cls=candidate, verify=False)
        em, ev, _ = sequential(p, {4: 3, 8: 0})
        require((wrong.memory, wrong.committed_loads) != (em, ev), 'mutant survives')
        mutants.append({'mutant': candidate.__name__, 'result': summary(wrong, st)})
    return dict(status='PASS', pure_source_queries=selections,
                random_runs=random_runs, driver_rounds=rounds, replays=replays,
                invalid_actions_rejected=rejects, external_mutants=mutants,
                examples=examples,
                scope='Author executable checks, not independent review or CPU benchmarks')


if __name__ == '__main__':
    print(json.dumps(main(), ensure_ascii=False, indent=2, sort_keys=True))
