#!/usr/bin/env python3
"""Scoped SSA value numbering and single-target edge-local PRE.

Program = {entry, params, blocks}; block = {phis, code, term}.
Instruction = (destination, operation, operands); operands are names or ints.
Phi = (destination, ((predecessor, operand), ...)); phis read simultaneously.
Terms: ('jump', block), ('branch', value, true_block, false_block),
       ('return', tuple(values)). Integer zero is false, nonzero is true.
All arithmetic is pure and total over mathematical integers.
"""
from copy import deepcopy
import json
import random

OPS = {'copy':1,'add':2,'sub':2,'mul':2,'lt':2,'eq':2}


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


def is_name(x): return type(x) is str and bool(x)


def names(args): return {x for x in args if type(x) is str}


def successors(block):
    term=block['term']
    if term[0]=='jump': return (term[1],)
    if term[0]=='branch': return tuple(dict.fromkeys(term[2:]))
    return ()


def reads_term(term):
    if term[0]=='branch':return (term[1],)
    if term[0]=='return':return term[1]
    return ()


def graph(program):
    require(type(program) is dict and set(program)=={'entry','params','blocks'},'program shape')
    blocks=program['blocks'];params=program['params'];entry=program['entry']
    require(type(blocks) is dict and blocks and all(is_name(b) for b in blocks),'blocks')
    require(type(params) is tuple and all(is_name(p) for p in params) and len(set(params))==len(params),'parameters')
    require(entry in blocks,'entry')
    pred={b:[] for b in blocks}
    for b,block in blocks.items():
        require(type(block) is dict and set(block)=={'phis','code','term'},'block shape')
        require(type(block['phis']) is tuple and type(block['code']) is tuple,'instruction lists')
        for x,op,args in block['code']:
            require(is_name(x) and op in OPS and type(args) is tuple and len(args)==OPS[op],'instruction')
            require(all(is_name(v) or type(v) is int for v in args),'operand')
        for x,slots in block['phis']:
            require(is_name(x) and type(slots) is tuple,'phi')
            require(all(type(s) is tuple and len(s)==2 and is_name(s[0]) and (is_name(s[1]) or type(s[1]) is int) for s in slots),'phi slot')
        t=block['term'];require(type(t) is tuple and t,'terminator')
        require((t[0]=='jump' and len(t)==2) or (t[0]=='branch' and len(t)==4 and (is_name(t[1]) or type(t[1]) is int)) or (t[0]=='return' and len(t)==2 and type(t[1]) is tuple and all(is_name(v) or type(v) is int for v in t[1])),'terminator shape')
        for c in successors(block):
            require(c in blocks,'successor')
            pred[c].append(b)
    require(not pred[entry] and not blocks[entry]['phis'],'artificial entry has no incoming edge/phi')
    reached=set();post=[];pending=[(entry,False)]
    while pending:
        b,leaving=pending.pop()
        if leaving:post.append(b);continue
        if b in reached:continue
        reached.add(b);pending.append((b,True))
        pending.extend((c,False)for c in reversed(successors(blocks[b])))
    require(reached==set(blocks),'unreachable block')
    for b,block in blocks.items():
        for _,slots in block['phis']:
            require(len(slots)==len(pred[b]) and {p for p,_ in slots}==set(pred[b]),'exact phi predecessor slots')
    return pred,list(reversed(post))


def dominators(program,pred=None):
    if pred is None:pred,_=graph(program)
    entry=program['entry'];all_blocks=set(program['blocks'])
    dom={b:({b} if b==entry else set(all_blocks)) for b in all_blocks};rounds=0
    while True:
        nxt={entry:{entry}}
        for b in program['blocks']:
            if b!=entry:nxt[b]={b}|set.intersection(*(dom[p] for p in pred[b]))
        rounds+=1
        if nxt==dom:break
        dom=nxt
    return dom,rounds


def validate_ssa(program):
    pred,rpo=graph(program);dom,rounds=dominators(program,pred)
    defs={p:(program['entry'],-2) for p in program['params']}
    for b,block in program['blocks'].items():
        for x,_ in block['phis']:
            require(x not in defs,'duplicate SSA name');defs[x]=(b,-1)
        for i,(x,_,_) in enumerate(block['code']):
            require(x not in defs,'duplicate SSA name');defs[x]=(b,i)
    def available(v,b,index):
        if type(v) is int:return
        require(v in defs,'unknown SSA name')
        d,pos=defs[v];require(d in dom[b] and (d!=b or pos<index),'definition does not dominate use')
    for b,block in program['blocks'].items():
        for _,slots in block['phis']:
            for p,v in slots:available(v,p,len(program['blocks'][p]['code']))
        for i,(_,_,args) in enumerate(block['code']):
            for v in args:available(v,b,i)
        for v in reads_term(block['term']):available(v,b,len(block['code']))
    return pred,rpo,dom,rounds


def validate_mutable(program):
    pred,_=graph(program);require(all(not b['phis'] for b in program['blocks'].values()),'PRE requires phi-free IR')
    all_names=set(program['params'])|{x for b in program['blocks'].values() for x,_,_ in b['code']}
    entry=program['entry'];out={b:set(all_names) for b in program['blocks']}
    while True:
        new={};ins={}
        for b,block in program['blocks'].items():
            ins[b]=set(program['params']) if b==entry else set.intersection(*(out[p] for p in pred[b]))
            new[b]=ins[b]|{x for x,_,_ in block['code']}
        if new==out:break
        out=new
    for b,block in program['blocks'].items():
        have=set(ins[b])
        for x,_,args in block['code']:
            require(names(args)<=have,'possibly uninitialized operand');have.add(x)
        require(names(reads_term(block['term']))<=have,'possibly uninitialized terminator')
    return pred


def calculate(op,args):
    if op=='copy':return args[0]
    if op=='add':return args[0]+args[1]
    if op=='sub':return args[0]-args[1]
    if op=='mul':return args[0]*args[1]
    if op=='lt':return int(args[0]<args[1])
    if op=='eq':return int(args[0]==args[1])
    raise ValueError('unknown operation')


def execute(program,inputs,block_budget=1000,selected=None):
    graph(program);require(type(inputs) is dict and set(inputs)==set(program['params']) and all(type(v) is int for v in inputs.values()),'inputs')
    require(type(block_budget) is int and block_budget>=0,'budget')
    env=dict(inputs);b=program['entry'];previous=None;trace=[];counts={op:0 for op in OPS};selected_count=0
    def value(v):
        if type(v) is int:return v
        require(v in env,'uninitialized '+v);return env[v]
    for _ in range(block_budget):
        block=program['blocks'][b]
        updates={x:value(dict(slots)[previous]) for x,slots in block['phis']}
        env.update(updates)
        for x,op,args in block['code']:
            result=calculate(op,tuple(value(v) for v in args));counts[op]+=1
            if selected==(op,args):selected_count+=1
            env[x]=result
        trace.append({'block':b,'state':dict(env)})
        t=block['term']
        if t[0]=='return':return {'status':'returned','value':tuple(value(v)for v in t[1]),'trace':trace,'counts':counts,'selected_count':selected_count}
        c=t[1] if t[0]=='jump' else (t[2] if value(t[1])!=0 else t[3])
        previous,b=b,c
    return {'status':'budget_exhausted','trace':trace,'counts':counts,'selected_count':selected_count}


def scoped_value_numbering(program):
    pred,rpo,dom,rounds=validate_ssa(program);result=deepcopy(program)
    rank={b:i for i,b in enumerate(rpo)};children={b:[] for b in program['blocks']}
    for b in program['blocks']:
        if b==program['entry']:continue
        strict=dom[b]-{b};parent=max(strict,key=lambda d:len(dom[d]));children[parent].append(b)
    for c in children.values():c.sort(key=rank.__getitem__)
    representative={p:p for p in program['params']};table={};evidence=[];phi_evidence=[]
    def canonical(v):return v if type(v) is int else representative.get(v,v)
    def known(v):return type(v) is int or v in representative
    def process(b):
        introduced=[];local_phis={}
        for x,slots in program['blocks'][b]['phis']:
            available=all(known(v) for _,v in slots)
            actual=tuple(sorted((p,canonical(v)) for p,v in slots))
            if available and len({v for _,v in actual})==1:
                representative[x]=actual[0][1];reason='all inputs equal'
            elif available and actual in local_phis:
                representative[x]=local_phis[actual];reason='same-block phi'
            else:
                representative[x]=x;reason='fresh unresolved phi' if not available else 'fresh phi'
                if available:local_phis[actual]=x
            phi_evidence.append({'block':b,'name':x,'inputs_known':available,'representative':representative[x],'reason':reason,'slots':actual})
        code=[]
        for x,op,args in program['blocks'][b]['code']:
            new_args=tuple(canonical(v)for v in args)
            if op=='copy':
                representative[x]=new_args[0];code.append((x,op,new_args));continue
            key=(op,new_args)
            if key in table:
                representative[x]=table[key];code.append((x,'copy',(table[key],)))
                evidence.append({'block':b,'destination':x,'key':key,'carrier':table[key]})
            else:
                representative[x]=x;table[key]=x;introduced.append(key);code.append((x,op,new_args))
        result['blocks'][b]['code']=tuple(code)
        return introduced
    pending=[(program['entry'],None)]
    while pending:
        b,trail=pending.pop()
        if trail is not None:
            for key in trail:del table[key]
        else:
            trail=process(b);pending.append((b,trail))
            pending.extend((c,None)for c in reversed(children[b]))
    # Preserve every source definition. Final uses may follow solved copy/phi aliases.
    for block in result['blocks'].values():
        block['phis']=tuple((x,tuple((p,canonical(v)) for p,v in slots))for x,slots in block['phis'])
        t=block['term']
        if t[0]=='branch':block['term']=('branch',canonical(t[1]),t[2],t[3])
        elif t[0]=='return':block['term']=('return',tuple(canonical(v)for v in t[1]))
    validate_ssa(result)
    return {'program':result,'representatives':representative,'replacements':evidence,'phis':phi_evidence,'dominator_rounds':rounds,'order':rpo}


def availability(program,expression):
    pred=validate_mutable(program);op,args=expression
    require(op in OPS and op!='copy' and len(args)==OPS[op],'selected expression')
    used=names(args);entry=program['entry'];out={b:True for b in program['blocks']};rounds=0
    def transfer(block,available):
        for x,o,a in block['code']:
            if x in used:available=False
            if (o,a)==expression and x not in used:available=True
        return available
    while True:
        ins={b:(False if b==entry else all(out[p]for p in pred[b]))for b in program['blocks']}
        new={b:transfer(block,ins[b]) for b,block in program['blocks'].items()};rounds+=1
        if new==out:return {'in':ins,'out':new,'rounds':rounds}
        out=new


def edge_local_pre(program,target):
    pred=validate_mutable(program)
    require(target in program['blocks'] and program['blocks'][target]['code'],'target entry instruction')
    destination,op,args=program['blocks'][target]['code'][0]
    require(op!='copy','target must compute an operation')
    expression=(op,args);facts=availability(program,expression)
    good=[p for p in pred[target] if facts['out'][p]]
    missing=[p for p in pred[target] if not facts['out'][p]]
    if not good:return {'status':'no_available_edge','program':deepcopy(program),'expression':expression,'availability':facts,'insertions':[],'carrier_sites':[]}
    result=deepcopy(program);taken=set(program['params'])|set(program['blocks'])|{x for b in program['blocks'].values()for x,_,_ in b['code']}
    counters={}
    def fresh(prefix):
        i=counters.get(prefix,0)
        while prefix+str(i) in taken:i+=1
        counters[prefix]=i+1
        x=prefix+str(i);taken.add(x);return x
    carrier=fresh('__pre_value_');sites=[];insertions=[]
    for b,block in program['blocks'].items():
        code=[]
        for i,(x,o,a) in enumerate(block['code']):
            if b==target and i==0:code.append((x,'copy',(carrier,)))
            elif (o,a)==expression:
                code.extend([(carrier,o,a),(x,'copy',(carrier,))]);sites.append((b,i))
            else:code.append((x,o,a))
        result['blocks'][b]['code']=tuple(code)
    for p in missing:
        critical=len(successors(program['blocks'][p]))>1 and len(pred[target])>1
        bridge=fresh('__pre_edge_')
        result['blocks'][bridge]={'phis':(),'code':((carrier,op,args),),'term':('jump',target)}
        t=result['blocks'][p]['term']
        if t[0]=='jump':result['blocks'][p]['term']=('jump',bridge)
        else:result['blocks'][p]['term']=('branch',t[1],bridge if t[2]==target else t[2],bridge if t[3]==target else t[3])
        insertions.append({'from':p,'to':target,'block':bridge,'critical':critical})
    validate_mutable(result)
    return {'status':'fully_redundant' if not missing else 'partially_redundant','program':result,'expression':expression,'availability':facts,'carrier':carrier,'carrier_sites':sites,'insertions':insertions}


def block(code=(),term=('return',()),phis=()):return {'phis':tuple(phis),'code':tuple(code),'term':term}


def value_example():
    return {'entry':'E','params':('a','b','c','p'),'blocks':{
        'E':block((('t','add',('a','b')),('u','copy',('a',))),('branch','p','T','F')),
        'T':block((('x','add',('u','b')),('left','mul',('x','c'))),('jump','J')),
        'F':block((('y','add',('a','b')),('right','mul',('y','c'))),('jump','J')),
        'J':block((('r','add',('z','w')),('k','add',('w','z')),('s','add',('t','t'))),('return',('r','k','s')),
                  (('z',(('T','left'),('F','right'))),('w',(('T','left'),('F','right')))))}}


def pre_example():
    return {'entry':'E','params':('a','b','p','q'),'blocks':{
        'E':block((),('branch','p','T','F')),
        'T':block((('x','add',('a','b')),('x','copy',(0,))),('jump','J')),
        'F':block((),('branch','q','J','X')),
        'J':block((('z','add',('a','b')),('answer','mul',('z',2))),('return',('answer',))),
        'X':block((),('return',(-1,)))}}



def value_loop():
    return {'entry':'E','params':('n',),'blocks':{
        'E':block((),('jump','H')),
        'H':block((('test','lt',('i','n')),),('branch','test','B','X'),
                  (('i',(('E',0),('B','next'))),('sum',(('E',0),('B','acc'))))),
        'B':block((('one','add',('i',1)),('again','add',('i',1)),
                   ('acc','add',('sum','again')),('next','add',('i',1))),('jump','H')),
        'X':block((),('return',('sum',)))}}


def edge_loop(kill=False):
    # The first visit uses E->H->J. Later visits come directly from B->J.
    # Keeping the final test in B makes entry H->J an actual critical edge.
    body=(('sum','add',('sum','z')),('i','add',('i',1)))
    if kill:body+=(('b','add',('b',1)),)
    return {'entry':'E','params':('a','b','n'),'blocks':{
        'E':block((('i','copy',(0,)),('sum','copy',(0,))),('jump','H')),
        'H':block((('test','lt',('i','n')),),('branch','test','J','X')),
        'J':block((('z','add',('a','b')),),('jump','B')),
        'B':block(body+(('test','lt',('i','n')),),('branch','test','J','X')),
        'X':block((),('return',('sum',)))}}


def macro_trace(run,original):
    visible=set(original['params'])|{x for b in original['blocks'].values()for x,_,_ in b['code']}|{x for b in original['blocks'].values()for x,_ in b['phis']}
    return [{'block':row['block'],'state':{k:v for k,v in row['state'].items() if k in visible}}
            for row in run['trace'] if row['block'] in original['blocks']]


def random_ssa(rng,diamonds):
    blocks={};params=('a','b','c')+tuple('p'+str(i)for i in range(diamonds));current='E';previous='a'
    for i in range(diamonds):
        t,f,j='T'+str(i),'F'+str(i),'J'+str(i);v='v'+str(i);cp='cp'+str(i)
        blocks[current]=block(( (cp,'copy',(previous,)),(v,'add',(cp,'b')) ),('branch','p'+str(i),t,f),blocks.get(current,{}).get('phis',()))
        left,right='left'+str(i),'right'+str(i);lx,rx='lx'+str(i),'rx'+str(i)
        blocks[t]=block(((lx,'add',(previous,'b')),(left,'mul',(lx,rng.randrange(-3,4)))),('jump',j))
        blocks[f]=block(((rx,'add',(cp,'b')),(right,'mul',(rx,rng.randrange(-3,4)))),('jump',j))
        z,w='z'+str(i),'w'+str(i)
        blocks[j]=block((),('return',(z,w)),((z,((t,left),(f,right))),(w,((t,left),(f,right)))))
        current,previous=j,w
    return {'entry':'E','params':params,'blocks':blocks}


def encode(v):
    if type(v) is dict:return {k:encode(x)for k,x in v.items()}
    if type(v) in (tuple,list):return [encode(x)for x in v]
    if type(v) is set:return [encode(x)for x in sorted(v)]
    return v


def self_test():
    value=value_example();numbered=scoped_value_numbering(value)
    samples=[]
    for p in (0,1):
        inputs={'a':2,'b':3,'c':4,'p':p};before=execute(value,inputs);after=execute(numbered['program'],inputs)
        require(before['value']==after['value']==(40,40,10),'value example')
        require(sum(v for k,v in before['counts'].items()if k!='copy')==6 and sum(v for k,v in after['counts'].items()if k!='copy')==4,'dynamic arithmetic')
        require(before['trace']==after['trace'],'all source states preserved')
        samples.append({'inputs':inputs,'before':before,'after':after})
    source=pre_example();pre=edge_local_pre(source,'J');cases=[]
    for p,q in ((1,0),(0,1),(0,0)):
        inp={'a':2,'b':3,'p':p,'q':q};before=execute(source,inp,selected=pre['expression']);after=execute(pre['program'],inp,selected=pre['expression'])
        require(before['value']==after['value'],'PRE value')
        require(after['selected_count']<=before['selected_count'],'PRE count')
        require(before['trace']==macro_trace(after,source),'PRE macro states')
        cases.append({'inputs':inp,'before':before,'after':after})
    loops=[];loop=value_loop();opt=scoped_value_numbering(loop)
    for n in range(13):
        original=execute(loop,{'n':n});changed=execute(opt['program'],{'n':n})
        require(original['value']==changed['value']==(n*(n+1)//2,),'loop values')
        require(original['trace']==changed['trace'],'SSA loop states')
        require(original['counts']['add']==4*n and changed['counts']['add']==2*n,'SSA loop add counts')
        if n in (0,3):loops.append({'n':n,'before':original,'after':changed})
    pre_loops=[];loop_checks=0
    for kill in (False,True):
        program=edge_loop(kill);transformed=edge_local_pre(program,'J')
        require(transformed['status']==('no_available_edge' if kill else 'partially_redundant'),'loop kill placement')
        for n in range(13):
            inp={'a':2,'b':3,'n':n};before=execute(program,inp,selected=transformed['expression']);after=execute(transformed['program'],inp,selected=transformed['expression'])
            expected=5*n+(n*(n-1)//2 if kill else 0)
            require(before['value']==after['value']==(expected,),'PRE loop output')
            require(before['trace']==macro_trace(after,program),'PRE loop states')
            require(after['selected_count']==(n if kill else int(n>0)),'PRE loop counts')
            loop_checks+=1
            if n in (0,3):pre_loops.append({'kill':kill,'n':n,'transformation':transformed,'before':before,'after':after})
    # Broken sibling reuse reads a name never defined on the taken branch.
    sibling={'entry':'E','params':('a','b','p'),'blocks':{
        'E':block((),('branch','p','T','F')),
        'T':block((('x','add',('a','b')),),('jump','J')),
        'F':block((('y','add',('a','b')),),('jump','J')),
        'J':block((),('return',('z',)),(('z',(('T','x'),('F','y'))),))}}
    safe=scoped_value_numbering(sibling);require(not safe['replacements'],'siblings not carriers')
    bad=deepcopy(sibling);bad['blocks']['F']['code']=(('y','copy',('x',)),)
    try:execute(bad,{'a':2,'b':3,'p':0})
    except ValueError as error:sibling_failure=str(error)
    else:raise ValueError('sibling fault survived')
    # Removing only the carrier saves, while retaining the old computations, is wrong.
    missing=deepcopy(pre['program']);temp=pre['carrier']
    missing['blocks']['T']['code']=source['blocks']['T']['code']
    try:execute(missing,{'a':2,'b':3,'p':1,'q':0})
    except ValueError as error:carrier_failure=str(error)
    else:raise ValueError('carrier fault survived')
    wrong_place=deepcopy(pre['program']);bridge=pre['insertions'][0]['block']
    wrong_place['blocks']['F']['code']=wrong_place['blocks'][bridge]['code']
    wrong_place['blocks']['F']['term']=source['blocks']['F']['term'];del wrong_place['blocks'][bridge]
    wrong=execute(wrong_place,{'a':2,'b':3,'p':0,'q':0},selected=pre['expression'])
    require(wrong['value']==(-1,) and wrong['selected_count']==1,'critical bypass extra computation')
    # Collisions with generated-name prefixes must preserve every original name.
    collision=deepcopy(source);collision['params']+=('__pre_value_0',)
    collision['blocks']['__pre_edge_0']=collision['blocks'].pop('X')
    collision['blocks']['F']['term']=('branch','q','J','__pre_edge_0')
    renamed=edge_local_pre(collision,'J')
    require(renamed['carrier']=='__pre_value_1' and renamed['insertions'][0]['block']=='__pre_edge_1','fresh-name collisions')
    inputs={'a':2,'b':3,'p':1,'q':0,'__pre_value_0':99}
    before=execute(collision,inputs);after=execute(renamed['program'],inputs)
    require(before['trace']==macro_trace(after,collision),'prefix name still observable')
    # A structural unconditional back edge diverges; compare finite macro prefixes.
    divergent=edge_loop();divergent['blocks']['B']['term']=('branch',1,'J','X')
    transformed=edge_local_pre(divergent,'J');inputs={'a':2,'b':3,'n':1}
    before=execute(divergent,inputs,120,transformed['expression'])
    after=execute(transformed['program'],inputs,121,transformed['expression'])
    require(before['status']==after['status']=='budget_exhausted' and before['trace']==macro_trace(after,divergent),'divergent macro prefixes')
    rng=random.Random(21021);ssa_cases=0;pre_cases=0
    for _ in range(100):
        program=random_ssa(rng,rng.randrange(1,5));transformed=scoped_value_numbering(program)
        for _ in range(10):
            inp={p:rng.randrange(-4,5) if not p.startswith('p') else rng.randrange(2)for p in program['params']}
            before=execute(program,inp);after=execute(transformed['program'],inp)
            require(before['value']==after['value'] and before['trace']==after['trace'],'generated SSA equivalence');ssa_cases+=1
    for _ in range(400):
        inp={'a':rng.randrange(-20,21),'b':rng.randrange(-20,21),'p':rng.randrange(2),'q':rng.randrange(2)}
        before=execute(source,inp,selected=pre['expression']);after=execute(pre['program'],inp,selected=pre['expression'])
        require(before['value']==after['value'] and before['trace']==macro_trace(after,source) and after['selected_count']<=before['selected_count'],'PRE family');pre_cases+=1
    return {'status':'PASS','value_numbering':numbered,'value_runs':samples,'pre':pre,'pre_runs':cases,
            'value_loop_transformation':opt,'value_loops':loops,'pre_loops':pre_loops,
            'counterexamples':{'sibling_reuse':sibling_failure,'missing_carrier':carrier_failure,'critical_bypass':wrong},
            'regressions':{'generated_ssa_runs':ssa_cases,'pre_branch_runs':pre_cases,'value_loop_inputs':13,'pre_loop_inputs':loop_checks}}


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