#!/usr/bin/env python3
"""OPT-8: finite dataflow, actual IR rewrites, SCCP, and independent checks.
Python 3.10+, standard library only. No assert is used for acceptance.
"""
from collections import deque
from copy import deepcopy
import argparse, json

def check(condition, message):
    if not condition: raise RuntimeError(message)

def var(x): return isinstance(x, str)
def uses(i):
    return {x for x in i[2:] if var(x)}
def dest(i): return i[0]
def rhs(i): return (i[1], *i[2:])
def val(x, env): return env[x] if var(x) else x

def op(name, xs):
    if name == 'copy': return xs[0]
    if name == 'add': return xs[0] + xs[1]
    if name == 'mul': return xs[0] * xs[1]
    if name == 'eq': return xs[0] == xs[1]
    raise ValueError(name)

BASE = {
'E': {'code':[('x','add','a','b'),('y','copy','x')], 'term':('br','p','T','F')},
'T': {'code':[('u','add','a','b')], 'term':('jmp','J')},
'F': {'code':[('u','add','a','b')], 'term':('jmp','J')},
'J': {'code':[('z','add','a','b'),('dead','add','z',1),('r','add','u','y')], 'term':('ret','r')},
}
INPUTS = {'a','b','p'}
def successors(b):
    t=b['term']; return list(t[2:]) if t[0]=='br' else ([t[1]] if t[0]=='jmp' else [])
def termuses(b):
    t=b['term']; return {t[1]} if t[0] in ('br','ret') and var(t[1]) else set()
def graph(P):
    seen=set(); todo=['E']
    while todo:
        b=todo.pop()
        if b in seen: continue
        seen.add(b); todo += successors(P[b])
    names=[b for b in P if b in seen]
    pred={b:[] for b in names}
    for b in names:
        for c in successors(P[b]): pred[c].append(b)
    return names,pred

def run(P, inputs):
    env=dict(inputs); b='E'; trace=[]; counts={'add':0,'mul':0,'copy':0,'eq':0,'branches':0}
    # The fixture is acyclic. Detect accidental graph changes rather than treating fuel as divergence.
    visited=set()
    while True:
        check(b not in visited,'fixture acquired a cycle'); visited.add(b)
        trace.append(b)
        for i in P[b]['code']:
            env[i[0]]=op(i[1],[val(x,env) for x in i[2:]])
            counts[i[1]] += 1
        t=P[b]['term']
        if t[0]=='ret': return {'value':val(t[1],env),'path':trace,'counts':counts}
        if t[0]=='jmp': b=t[1]
        else:
            counts['branches']+=1; b=t[2] if val(t[1],env) else t[3]

def live(P):
    names,pred=graph(P); IN={b:set() for b in names}; OUT=deepcopy(IN)
    visits=0
    q=deque(reversed(names)); queued=set(q)
    while q:
        b=q.popleft(); queued.remove(b); visits+=1
        out=set().union(*(IN[s] for s in successors(P[b]))) if successors(P[b]) else set()
        cur=out|termuses(P[b])
        for i in reversed(P[b]['code']): cur=(cur-{dest(i)})|uses(i)
        OUT[b]=out
        if cur!=IN[b]:
            IN[b]=cur
            for p in pred[b]:
                if p not in queued:q.append(p);queued.add(p)
    return IN,OUT,visits

def reaching(P):
    names,pred=graph(P); byvar={}; labels={}
    for b in names:
        for k,i in enumerate(P[b]['code']):
            d=f'{b}.{k+1}'; labels[(b,k)]=d; byvar.setdefault(dest(i),set()).add(d)
    boundary={f'input:{x}' for x in INPUTS}
    for x in INPUTS: byvar.setdefault(x,set()).add(f'input:{x}')
    IN={b:set() for b in names}; OUT=deepcopy(IN); rounds=[]
    while True:
        changed=False
        for b in names:
            new=set().union(*(OUT[p] for p in pred[b])) if pred[b] else set()
            if b=='E':new|=boundary
            cur=set(new)
            for k,i in enumerate(P[b]['code']):cur=(cur-byvar[dest(i)])|{labels[(b,k)]}
            if cur!=OUT[b] or new!=IN[b]:changed=True;IN[b]=new;OUT[b]=cur
        rounds.append({b:sorted(OUT[b]) for b in names})
        if not changed:break
    return IN,OUT,rounds

def expression_universe(P):
    return {rhs(i) for b in P.values() for i in b['code'] if i[1] in ('add','mul','eq')}
def expression_vars(e):return {a for a in e[1:] if var(a)}
def available(P):
    names,pred=graph(P); U=expression_universe(P)
    IN={b:(set() if b=='E' else set(U)) for b in names};OUT={b:set(U) for b in names}
    while True:
        changed=False
        for b in names:
            new=set() if b=='E' else set.intersection(*(OUT[p] for p in pred[b]))
            cur=set(new)
            for i in P[b]['code']:
                cur={e for e in cur if dest(i) not in expression_vars(e)}
                if i[1] in ('add','mul','eq') and dest(i) not in uses(i):cur.add(rhs(i))
            if cur!=OUT[b] or new!=IN[b]:IN[b]=new;OUT[b]=cur;changed=True
        if not changed:return IN,OUT

def copies(P):
    names,pred=graph(P)
    U={(i[0],i[2]) for b in P.values() for i in b['code'] if i[1]=='copy' and var(i[2]) and i[0]!=i[2]}
    IN={b:(set() if b=='E' else set(U)) for b in names};OUT={b:set(U) for b in names}
    def transfer(facts,i):
        cur={p for p in facts if i[0] not in p}
        if i[1]=='copy' and var(i[2]) and i[0]!=i[2]:cur.add((i[0],i[2]))
        return cur
    while True:
        changed=False
        for b in names:
            new=set() if b=='E' else set.intersection(*(OUT[p] for p in pred[b]))
            cur=set(new)
            for i in P[b]['code']:cur=transfer(cur,i)
            if new!=IN[b] or cur!=OUT[b]:IN[b]=new;OUT[b]=cur;changed=True
        if not changed:break
    return IN,OUT,transfer

def copy_propagate(P):
    IN,OUT,transfer=copies(P); Q=deepcopy(P)
    # Analyze and transform exactly the entry-reachable subgraph; preserve other blocks verbatim.
    for b in IN:
        facts=set(IN[b]); newcode=[]
        for i in P[b]['code']:
            lookup=dict(facts)
            # Use one simultaneously justified replacement per original operand.
            newcode.append((i[0],i[1],*(lookup.get(x,x) if var(x) else x for x in i[2:])))
            facts=transfer(facts,i)
        Q[b]['code']=newcode
        t=P[b]['term'];lookup=dict(facts)
        if t[0] in ('br','ret'):Q[b]['term']=(t[0],lookup.get(t[1],t[1]),*t[2:])
    return Q

def cse_fixture(P):
    # Conservative anchored CSE: only reuse E.1's a+b in blocks it dominates,
    # while availability certifies operands unchanged and no write kills x.
    Q=deepcopy(P);IN,OUT=available(P)
    for b in ('T','F','J'):
        check(('add','a','b') in IN[b],'CSE missing availability')
        for k,i in enumerate(Q[b]['code']):
            if rhs(i)==('add','a','b'):Q[b]['code'][k]=(i[0],'copy','x')
    check(all(i[0]!='x' for b in ('T','F','J') for i in P[b]['code']),'anchor overwritten')
    return Q

def dce(P):
    Q=deepcopy(P);removed=[]
    while True:
        IN,OUT,_=live(Q);change=[]
        # Unreachable blocks have no liveness entries and are deliberately retained unchanged.
        for b in IN:
            cur=OUT[b]|termuses(Q[b]);kept=[]
            for i in reversed(Q[b]['code']):
                if i[0] not in cur:
                    # This fixture's operations are total, pure, deterministic mathematical operations.
                    change.append((b,i));continue
                kept.append(i);cur=(cur-{i[0]})|uses(i)
            Q[b]['code']=list(reversed(kept))
        removed.append(change)
        if not change:return Q,removed

BOT=('bottom',);TOP=('top',)
def join(a,b):
    if a==BOT:return b
    if b==BOT:return a
    return a if a==b else TOP

def sccp(branch='constant',false_value=None):
    # SSA fixture with explicit phi incoming-edge identities.
    B={
'E':{'code':[('a','copy',4),('b','copy',6),('c','eq','a',4)],'term':('br','c' if branch=='constant' else 'p','T','F')},
'T':{'code':[('t','add','a','b')],'term':('jmp','J')},
'F':{'code':[('f','copy','q' if false_value is None else false_value)],'term':('jmp','J')},
'J':{'phis':[('v',{'T':'t','F':'f'})],'code':[('w','mul','v',2),('d','eq','w',20)],'term':('br','d','R','B')},
'R':{'code':[],'term':('ret','w')},'B':{'code':[],'term':('ret',99)}}
    values={'p':TOP,'q':TOP};owner={};user={};phis={};instructions={}
    for b,z in B.items():
        for v,inc in z.get('phis',[]):
            values[v]=BOT;owner[v]=b;phis[v]=inc
            for x in inc.values():user.setdefault(x,set()).add(v)
        for i in z['code']:
            values[i[0]]=BOT;owner[i[0]]=b;instructions[i[0]]=i
            for x in uses(i):user.setdefault(x,set()).add(i[0])
        if z['term'][0]=='br':user.setdefault(z['term'][1],set()).add('@'+b)
    edges=set();reachable=set();q=deque([('edge',('START','E'))]);events=[]
    def abstract(x):return values[x] if var(x) else ('const',x)
    def evaluate(v):
        b=owner[v]
        if b not in reachable:return
        if v in phis:
            result=BOT
            for p,x in phis[v].items():
                if (p,b) in edges:result=join(result,abstract(x))
        else:
            i=instructions[v];args=[abstract(x) for x in i[2:]]
            if BOT in args:result=BOT
            elif TOP in args:result=TOP
            else:result=('const',op(i[1],[a[1] for a in args]))
        nv=join(values[v],result)
        if nv!=values[v]:
            values[v]=nv;events.append({'value':v,'abstract':nv})
            for u in sorted(user.get(v,())):q.append(('value',u))
    def branch_visit(b):
        if b not in reachable:return
        t=B[b]['term']
        if t[0]=='jmp':q.append(('edge',(b,t[1])))
        if t[0]=='br':
            c=abstract(t[1])
            if c==TOP:
                for dst in t[2:]:q.append(('edge',(b,dst)))
            elif c!=BOT:q.append(('edge',(b,t[2] if c[1] else t[3])))
    while q:
        kind,item=q.popleft()
        if kind=='edge':
            if item in edges:continue
            edges.add(item);p,b=item;fresh=b not in reachable;reachable.add(b)
            events.append({'edge':list(item)})
            for v,_ in B[b].get('phis',[]):evaluate(v)
            if fresh:
                for i in B[b]['code']:evaluate(i[0])
            branch_visit(b)
        elif item.startswith('@'):branch_visit(item[1:])
        else:evaluate(item)
    return {'values':values,'edges':sorted(edges),'reachable':sorted(reachable),'events':events}

def loop(a,b,n,hoist=False):
    i=s=0;count=0
    if hoist:h=a*b;count+=1
    states=[]
    while i<n:
        if hoist:t=h
        else:t=a*b;count+=1
        s+=t;i+=1;states.append([i,s])
    return {'value':s,'multiplications':count,'states':states}

def dump_sets(d):return {k:sorted(v,key=str) for k,v in d.items()}
def main():
    parser=argparse.ArgumentParser();parser.add_argument('--out');args=parser.parse_args()
    RIN,ROUT,rounds=reaching(BASE);LIN,LOUT,visits=live(BASE);AIN,AOUT=available(BASE)
    check(RIN['J']=={'input:a','input:b','input:p','E.1','E.2','T.1','F.1'},'RD join')
    check(LIN=={'E':{'a','b','p'},'T':{'a','b','y'},'F':{'a','b','y'},'J':{'a','b','u','y'}},'liveness table')
    check(AIN['J']=={('add','a','b')},'availability join')
    cse=cse_fixture(BASE);cp=copy_propagate(cse);opt,deleted=dce(cp)
    check(sum(len(b['code']) for b in opt.values())==2,'two live assignments expected')
    count=0
    for a in range(-4,5):
        for b in range(-4,5):
            for p in (False,True):
                outcomes=[run(P,dict(a=a,b=b,p=p)) for P in (BASE,cse,cp,opt)]
                check(all(o['value']==2*(a+b) for o in outcomes),'optimization output')
                check(len({tuple(o['path']) for o in outcomes})==1,'control changed')
                check(outcomes[0]['counts']['add']==5 and outcomes[-1]['counts']['add']==2,'operation budget')
                count+=1
    # Unreachable generator must not pollute RD.
    unreachable=deepcopy(BASE);unreachable['U']={'code':[('u','copy',99)],'term':('jmp','J')}
    check('U.1' not in reaching(unreachable)[0]['J'],'unreachable pollution')
    unreachable_cp=copy_propagate(unreachable);unreachable_dce,_=dce(unreachable)
    check(unreachable_cp['U']==unreachable['U'] and unreachable_dce['U']==unreachable['U'],'preserve unreachable code')
    for p in (False,True):
        expected_value=run(BASE,{'a':2,'b':3,'p':p})['value']
        check(run(unreachable_cp,{'a':2,'b':3,'p':p})['value']==expected_value,'copy with unreachable U')
        check(run(unreachable_dce,{'a':2,'b':3,'p':p})['value']==expected_value,'DCE with unreachable U')
    empty={'E':{'code':[],'term':('ret',0)}}
    check(live(empty)[0]=={'E':set()} and available(empty)[0]=={'E':set()} and copies(empty)[0]=={'E':set()},'empty fact universes')
    check(copy_propagate(empty)==empty and dce(empty)[0]==empty,'empty transform input')
    # Self-defining expression is killed after it was computed using OLD a.
    selfdef={'E':{'code':[('a','add','a','b')],'term':('ret','a')}}
    check(not available(selfdef)[1]['E'],'self-dependent gen')
    # Killing either side invalidates a copy fact.
    stale={'E':{'code':[('x','copy',1),('y','copy','x'),('x','copy',9)],'term':('ret','y')}}
    check(run(copy_propagate(stale),{})['value']==1,'copy right-side kill')
    # Separate two-pass test: a predecessor's definition dies only after successor deletion.
    chain={'E':{'code':[('a','copy',3)],'term':('jmp','T')},'T':{'code':[('b','add','a',1)],'term':('ret',0)}}
    chainopt,chainrounds=dce(chain)
    check([len(x) for x in chainrounds]==[1,1,0],'iterated DCE')
    S=sccp();check(S['values']['w']==('const',20) and S['reachable']==['E','J','R','T'],'SCCP main')
    both=sccp('input',14);same=sccp('input',10)
    check(both['values']['v']==TOP and both['values']['w']==TOP,'SCCP conflicting phi')
    check(same['values']['v']==('const',10),'SCCP equal phi')
    loopcount=0
    for a in range(-3,4):
        for b in range(-3,4):
            for n in range(6):
                x,y=loop(a,b,n),loop(a,b,n,True)
                check(x['value']==y['value']==n*a*b,'LICM result');check(x['states']==y['states'],'LICM states');loopcount+=1
    # Compute negative witnesses in explicitly separate, richer teaching semantics.
    def observed(thunk):
        try:return {'return':thunk()}
        except ZeroDivisionError:return {'trap':'divide by zero'}
        except KeyError as e:return {'error':'undefined '+str(e.args[0])}
    wrong_stale=deepcopy(stale);wrong_stale['E']['term']=('ret','x')
    missing={'E':{'code':[],'term':('br','p','T','F')},
             'T':{'code':[('t','add','a','b')],'term':('jmp','J')},
             'F':{'code':[('v','add','a','b')],'term':('jmp','J')},
             'J':{'code':[],'term':('ret','t')}}
    def zero_trip_bad():
        h=1//0
        return 0
    def dead_trap():
        dead=1//0
        return 7
    witnesses={'stale_copy':[run(stale,{})['value'],run(wrong_stale,{})['value']],
               'missing_branch_carrier':observed(lambda:run(missing,{'p':False,'a':2,'b':3})['value']),
               'dead_trap':{'source':observed(dead_trap),'incorrect_deleted':observed(lambda:7)},
               'zero_trip_hoist':{'source':observed(lambda:0),'incorrect_hoisted':observed(zero_trip_bad)}}
    check(witnesses['stale_copy']==[1,9],'stale witness')
    check(witnesses['missing_branch_carrier']=={'error':'undefined t'},'missing carrier witness')
    check(witnesses['dead_trap']['source']=={'trap':'divide by zero'},'trap witness')
    # The RD and AE solvers also encounter a genuine back edge.
    cyc={'E':{'code':[('x','copy',0),('t','add','a','b')],'term':('jmp','H')},
         'H':{'code':[],'term':('br','p','L','X')},
         'L':{'code':[('x','add','x',1)],'term':('jmp','H')},
         'X':{'code':[],'term':('ret','x')}}
    cIN,cOUT,cr=reaching(cyc);aeIN,aeOUT=available(cyc)
    check({'E.1','L.1'} <= cIN['H'],'RD loop backedge')
    check(('add','a','b') in aeIN['H'],'must greatest fixed point in loop')
    phi_live={'T':({'v'}-{'v'})|{'x'},'F':({'v'}-{'v'})|{'y'}}
    check(phi_live=={'T':{'x'},'F':{'y'}},'phi edge liveness')
    result={'schema':'OPT-8-v1','finite_validation_inputs':count,'loop_validation_inputs':loopcount,
            'rd_in':dump_sets(RIN),'rd_out':dump_sets(ROUT),'rd_rounds':rounds,
            'live_in':dump_sets(LIN),'live_out':dump_sets(LOUT),'liveness_worklist_visits':visits,
            'available_in':dump_sets(AIN),'available_out':dump_sets(AOUT),
            'ir':{'source':BASE,'after_cse':cse,'after_copy':cp,'after_dce':opt},
            'main_execution':{k:run(P,{'a':2,'b':3,'p':False}) for k,P in [('source',BASE),('optimized',opt)]},
            'dce_removed_by_pass':deleted,'dce_chain_pass_sizes':[len(x) for x in chainrounds],
            'sccp':S,'sccp_migration_conflict':both,'sccp_migration_equal':same,
            'licm_n3':{'source':loop(2,3,3),'hoisted':loop(2,3,3,True)},
            'licm_zero':{'source':loop(2,3,0),'hoisted':loop(2,3,0,True)},
            'unreachable_block_regression':{'copy_preserved':unreachable_cp['U']==unreachable['U'],'dce_preserved':unreachable_dce['U']==unreachable['U'],'both_reachable_paths_preserved':True},'empty_fact_universes':'PASS','negative_witnesses':witnesses,'rd_loop_in':dump_sets(cIN),'ae_loop_in':dump_sets(aeIN),'phi_live_edges':dump_sets(phi_live),
            'limits':'Finite tests supplement proofs. CSE is the documented anchored subset; no LLVM, machine timings, alias analysis, exceptions, or general optimizer certification.'}
    encoded=json.dumps(result,ensure_ascii=False,indent=2)
    if args.out:
        with open(args.out,'w') as f:f.write(encoded+'\n')
    print(f'PASS: {count} diamond inputs, {loopcount} loop inputs, RD/live/AE/copy/DCE/SCCP and boundary tests')
if __name__=='__main__':main()
