#!/usr/bin/env python3
"""Exact-integer teaching machine; not a RISC-V implementation.
In-order rename, tagged ready issue, one result port, ordered retirement.
No loads, branches, devices, concurrency, or speculative memory ordering.
Standard library, stdout only; explicit checks survive python -O.
"""
from collections import deque
from dataclasses import dataclass
import json
import random


def need(ok, reason):
    if not ok:
        raise AssertionError(reason)


@dataclass(frozen=True)
class Instruction:
    op: str
    dst: object = None
    src: tuple = ()
    immediate: object = None


UNITS = {'CONST': 'ALU', 'ADD': 'ALU', 'MUL': 'MUL', 'DIV': 'DIV', 'STORE': 'STORE'}
DEFAULT_LATENCY = {'ALU': 1, 'MUL': 6, 'DIV': 2, 'STORE': 1}


def validate(program, logical):
    program = tuple(program)
    for ins in program:
        if not isinstance(ins, Instruction) or ins.op not in UNITS:
            raise ValueError('unknown instruction')
        arity = 0 if ins.op == 'CONST' else 1 if ins.op == 'STORE' else 2
        if type(ins.src) is not tuple or len(ins.src) != arity:
            raise ValueError('source arity')
        if any(type(s) is not int or not 0 <= s < logical for s in ins.src):
            raise ValueError('source register')
        if ins.op == 'STORE':
            if ins.dst is not None or type(ins.immediate) is not int or ins.immediate < 0:
                raise ValueError('store uses a nonnegative immediate slot address')
        else:
            if type(ins.dst) is not int or not 0 <= ins.dst < logical:
                raise ValueError('destination register')
            if ins.op == 'CONST' and type(ins.immediate) is not int:
                raise ValueError('integer constant')
            if ins.op != 'CONST' and ins.immediate is not None:
                raise ValueError('unexpected immediate')
    return program


def evaluate(ins, operands):
    if ins.op == 'CONST': return ins.immediate, None
    if ins.op == 'ADD': return operands[0] + operands[1], None
    if ins.op == 'MUL': return operands[0] * operands[1], None
    if ins.op == 'DIV':
        return (None, 'division-by-zero') if operands[1] == 0 else (operands[0] // operands[1], None)
    return operands[0], None


@dataclass
class Entry:
    serial: int
    pc: int
    ins: Instruction
    sources: tuple
    destination: object
    stale: object
    phase: str = 'WAIT'
    result: object = None
    error: object = None


class Machine:
    def __init__(self, program, registers, memory=None, physical=None, capacity=8, latency=None, record=False):
        self.initial = list(registers)
        self.logical = len(self.initial)
        if not self.logical or any(type(x) is not int for x in self.initial):
            raise ValueError('nonempty exact-integer registers')
        self.program = validate(program, self.logical)
        self.memory = {} if memory is None else dict(memory)
        if any(type(a) is not int or a < 0 or type(v) is not int for a,v in self.memory.items()):
            raise ValueError('integer memory slots')
        self.initial_memory = dict(self.memory)
        self.physical = self.logical + 5 if physical is None else physical
        if type(self.physical) is not int or self.physical < self.logical + 1:
            raise ValueError('at least one extra physical register')
        if type(capacity) is not int or capacity < 1: raise ValueError('positive ROB capacity')
        self.capacity = capacity
        self.latency = DEFAULT_LATENCY.copy() if latency is None else dict(latency)
        if set(self.latency) != set(DEFAULT_LATENCY) or any(type(v) is not int or v < 1 for v in self.latency.values()):
            raise ValueError('all four positive integer unit latencies')
        self.values = self.initial + [None] * (self.physical - self.logical)
        self.ready = [True] * self.logical + [False] * (self.physical - self.logical)
        self.spec = list(range(self.logical))
        self.committed = list(self.spec)
        self.free = set(range(self.logical, self.physical))
        self.rob = deque()
        self.units = {}  # unit -> (eligible cycle, entry, captured result, captured exception)
        self.fetch = self.retired = self.cycle = self.serial = 0
        self.trap = None
        self.record = record
        self.trace = []
        self.last_events = []
        self.stalls = {'ROB': 0, 'PHYSICAL': 0}
        self.check()

    def event(self, kind, **fields):
        self.last_events.append(dict(kind=kind, **fields))

    def architectural(self):
        return [self.values[p] for p in self.committed], dict(self.memory)

    def dispatch(self):
        if self.fetch == len(self.program): return
        if len(self.rob) == self.capacity:
            self.stalls['ROB'] += 1; return
        ins = self.program[self.fetch]
        if ins.dst is not None and not self.free:
            self.stalls['PHYSICAL'] += 1; return
        # All source tags are read before replacing the destination map, including r <- r+r.
        sources = tuple(self.spec[s] for s in ins.src)
        destination = stale = None
        if ins.dst is not None:
            stale = self.spec[ins.dst]
            destination = min(self.free)
            self.free.remove(destination)
            self.ready[destination] = False
            self.values[destination] = None
            self.spec[ins.dst] = destination
        e = Entry(self.serial, self.fetch, ins, sources, destination, stale)
        self.serial += 1
        self.rob.append(e)
        self.event('rename', pc=e.pc, sources=list(sources), destination=destination, stale=stale)
        self.fetch += 1

    def issue(self):
        # deque iteration preserves age; a blocked older instruction does not block other units.
        for e in self.rob:
            unit = UNITS[e.ins.op]
            if e.phase == 'WAIT' and unit not in self.units and all(self.ready[p] for p in e.sources):
                operands = tuple(self.values[p] for p in e.sources)
                result, error = evaluate(e.ins, operands)
                e.phase = 'RUN'
                at = self.cycle + self.latency[unit]
                self.units[unit] = (at, e, result, error)
                self.event('issue', pc=e.pc, operands=list(operands), unit=unit, eligible=at)
                return

    def publish(self):
        eligible = [(value[1].serial, unit) for unit,value in self.units.items() if value[0] <= self.cycle]
        if not eligible: return
        _,unit = min(eligible)
        _,e,result,error = self.units.pop(unit)
        need(e.phase == 'RUN' and any(x is e for x in self.rob), 'live producer identity')
        e.result, e.error, e.phase = result, error, 'DONE'
        if e.destination is not None and error is None:
            self.values[e.destination] = result
            self.ready[e.destination] = True
        # A STORE result is only prepared data. Memory is untouched here.
        self.event('publish', pc=e.pc, destination=e.destination, result=result, error=error)

    def retire(self):
        if not self.rob or self.rob[0].phase != 'DONE': return
        e = self.rob[0]
        if e.error is not None:
            self.trap = dict(pc=e.pc, cause=e.error)
            killed = [x.pc for x in self.rob]
            self.units.clear()  # cancel outstanding completion records before reusing any tag
            self.rob.clear()
            self.spec = list(self.committed)
            self.free = set(range(self.physical)) - set(self.committed)
            for p in self.free: self.values[p] = None; self.ready[p] = False
            self.event('trap', pc=e.pc, cause=e.error, cancelled=killed)
            return
        self.rob.popleft()
        need(e.pc == self.retired, 'ordered prefix')
        if e.destination is not None:
            need(self.committed[e.ins.dst] == e.stale, 'stale tag is preceding committed version')
            self.committed[e.ins.dst] = e.destination
            self.free.add(e.stale)
            self.values[e.stale] = None
            self.ready[e.stale] = False
        else:
            self.memory[e.ins.immediate] = e.result
        self.retired += 1
        self.event('retire', pc=e.pc, destination=e.destination, released=e.stale,
                   store=None if e.ins.op != 'STORE' else [e.ins.immediate,e.result])

    def step(self):
        if self.trap is not None or (self.fetch == len(self.program) and not self.rob):
            return False
        self.cycle += 1
        self.last_events = []
        self.retire()
        if self.trap is None:
            self.publish(); self.issue(); self.dispatch()
        self.check()
        if self.record: self.trace.append(dict(cycle=self.cycle, events=self.last_events))
        return True

    def check(self):
        need(len(set(self.committed)) == self.logical, 'committed map injective')
        need(len(set(self.spec)) == self.logical, 'speculative map injective')
        need(not set(self.committed) & self.free, 'committed tags not free')
        need(not set(self.spec) & self.free, 'speculative tags not free')
        expected = list(self.committed)
        allocated = set(self.committed)
        for e in self.rob:
            need(not set(e.sources) & self.free, 'source lifetime')
            if e.destination is not None:
                need(e.destination not in allocated and e.stale == expected[e.ins.dst], 'single fresh version')
                allocated.add(e.destination)
                expected[e.ins.dst] = e.destination
        need(expected == self.spec, 'map equals ordered pending writes')
        need(allocated | self.free == set(range(self.physical)) and not allocated & self.free,
             'no leaked or double-owned physical register')
        need(all(self.ready[p] for p in self.committed), 'committed values ready')
        need(len(self.rob) <= self.capacity, 'ROB capacity')
        need(len(self.units) <= 4, 'one operation per nonpipelined unit')
        if self.trap is not None: need(not self.rob and not self.units, 'flush cancels everything')

    def run(self):
        # Defensive test watchdog, not a advertised cycle bound.
        budget = (len(self.program)+1)**2 * (max(self.latency.values())+10)
        while self.step():
            need(self.cycle <= budget, 'unexpected lack of progress')
        return self


def sequential(program, registers, memory):
    regs,mem = list(registers),dict(memory)
    committed = []
    for pc,ins in enumerate(program):
        args = [regs[s] for s in ins.src]
        # Independent instruction switch for the sequential specification.
        if ins.op == 'DIV' and args[1] == 0:
            return regs,mem,dict(pc=pc,cause='division-by-zero'),committed
        if ins.op == 'CONST': value = ins.immediate
        elif ins.op == 'ADD': value = sum(args)
        elif ins.op == 'MUL': value = args[0] * args[1]
        elif ins.op == 'DIV': value = args[0] // args[1]
        else: value = args[0]
        if ins.op == 'STORE': mem[ins.immediate] = value
        else: regs[ins.dst] = value
        committed.append((list(regs),dict(mem)))
    return regs,mem,None,committed


def verify(machine):
    regs,mem,trap,prefixes = sequential(machine.program,machine.initial,machine.initial_memory)
    while machine.step():
        need(machine.cycle <= (len(machine.program)+1)**2 * (max(machine.latency.values())+10),
             'test watchdog: unexpected lack of progress')
        expected = prefixes[machine.retired-1] if machine.retired else (machine.initial,machine.initial_memory)
        need(machine.architectural() == expected, 'every visible state is the sequential retired prefix')
    need(machine.architectural() == (regs,mem) and machine.trap == trap, 'sequential terminal semantics')
    return machine


def examples():
    initial=[0,2,3,4,0,9]
    normal=[Instruction('MUL',1,(2,3)),Instruction('ADD',4,(1,2)),Instruction('CONST',1,(),7),
            Instruction('ADD',5,(1,3)),Instruction('STORE',None,(4,),0)]
    fault=[Instruction('MUL',1,(2,3)),Instruction('DIV',4,(2,0)),Instruction('CONST',1,(),99),
           Instruction('STORE',None,(3,),0),Instruction('ADD',5,(2,3))]
    machines=[]
    for name,program in [('normal',normal),('fault',fault)]:
        d=verify(Machine(program,initial,{0:31},physical=11,record=True))
        machines.append(dict(name=name,cycles=d.cycle,registers=d.architectural()[0],memory=d.memory,
                             trap=d.trap,committed_map=d.committed,speculative_map=d.spec,free=sorted(d.free),trace=d.trace))
    variants=[]
    for p,q in [(7,1),(7,8),(8,2),(11,8)]:
        d=verify(Machine(normal,initial,{0:31},physical=p,capacity=q))
        variants.append(dict(physical=p,capacity=q,cycles=d.cycle,stalls=d.stalls))
    return machines,variants


def randomized():
    rng=random.Random(6211026);cases=cycles=0
    for _ in range(1600):
        logical=rng.randrange(1,7);initial=[rng.randrange(-4,5) for _ in range(logical)]
        program=[]
        for _ in range(rng.randrange(18)):
            op=rng.choice(list(UNITS));dst=rng.randrange(logical);src=lambda:rng.randrange(logical)
            if op=='STORE':ins=Instruction(op,None,(src(),),rng.randrange(3))
            elif op=='CONST':ins=Instruction(op,dst,(),rng.randrange(-6,7))
            else:ins=Instruction(op,dst,(src(),src()))
            program.append(ins)
        lat={u:rng.randrange(1,8) for u in DEFAULT_LATENCY}
        d=verify(Machine(program,initial,{0:17},physical=logical+rng.randrange(1,7),
                         capacity=rng.randrange(1,9),latency=lat))
        cases+=1;cycles+=d.cycle
    return dict(programs=cases,cycles=cycles)


def boundaries():
    d=verify(Machine([Instruction('ADD',0,(0,0)),Instruction('ADD',0,(0,0))],[3],physical=2))
    need(d.architectural()[0]==[12], 'self-source before destination map')
    d=verify(Machine([Instruction('DIV',0,(0,1))],[-7,3],physical=3))
    need(d.architectural()[0]==[-3,3], 'floor division, not truncation')
    late=[Instruction('DIV',0,(1,2)),Instruction('MUL',1,(1,1))]
    d=verify(Machine(late,[8,2,0],physical=5,latency={'ALU':1,'MUL':30,'DIV':2,'STORE':1},record=True))
    need(d.trap is not None and not d.units and d.architectural()[0]==[8,2,0], 'cancel unfinished long operation')
    failures=[]
    for name,fn in [('no spare register',lambda:Machine([], [0],physical=1)),
                    ('zero ROB',lambda:Machine([], [0],capacity=0)),
                    ('zero latency',lambda:Machine([], [0],latency={'ALU':0,'MUL':1,'DIV':1,'STORE':1})),
                    ('bool register',lambda:Machine([], [True])),
                    ('bad source',lambda:Machine([Instruction('ADD',0,(0,1))],[1]))]:
        try:fn()
        except ValueError:failures.append(name)
        else:raise AssertionError('invalid input accepted')
    return dict(self_source_result=12,negative_division=-3,cancelled_long_operation_cycle=d.cycle,
                cancelled_long_operation_trace=d.trace,rejected=failures)


def counterexamples():
    initial=[0,2,3,4,0,9]
    normal=[Instruction('MUL',1,(2,3)),Instruction('ADD',4,(1,2)),Instruction('CONST',1,(),7),
            Instruction('ADD',5,(1,3)),Instruction('STORE',None,(4,),0)]
    fault=[Instruction('MUL',1,(2,3)),Instruction('DIV',4,(2,0)),Instruction('CONST',1,(),99),
           Instruction('STORE',None,(3,),0),Instruction('ADD',5,(2,3))]
    out={}
    d=Machine([Instruction('ADD',0,(0,0))],[3],physical=2)
    d.step();e=d.rob[0];e.sources=(e.destination,e.destination)  # MUTANT: map updated before source capture.
    for _ in range(8):d.step()
    need(e.phase=='WAIT' and not d.ready[e.destination], 'self-dependent mutant remains blocked')
    out['destination_before_sources']=dict(phase=e.phase,source_tags=list(e.sources),ready=False)
    d=Machine(normal,initial,{0:31})
    while d.cycle<5:d.step()
    e=next(e for e in d.rob if e.pc==1)
    e.sources=(d.spec[1],d.spec[2])  # MUTANT: reread current names instead of captured producer tags.
    d.run();need(d.architectural()[0][4]==10 and d.memory[0]==10, 'wrong version changes both outputs')
    out['reread_latest_name']=dict(r4=10,memory0=10,correct_r4=15)
    d=Machine(fault,initial,{0:31})
    while d.cycle<7:d.step()
    store=next(e for e in d.rob if e.ins.op=='STORE')
    need(store.phase=='DONE' and d.memory[0]==31,'store prepared but not visible')
    d.memory[store.ins.immediate]=store.result  # MUTANT: write at execution completion.
    d.run();need(d.trap is not None and d.memory[0]==4,'early store survives flush incorrectly')
    out['store_before_retirement']=dict(memory0=4,correct_memory0=31,trap=d.trap)
    d=Machine(fault,initial,{0:31})
    while d.cycle<5:d.step()
    need(any(e.error for e in d.rob) and d.retired==0,'younger exception detected first')
    older_phase=d.rob[0].phase
    d.trap=dict(pc=1,cause='division-by-zero')  # MUTANT: flush at discovery, discarding older work.
    d.units.clear();d.rob.clear();d.spec=list(d.committed)
    d.free=set(range(d.physical))-set(d.committed)
    for p in d.free:d.values[p]=None;d.ready[p]=False
    d.check()
    out['report_first_detected_exception']=dict(observed_r1=d.architectural()[0][1],required_r1=12,
                                              discarded_older_multiply_phase=older_phase)
    d=Machine(normal,initial,{0:31})
    while d.cycle<5:d.step()
    younger=next(e for e in d.rob if e.pc==2)
    d.free.add(younger.stale)  # MUTANT: free old version at younger completion, not retirement.
    try:d.check()
    except AssertionError as exc:out['free_stale_at_completion']=str(exc)
    else:raise AssertionError('early free escaped lifetime check')
    return out


def main():
    traces,variants=examples()
    print(json.dumps(dict(examples=traces,capacity_migrations=variants,random=randomized(),boundaries=boundaries(),
                         counterexamples=counterexamples()),
                     ensure_ascii=False,indent=2))


if __name__=='__main__':main()
