#!/usr/bin/env python3
"""Teaching register allocation: 2 allocatable + 2 reserved physical registers.
Pure integer IR. No network, external packages, or production machine-code claim.
All required checks remain active under python -O.
"""
from collections import deque
from itertools import combinations
import json

REGS=('R0','R1')
SCRATCH=('S0','S1')

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

def names(args):
    return {a for a in args if isinstance(a,str)}

def use_def(ins):
    op=ins[0]
    if op=='input': return set(),{ins[1]}
    if op in ('const','mov'): return names(ins[2:]),{ins[1]}
    if op in ('add','sub'): return names(ins[2:]),{ins[1]}
    if op=='call': return names(ins[2:]),{ins[1]}
    if op=='pcopy': return names([s for d,s in ins[1]]),{d for d,s in ins[1]}
    if op=='br': return names(ins[1:2]),set()
    if op=='ret': return names(ins[1:2]),set()
    if op=='jmp': return set(),set()
    raise ValueError('unknown IR opcode')

def successors(block):
    last=block[-1]
    if last[0]=='jmp':return [last[1]]
    if last[0]=='br':return list(dict.fromkeys(last[2:]))
    check(last[0]=='ret','block requires terminator')
    return []

def reachable(program,entry='E'):
    check(entry in program,'entry absent');seen=set();todo=[entry]
    while todo:
        b=todo.pop()
        if b in seen:continue
        check(b in program and program[b],'empty or missing block')
        seen.add(b);todo.extend(successors(program[b]))
    return [b for b in program if b in seen]

def liveness(program,entry='E'):
    blocks=reachable(program,entry);points=[(b,i) for b in blocks for i in range(len(program[b]))]
    edges={}
    for b,i in points:
        edges[b,i]=[(b,i+1)] if i+1<len(program[b]) else [(c,0) for c in successors(program[b])]
    inside={p:set() for p in points};outside={p:set() for p in points};changed=True
    while changed:
        changed=False
        for p in reversed(points):
            u,d=use_def(program[p[0]][p[1]]);o=set().union(*(inside[q] for q in edges[p]));v=u|(o-d)
            if o!=outside[p] or v!=inside[p]:outside[p]=o;inside[p]=v;changed=True
    return blocks,inside,outside

def validate(program,entry='E'):
    """Finite syntax/definite-assignment checker; loops allowed, unreachable dropped."""
    blocks=reachable(program,entry);allvars=set();pred={b:[] for b in blocks}
    for b in blocks:
        for i,ins in enumerate(program[b]):
            check((ins[0] in ('br','ret','jmp'))==(i==len(program[b])-1),'terminator position')
            check(isinstance(ins,(tuple,list)) and bool(ins),'invalid instruction')
            op=ins[0];arity={'input':3,'const':3,'mov':3,'add':4,'sub':4,'call':3,'pcopy':2,'br':4,'ret':2,'jmp':2}
            check(op in arity and len(ins)==arity[op],'instruction arity')
            u,d=use_def(ins);allvars|=u|d
            check(all(isinstance(x,str) and x and not x.startswith('$R') for x in u|d),'virtual name uses reserved precolor namespace')
            if op=='const':check(type(ins[2]) is int,'const needs mathematical integer')
            if op=='input':check(isinstance(ins[2],str),'input key')
            operands=ins[2:] if op in ('mov','add','sub','call') else ins[1:2] if op in ('br','ret') else [v for _,v in ins[1]] if op=='pcopy' else []
            check(all(type(v) is int or isinstance(v,str) for v in operands),'invalid operand')
            if ins[0]=='pcopy':check(len(d)==len(ins[1]),'duplicate parallel-copy destination')
        for c in successors(program[b]):pred[c].append(b)
    din={b:(set() if b==entry else allvars.copy()) for b in blocks};dout={b:allvars.copy() for b in blocks};changed=True
    while changed:
        changed=False
        for b in blocks:
            incoming=set() if b==entry else set.intersection(*(dout[p] for p in pred[b]))
            out=incoming.copy()
            for ins in program[b]:out|=use_def(ins)[1]
            if din[b]!=incoming or dout[b]!=out:din[b]=incoming;dout[b]=out;changed=True
    for b in blocks:
        state=din[b].copy()
        for ins in program[b]:
            u,d=use_def(ins);check(u<=state,'undefined use at '+b+': '+repr(ins));state|=d
    return {b:program[b] for b in blocks}

def interference(program,k=2,entry='E'):
    blocks,inside,outside=liveness(program,entry);fixed={'$R'+str(j):j for j in range(k)}
    g={x:set() for x in fixed};moves=[]
    def edge(x,y):
        if x!=y:g.setdefault(x,set()).add(y);g.setdefault(y,set()).add(x)
    for x,y in combinations(fixed,2):edge(x,y)
    for b in blocks:
        for i,ins in enumerate(program[b]):
            u,d=use_def(ins)
            for x in u|d:g.setdefault(x,set())
            # Distinct physical destinations make the parallel-copy interface total.
            # This deliberately also separates destinations dead after this copy.
            if ins[0]=='pcopy':
                for x,y in combinations(d,2):edge(x,y)
            pairs=ins[1] if ins[0]=='pcopy' else [(ins[1],ins[2])] if ins[0]=='mov' else []
            for x,y in pairs:
                if isinstance(y,str):moves.append((x,y))
            source=dict(pairs)
            for x in d:
                excluded={x}
                if isinstance(source.get(x),str):excluded.add(source[x])
                for y in outside[b,i]-excluded:edge(x,y)
            if ins[0]=='call':
                for x in outside[b,i]-d:edge(x,'$R0')
    return g,fixed,moves

def locations_from_colors(g,colors,representative=None):
    representative=representative or (lambda x:x)
    slots={};out={}
    for x in sorted(g):
        if x.startswith('$R'):continue
        y=representative(x)
        if y in colors:out[x]='R'+str(colors[y])
        else:
            if y not in slots:slots[y]='slot'+str(len(slots))
            out[x]=slots[y]
    return out

def color_graph(g,fixed,k=2):
    active=set(g);stack=[];trace=[]
    while active-set(fixed):
        remaining=active-set(fixed);degree={x:len(g[x]&active) for x in remaining}
        low=sorted(x for x in remaining if degree[x]<k)
        x=low[0] if low else min(remaining,key=lambda v:(-degree[v],v))
        stack.append(x);trace.append(['simplify' if low else 'potential-spill',x,degree[x]]);active.remove(x)
    colors=dict(fixed)
    for x in reversed(stack):
        used={colors[y] for y in g[x] if y in colors};free=min(set(range(k))-used,default=None)
        if free is not None:colors[x]=free;trace.append(['color',x,free])
        else:trace.append(['actual-spill',x])
    return locations_from_colors(g,colors),trace

def coalesce_graph(g,fixed,moves,k=2):
    """Rescanning IRC variant; each topology change retries blocked moves.
    Precolored representatives have infinite degree for conservative tests.
    """
    graph={x:set(ns) for x,ns in g.items()};alias={x:x for x in g};pending=list(moves);stack=[];trace=[]
    def root(x):
        while alias[x]!=x:x=alias[x]
        return x
    def degree(x):return float('inf') if x in fixed else len(graph[x])
    def remove(x,why):
        stack.append((x,graph[x].copy()));trace.append([why,x,len(graph[x])])
        for y in graph[x]:graph[y].remove(x)
        del graph[x]
    while set(graph)-set(fixed):
        # Normalize aliases and remove satisfied/constrained move preferences.
        todo=[]
        for a,b in pending:
            a,b=root(a),root(b)
            if a==b:trace.append(['satisfied',a]);continue
            if b in graph.get(a,set()) or (a in fixed and b in fixed):trace.append(['constrained',a,b]);continue
            todo.append((a,b))
        pending=todo;related={x for pair in pending for x in pair};ordinary=set(graph)-set(fixed)
        low=sorted(x for x in ordinary if degree(x)<k and x not in related)
        if low:remove(low[0],'simplify');continue
        chosen=None
        for j,(a,b) in enumerate(pending):
            if b in fixed:a,b=b,a
            if a in fixed:
                safe=all(degree(t)<k or t in fixed or t in graph[a] for t in graph[b])
            else:
                safe=sum(degree(t)>=k for t in graph[a]|graph[b])<k
            if safe:chosen=(j,a,b);break
        if chosen:
            j,a,b=chosen;pending.pop(j);trace.append(['coalesce',a,b]);alias[b]=a
            neighbors=(graph[a]|graph[b])-{a,b}
            for t in list(graph[b]):graph[t].discard(b)
            del graph[b];graph[a]=neighbors
            for t in neighbors:graph[t].add(a)
            continue
        low=sorted(x for x in ordinary if degree(x)<k)
        x=low[0] if low else min(ordinary,key=lambda v:(-degree(v),v))
        frozen=[pair for pair in pending if x in pair];pending=[pair for pair in pending if x not in pair]
        if frozen:trace.append(['freeze',x,len(frozen)])
        remove(x,'simplify' if low else 'potential-spill')
    colors=dict(fixed)
    for x,neighbors in reversed(stack):
        used={colors[root(y)] for y in neighbors if root(y) in colors};free=min(set(range(k))-used,default=None)
        if free is not None:colors[x]=free;trace.append(['color',x,free])
        else:trace.append(['actual-spill',x])
    return locations_from_colors(g,colors,root),trace,{x:root(x) for x in sorted(g)}

def intervals(program,entry='E'):
    blocks,inside,outside=liveness(program,entry);positions={};calls=[];ordinal=0
    for b in blocks:
        for i,ins in enumerate(program[b]):
            u,d=use_def(ins)
            for x in inside[b,i]|u:positions.setdefault(x,[]).append(2*ordinal)
            for x in outside[b,i]|d:positions.setdefault(x,[]).append(2*ordinal+1)
            if ins[0]=='call':calls.append(2*ordinal+1)
            ordinal+=1
    ranges={x:(min(ps),max(ps)+1) for x,ps in positions.items()}
    allowed={x:({1} if any(a<c<b for c in calls) else {0,1}) for x,(a,b) in ranges.items()}
    return ranges,allowed

def linear_scan(ranges,allowed=None,k=2):
    allowed=allowed or {x:set(range(k)) for x in ranges};active=[];colors={};spilled=set();trace=[]
    for x in sorted(ranges,key=lambda v:(ranges[v][0],v)):
        start,end=ranges[x];active=[y for y in active if ranges[y][1]>start]
        free=min(allowed[x]-{colors[y] for y in active},default=None)
        if free is not None:colors[x]=free;active.append(x);trace.append(['assign',x,free])
        else:
            eligible=[y for y in active if colors[y] in allowed[x]]
            victim=max(eligible,key=lambda v:(ranges[v][1],v)) if eligible else None
            if victim is not None and ranges[victim][1]>end:
                colors[x]=colors.pop(victim);spilled.add(victim);active.remove(victim);active.append(x);trace.append(['evict-whole',victim,x])
            else:spilled.add(x);trace.append(['spill-whole',x])
    return locations_from_colors(ranges,colors),trace

def check_locations(g,fixed,loc):
    for x,neighbors in g.items():
        lx='R'+str(fixed[x]) if x in fixed else loc[x]
        for y in neighbors:
            ly='R'+str(fixed[y]) if y in fixed else loc[y]
            check(lx!=ly,'interference collision '+x+'/'+y)
    return True

def split_regions(program,entry='E'):
    """One region per original block. Edge adapters have simultaneous copies."""
    blocks,inside,outside=liveness(program,entry);out={};adapters={};mapping={}
    used_names=set().union(*(set.union(*use_def(ins)) for b in blocks for ins in program[b]));used_labels=set(program);edge_names={}
    def fresh(base,used):
        name=base;counter=0
        while name in used:counter+=1;name=base+'#'+str(counter)
        used.add(name);return name
    for b in blocks:
        vs=set().union(*(set.union(*use_def(ins)) for ins in program[b]))|inside[b,0]
        # Values merely live through the block also need a regional name.
        vs|=outside[b,len(program[b])-1]
        mapping[b]={x:fresh(x+'@'+b,used_names) for x in sorted(vs)}
    for b in blocks:
        code=[]
        def operand(x):return mapping[b][x] if isinstance(x,str) else x
        for ins in program[b]:
            op=ins[0]
            if op=='input':new=(op,operand(ins[1]),ins[2])
            elif op=='pcopy':new=(op,[(operand(d),operand(s)) for d,s in ins[1]])
            elif op in ('br','jmp'):
                targets=ins[2:] if op=='br' else ins[1:];labels=[]
                for c in targets:
                    if (b,c) not in edge_names:edge_names[b,c]=fresh('edge:'+b+'>'+c,used_labels)
                    e=edge_names[b,c];labels.append(e)
                    copies=[(mapping[c][x],mapping[b][x]) for x in sorted(inside[c,0])]
                    adapters[e]=[('pcopy',copies),('jmp',c)] if copies else [('jmp',c)]
                new=(op,operand(ins[1]),*labels) if op=='br' else (op,*labels)
            elif op=='ret':new=(op,operand(ins[1]))
            else:new=(op,operand(ins[1]),*(operand(x) for x in ins[2:]))
            code.append(new)
        out[b]=code
    out.update(adapters);return out,mapping

def constant_recipes(program):
    definitions={}
    for b in reachable(program):
        for ins in program[b]:
            for x in use_def(ins)[1]:definitions.setdefault(x,[]).append(ins)
    return {x:ds[0][2] for x,ds in definitions.items() if len(ds)==1 and ds[0][0]=='const' and type(ds[0][2]) is int}

def lower(program,loc,rematerialize=False):
    """Private slots; S0/S1 never hold inter-instruction logical values."""
    recipes=constant_recipes(program) if rematerialize else {};out={};cycle_counter=0
    # Coalesced aliases sharing a spilled home require a family-level proof.
    # This prototype uses only recipes whose home belongs to this one name.
    counts={}
    for h in loc.values():counts[h]=counts.get(h,0)+1
    recipes={x:v for x,v in recipes.items() if counts[loc[x]]==1}
    used_slots=set(loc.values())
    def home(x):return loc[x] if isinstance(x,str) else x
    for b in reachable(program):
        code=[]
        def read(x,scratch):
            if not isinstance(x,str):return x
            h=loc[x]
            if h in REGS:return h
            if x in recipes:code.append(('IMM',scratch,recipes[x]))
            else:code.append(('LD',scratch,h))
            return scratch
        def write(dst,src):
            h=loc[dst]
            if h in REGS:
                if h!=src:code.append(('MOV',h,src))
            else:code.append(('ST',h,src))
        def transfer(dst,src):
            if dst==src:return
            if isinstance(src,str) and src not in REGS+SCRATCH:code.append(('LD','S0',src));src='S0'
            if dst in REGS:code.append(('MOV',dst,src))
            else:code.append(('ST',dst,src))
        for ins in program[b]:
            op=ins[0]
            if op=='input':
                dst=loc[ins[1]];reg=dst if dst in REGS else 'S0';code.append(('IN',reg,ins[2]));write(ins[1],reg)
            elif op=='const' and ins[1] in recipes and loc[ins[1]] not in REGS:pass
            elif op in ('const','mov'):
                src=read(ins[2],'S0');write(ins[1],src)
            elif op in ('add','sub'):
                a=read(ins[2],'S0');c=read(ins[3],'S1');dst=loc[ins[1]];reg=dst if dst in REGS else 'S0';code.append((op.upper(),reg,a,c));write(ins[1],reg)
            elif op=='call':
                a=read(ins[2],'S0')
                if a!='R0':code.append(('MOV','R0',a))
                code.append(('CALL2',));write(ins[1],'R0')
            elif op=='pcopy':
                # Rematerialization is restricted to const definitions, so a target
                # defined by pcopy is not a recipe. Recipe sources are immediates.
                todo={home(d):(recipes[s] if isinstance(s,str) and s in recipes and loc[s] not in REGS else home(s)) for d,s in ins[1] if home(d)!=home(s)}
                while todo:
                    srcs=set(todo.values());safe=min((d for d in todo if d not in srcs),default=None)
                    if safe is not None:d=safe;transfer(d,todo.pop(d))
                    else:
                        d=min(todo);temp='@cycle'+str(cycle_counter);cycle_counter+=1
                        while temp in used_slots:temp='@cycle'+str(cycle_counter);cycle_counter+=1
                        used_slots.add(temp);transfer(temp,d)
                        todo={x:(temp if y==d else y) for x,y in todo.items()}
            elif op=='br':code.append(('BR',read(ins[1],'S0'),ins[2],ins[3]))
            elif op=='jmp':code.append(('JMP',ins[1]))
            elif op=='ret':code.append(('RET',read(ins[1],'S0')))
        out[b]=code
    return out

def run_source(program,inputs,entry='E',limit=100000):
    env={};b=entry;branches=[];steps=0
    def value(x):return env[x] if isinstance(x,str) else x
    while True:
        for ins in program[b]:
            steps+=1
            if steps>limit:raise RuntimeError('source step limit; no termination claim')
            op=ins[0]
            if op=='input':env[ins[1]]=inputs[ins[2]]
            elif op in ('const','mov'):env[ins[1]]=value(ins[2])
            elif op=='add':env[ins[1]]=value(ins[2])+value(ins[3])
            elif op=='sub':env[ins[1]]=value(ins[2])-value(ins[3])
            elif op=='call':env[ins[1]]=2*value(ins[2])
            elif op=='pcopy':env.update({d:value(s) for d,s in ins[1]})
            elif op=='br':chosen=bool(value(ins[1]));branches.append(chosen);b=ins[2] if chosen else ins[3];break
            elif op=='jmp':b=ins[1];break
            elif op=='ret':return {'value':value(ins[1]),'branches':branches,'steps':steps}

def run_target(program,inputs,entry='E',limit=100000):
    regs={};slots={};b=entry;branches=[];counts={};steps=0
    def value(x):return regs[x] if isinstance(x,str) else x
    while True:
        for ins in program[b]:
            op=ins[0];steps+=1;counts[op]=counts.get(op,0)+1
            if steps>limit:raise RuntimeError('target step limit; no termination claim')
            if op=='IN':regs[ins[1]]=inputs[ins[2]]
            elif op in ('MOV','IMM'):regs[ins[1]]=value(ins[2])
            elif op=='LD':regs[ins[1]]=slots[ins[2]]
            elif op=='ST':slots[ins[1]]=value(ins[2])
            elif op=='ADD':regs[ins[1]]=value(ins[2])+value(ins[3])
            elif op=='SUB':regs[ins[1]]=value(ins[2])-value(ins[3])
            elif op=='CALL2':
                result=2*regs['R0'];regs['R0']=result;regs['S0']=987654321;regs['S1']=-987654321
            elif op=='BR':chosen=bool(value(ins[1]));branches.append(chosen);b=ins[2] if chosen else ins[3];break
            elif op=='JMP':b=ins[1];break
            elif op=='RET':return {'value':value(ins[1]),'branches':branches,'counts':counts,'steps':steps}
            else:raise ValueError('unknown target opcode')

def main_program():
    return {'E':[('input','a','a'),('input','flag','flag'),('add','t2','a',1),('add','t3','a',2),('br','flag','T','F')],
            'T':[('add','t4','t2','t3'),('jmp','J')],
            'F':[('sub','t4','t2','t3'),('jmp','J')],
            'J':[('mov','v','t4'),('call','t5','v'),('add','y','a','t5'),('ret','y')]}

def compile_all(program):
    program=validate(program);g,f,moves=interference(program);ranges,allowed=intervals(program)
    out={}
    for kind in ('color','linear','coalesce'):
        aliases=None
        if kind=='color':loc,trace=color_graph(g,f)
        elif kind=='linear':loc,trace=linear_scan(ranges,allowed)
        else:loc,trace,aliases=coalesce_graph(g,f,moves)
        check_locations(g,f,loc);out[kind]={'locations':loc,'trace':trace,'target':lower(program,loc)}
        if aliases is not None:out[kind]['aliases']=aliases
    return out

def main():
    p=main_program();compiled=compile_all(p);values=[];checks=0
    for a in range(-20,21):
        for flag in (0,1):
            inputs={'a':a,'flag':flag};s=run_source(p,inputs)
            for c in compiled.values():
                t=run_target(c['target'],inputs);check((t['value'],t['branches'])==(s['value'],s['branches']),'main compilation mismatch');checks+=1
            if a==3:values.append({'input':inputs,'source':s,'target_counts':{kind:run_target(c['target'],inputs)['counts'] for kind,c in compiled.items()}})
    split,mapping=split_regions(p);split_compiled=compile_all(split)
    for a in range(-20,21):
        for flag in (0,1):
            inp={'a':a,'flag':flag};s=run_source(p,inp);ss=run_source(split,inp);check((s['value'],s['branches'])==(ss['value'],ss['branches']),'split virtual mismatch')
            for c in split_compiled.values():
                t=run_target(c['target'],inp);check((t['value'],t['branches'])==(s['value'],s['branches']),'split target mismatch');checks+=1
    # Optimistic coloring: a 4-cycle has no initial degree<2, yet needs no spill.
    square={x:set() for x in 'abcd'}
    for x,y in [('a','b'),('b','c'),('c','d'),('d','a')]:square[x].add(y);square[y].add(x)
    square_loc,square_trace=color_graph(square,{},2);check(all(x in REGS for x in square_loc.values()),'optimistic square')
    triangle={x:set('abc')-{x} for x in 'abc'};tri_loc,tri_trace=color_graph(triangle,{},2);check(sum(v not in REGS for v in tri_loc.values())==1,'triangle pressure')
    # Equal-value move exclusion does not drop retained dead-write interference.
    dead={'E':[('const','x',7),('const','dead',9),('ret','x')]};dg,df,_=interference(dead);check('x' in dg['dead'],'retained dead write edge')
    wrong=lower(dead,{'x':'R0','dead':'R0'});check(run_target(wrong,{})['value']==9,'dead-write counterexample')
    # Constant rematerialization changes real instructions and memory traffic.
    rem={'E':[('const','c',7),('input','x','x'),('add','u','x','c'),('add','y','u','c'),('ret','y')]};loc={'c':'slot0','x':'R0','u':'R0','y':'R0'}
    ordinary=lower(rem,loc);rebuilt=lower(rem,loc,True)
    for x in range(-20,21):check(run_target(ordinary,{'x':x})['value']==run_target(rebuilt,{'x':x})['value']==x+14,'rematerialization')
    # Structural migrations, with their actual trajectories in the download.
    path={x:set() for x in 'aecbd'}
    for a,b in zip('aecb','ecbd'):path[a].add(b);path[b].add(a)
    irc_examples={}
    for pair in [('a','c'),('a','b')]:
        loc,trace,aliases=coalesce_graph(path,{},[pair]);check_locations(path,{},loc)
        check((loc[pair[0]]==loc[pair[1]])==(pair==('a','c')),'odd/even path move constraint')
        irc_examples['/'.join(pair)]={'locations':loc,'trace':trace,'aliases':aliases}
    adjacent,_=linear_scan({'a':(0,2),'b':(2,3)},k=1)
    check(adjacent=={'a':'R0','b':'R0'},'half-open adjacency')
    loop={'E':[('input','n','n'),('const','a',1),('const','b',2),('jmp','H')],
          'H':[('br','n','B','X')],
          'B':[('pcopy',[('a','b'),('b','a')]),('sub','n','n',1),('jmp','H')],
          'X':[('sub','r','a','b'),('ret','r')]}
    loop_checks=0
    for version in [loop,split_regions(loop)[0]]:
        for candidate in compile_all(version).values():
            for n in range(21):
                check(run_target(candidate['target'],{'n':n})['value']==(-1 if n%2==0 else 1),'loop edge transfer')
                loop_checks+=1
    shared={'E':[('const','c',7),('mov','d','c'),('ret','d')]}
    check(run_target(lower(shared,{'c':'slot0','d':'slot0'},True),{})['value']==7,'shared rematerialization home')
    report={'status':'PASS','physical_registers':list(REGS+SCRATCH),'allocatable':list(REGS),'source':p,'main':compiled,'examples':values,'source_target_checks':checks,'split_source':split,'split':split_compiled,'square':{'locations':square_loc,'trace':square_trace},'triangle':{'locations':tri_loc,'trace':tri_trace},'dead_write_wrong_return':9,'irc_path_examples':irc_examples,'half_open_adjacency':adjacent,'loop_target_checks':loop_checks,'rematerialization':{'ordinary':ordinary,'rebuilt':rebuilt,'ordinary_run':run_target(ordinary,{'x':3}),'rebuilt_run':run_target(rebuilt,{'x':3})}}
    print(json.dumps(report,ensure_ascii=False,indent=2))
if __name__=='__main__':main()
