#!/usr/bin/env python3
"""Canonical integer loops: affine recurrence and linear-function test replacement.

Expression: int, variable name, or (binary_op, left, right).
Program: params; init/body tuples of (destination, expression); guard; returns.
Every operation is pure/total over mathematical integers. No machine wraparound.
"""
from copy import deepcopy
import json
import random

ARITH=('add','sub','mul')
COMPARE=('lt','le','gt','ge','eq','ne')
REVERSE={'lt':'gt','le':'ge','gt':'lt','ge':'le','eq':'eq','ne':'ne'}


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


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


def read_names(e):
    if type(e)is int:return set()
    if name(e):return {e}
    require(type(e)is tuple and len(e)==3 and e[0]in ARITH+COMPARE,'expression shape')
    return read_names(e[1])|read_names(e[2])


def apply(op,a,b):
    if op=='add':return a+b
    if op=='sub':return a-b
    if op=='mul':return a*b
    if op=='lt':return int(a<b)
    if op=='le':return int(a<=b)
    if op=='gt':return int(a>b)
    if op=='ge':return int(a>=b)
    if op=='eq':return int(a==b)
    if op=='ne':return int(a!=b)
    raise ValueError('operator')


def evaluate(e,state,counts=None):
    if type(e)is int:return e
    if name(e):return state[e]
    a=evaluate(e[1],state,counts);b=evaluate(e[2],state,counts)
    if counts is not None:counts[e[0]]+=1
    return apply(e[0],a,b)


def validate(p):
    require(type(p)is dict and set(p)=={'params','init','guard','body','returns'},'program shape')
    require(type(p['params'])is tuple and all(name(x)for x in p['params']) and len(set(p['params']))==len(p['params']),'parameters')
    have=set(p['params']);all_names=set(have)
    for region in('init','body'):
        require(type(p[region])is tuple,'assignment sequence')
        if region=='body':
            require(type(p['guard'])is tuple and len(p['guard'])==3 and p['guard'][0]in COMPARE,'comparison guard')
            require(read_names(p['guard'])<=have,'uninitialized guard')
            require(type(p['returns'])is tuple and all(read_names(e)<=have for e in p['returns']),'return needs zero-trip definition')
        for instruction in p[region]:
            require(type(instruction)is tuple and len(instruction)==2 and name(instruction[0]),'assignment')
            x,e=instruction;require(read_names(e)<=have,'uninitialized expression');have.add(x);all_names.add(x)
    return all_names


def execute(p,inputs,iteration_budget=100):
    validate(p);require(type(inputs)is dict and set(inputs)==set(p['params']) and all(type(x)is int for x in inputs.values()),'integer inputs')
    require(type(iteration_budget)is int and iteration_budget>=0,'iteration budget')
    state=dict(inputs);counts={op:0 for op in ARITH+COMPARE};assignments=0
    for x,e in p['init']:state[x]=evaluate(e,state,counts);assignments+=1
    trace=[];iterations=0
    while True:
        trace.append(dict(state));condition=evaluate(p['guard'],state,counts)
        if not condition:
            value=tuple(evaluate(e,state,counts)for e in p['returns'])
            return {'status':'returned','value':value,'iterations':iterations,'headers':trace,'counts':counts,'assignments':assignments}
        if iterations==iteration_budget:
            return {'status':'budget_exhausted','iterations':iterations,'headers':trace,'counts':counts,'assignments':assignments}
        for x,e in p['body']:state[x]=evaluate(e,state,counts);assignments+=1
        iterations+=1


def fold(e):
    if type(e)is int or name(e):return e
    a,b=fold(e[1]),fold(e[2])
    return apply(e[0],a,b)if type(a)is int and type(b)is int else(e[0],a,b)


def substitute(e,old,new):
    if e==old:return new
    if type(e)is int or name(e):return e
    return (e[0],substitute(e[1],old,new),substitute(e[2],old,new))


def replace_expression(e,old,new):
    if e==old:return new,1
    if type(e)is int or name(e):return e,0
    a,na=replace_expression(e[1],old,new);b,nb=replace_expression(e[2],old,new)
    return (e[0],a,b),na+nb


def canonical(p,induction,expression):
    all_names=validate(p);require(name(induction)and induction not in p['params'],'local induction name')
    require(p['init']and p['init'][-1][0]==induction and sum(x==induction for x,_ in p['init'])==1,'induction initialized once at end of init')
    require(p['body']and p['body'][-1][0]==induction and sum(x==induction for x,_ in p['body'])==1,'one final induction update')
    update=p['body'][-1][1]
    require(type(update)is tuple and len(update)==3 and update[0]in('add','sub'),'affine update')
    if update[1]==induction:step=update[2]if update[0]=='add'else('sub',0,update[2])
    elif update[0]=='add'and update[2]==induction:step=update[1]
    else:raise ValueError('induction update shape')
    require(type(expression)is tuple and len(expression)==3 and expression[0]=='add','selected affine expression')
    product,offset=expression[1],expression[2]
    require(type(product)is tuple and len(product)==3 and product[0]=='mul'and type(product[1])is int and product[2]==induction,'literal slope times induction')
    slope=product[1];written={x for x,_ in p['body']};ready=set(p['params'])|{x for x,_ in p['init']}
    require(not(read_names(offset)|read_names(step))&written,'step and offset must be loop invariant')
    require((read_names(offset)|read_names(step))<=ready,'invariants initialized')
    return {'slope':slope,'offset':offset,'step':fold(step),'start':p['init'][-1][1],'names':all_names}


def strength_reduce(p,induction,expression):
    info=canonical(p,induction,expression);taken=set(info['names'])
    def fresh(prefix):
        i=0
        while prefix+str(i)in taken:i+=1
        x=prefix+str(i);taken.add(x);return x
    carrier=fresh('__iv_value_');delta=fresh('__iv_delta_');result=deepcopy(p);body=[];sites=[]
    for index,(x,e)in enumerate(p['body'][:-1]):
        rewritten,n=replace_expression(e,expression,carrier);body.append((x,rewritten))
        if n:sites.append({'instruction':index,'destination':x,'occurrences':n})
    require(sites,'no selected occurrence before update')
    result['init']+=((carrier,expression),(delta,fold(('mul',info['slope'],info['step']))))
    result['body']=tuple(body)+(p['body'][-1],(carrier,('add',carrier,delta)))
    validate(result)
    return {'program':result,'induction':induction,'expression':expression,'carrier':carrier,'delta':delta,'slope':info['slope'],'offset':info['offset'],'step':info['step'],'start':info['start'],'sites':sites}


def linear_test_replacement(p,induction,expression):
    # Re-establish the recurrence from the actual source; never trust a forged report.
    reduced=strength_reduce(p,induction,expression);slope=reduced['slope'];guard=p['guard'];written={x for x,_ in p['body']}
    def refuse(reason):return {'status':'refused','reason':reason,'program':None,'strength_reduction':reduced}
    if slope==0:return refuse('zero slope is not injective')
    if guard[1]!=induction:return refuse('guard must compare induction on the left')
    bound=guard[2]
    if read_names(bound)&written:return refuse('bound is not loop invariant')
    if any(induction in read_names(e)for _,e in reduced['program']['body'][:-2]):return refuse('induction still read in body')
    if any(induction in read_names(e)for e in p['returns']):return refuse('induction live at exit')
    result=deepcopy(reduced['program']);carrier=reduced['carrier'];taken=validate(result);n=0
    while '__iv_bound_'+str(n)in taken:n+=1
    limit='__iv_bound_'+str(n);new_bound=fold(('add',('mul',slope,bound),reduced['offset']))
    # The original induction initialization is last; no intervening source write
    # can change its RHS before substitution into the fresh carrier initializer.
    result['init']=p['init'][:-1]+((carrier,substitute(expression,induction,reduced['start'])),result['init'][-1],(limit,new_bound))
    result['body']=result['body'][:-2]+(result['body'][-1],)
    op=guard[0]if slope>0 else REVERSE[guard[0]];result['guard']=(op,carrier,limit)
    validate(result)
    return {'status':'accepted','program':result,'strength_reduction':reduced,'guard_before':guard,'guard_after':result['guard'],'bound_name':limit,'bound_expression':new_bound,'removed':induction}


def source_example(slope=5,relation='lt',step=3,live_exit=False):
    expression=('add',('mul',slope,'i'),'bias')
    p={'params':('start','limit','bias','seed','prior'),
       'init':(('sum','seed'),('j','prior'),('i','start')),
       'guard':(relation,'i','limit'),
       'body':(('j',expression),('sum',('add','sum','j')),('i',('add','i',step))),
       'returns':('sum','j','i')if live_exit else('sum','j')}
    return p,expression


def visible_headers(run,source,removed=None):
    names=validate(source)-({removed}if removed else set())
    return [{k:v for k,v in state.items()if k in names}for state in run['headers']]


def check_triplet(p,e,inputs,budget=100):
    reduced=strength_reduce(p,'i',e);lftr=linear_test_replacement(p,'i',e)
    a=execute(p,inputs,budget);b=execute(reduced['program'],inputs,budget)
    require(a['status']==b['status'] and a.get('value')==b.get('value'),'strength value/status')
    require(a['headers']==visible_headers(b,p),'strength source states')
    for source,target in zip(a['headers'],b['headers']):
        require(target[reduced['carrier']]==evaluate(e,source),'recurrence invariant')
    if lftr['status']=='accepted':
        c=execute(lftr['program'],inputs,budget)
        require(a['status']==c['status'] and a.get('value')==c.get('value'),'LFTR value/status')
        require(visible_headers(a,p,'i')==visible_headers(c,p,'i'),'LFTR retained states')
        require([evaluate(p['guard'],s)for s in a['headers']]==[evaluate(lftr['program']['guard'],s)for s in c['headers']],'guard truth sequence')
    else:c=None
    return {'source':a,'reduced':b,'eliminated':c,'strength':reduced,'test_replacement':lftr}


def self_test():
    examples=[]
    for slope,bias,start,limit in((5,-4,2,13),(-2,9,2,13),(5,-4,14,13)):
        p,e=source_example(slope);inputs={'start':start,'limit':limit,'bias':bias,'seed':0,'prior':99};row=check_triplet(p,e,inputs);row['inputs']=inputs;examples.append(row)
    require(examples[0]['source']['value']==(114,51),'main result')
    require(examples[1]['source']['value']==(-16,-13),'negative slope')
    require(examples[2]['source']['value']==(0,99),'zero-trip original j')
    require([x['counts']['mul']for x in(examples[0]['source'],examples[0]['reduced'],examples[0]['eliminated'])]==[4,1,2],'multiply accounting')
    p,e=source_example();wrong=deepcopy(examples[0]['strength']['program']);wrong['body']=(wrong['body'][-1],)+wrong['body'][:-1]
    phase=execute(wrong,examples[0]['inputs']);require(phase['value']==(174,66),'early recurrence update witness')
    wrong=deepcopy(examples[1]['test_replacement']['program']);wrong['guard']=('lt',)+wrong['guard'][1:]
    sign=execute(wrong,examples[1]['inputs']);require(sign['value']==(0,99),'negative test reversal witness')
    wrong=deepcopy(examples[0]['test_replacement']['program']);wrong['guard']=('ne',)+wrong['guard'][1:]
    inequality=execute(wrong,examples[0]['inputs'],12);require(inequality['status']=='budget_exhausted','strict inequality cannot become !=')
    p,e=source_example(0);zero=linear_test_replacement(p,'i',e);require(zero['status']=='refused','zero slope refusal')
    p,e=source_example(live_exit=True);live=linear_test_replacement(p,'i',e);require(live['reason']=='induction live at exit','exit liveness refusal')
    p,e=source_example();p['body']=(('j',e),('sum',('add','sum','i')),p['body'][-1]);body_live=linear_test_replacement(p,'i',e);require(body_live['reason']=='induction still read in body','body liveness refusal')
    p,e=source_example(relation='ne');unreachable=check_triplet(p,e,{'start':2,'limit':13,'bias':-4,'seed':0,'prior':99},30)
    require(unreachable['source']['status']=='budget_exhausted','unreachable != bound')
    p,e=source_example(step=0);stationary=check_triplet(p,e,{'start':2,'limit':13,'bias':-4,'seed':0,'prior':99},30)
    require(stationary['source']['status']=='budget_exhausted','zero step prefix')
    # Generated-name prefixes do not make existing variables invisible.
    p,e=source_example();p['params']+=('__iv_value_0','__iv_delta_0','__iv_bound_0')
    inp={'start':2,'limit':13,'bias':-4,'seed':0,'prior':99,'__iv_value_0':71,'__iv_delta_0':72,'__iv_bound_0':73}
    collision=check_triplet(p,e,inp);require(collision['strength']['carrier']=='__iv_value_1'and collision['test_replacement']['bound_name']=='__iv_bound_1','fresh names')
    rng=random.Random(22022);cases=0;returned=0;prefixes=0
    for _ in range(900):
        slope=rng.choice((-5,-2,-1,0,1,3,7));relation=rng.choice(COMPARE);step=rng.randrange(-4,5)
        p,e=source_example(slope,relation,step);inp={'start':rng.randrange(-8,9),'limit':rng.randrange(-8,9),'bias':rng.randrange(-5,6),'seed':rng.randrange(-5,6),'prior':rng.randrange(-5,6)}
        row=check_triplet(p,e,inp,25);cases+=1;returned+=row['source']['status']=='returned';prefixes+=row['source']['status']=='budget_exhausted'
    # A modular recurrence can be valid while its ordered comparison is invalid.
    i=0;source_modular=0
    while i<100:i=(i+1)%256;source_modular+=1
    h=0;wrong_modular=0
    while h<(4*100)%256:h=(h+4)%256;wrong_modular+=1
    require((source_modular,wrong_modular)==(100,36),'wraparound comparison witness')
    repeated=0.0
    for _ in range(6):repeated+=0.1
    require(repeated!=0.1*6,'floating recurrence witness')
    return {'status':'PASS','examples':examples,'counterexamples':{'early_update':phase,'missing_sign_reversal':sign,'changed_comparison_kind':inequality,'zero_slope_refusal':zero['reason'],'live_exit_refusal':live['reason'],'body_read_refusal':body_live['reason'],'modular_iterations':[source_modular,wrong_modular],'floating_values':[repeated,0.1*6]},'prefixes':{'unreachable_inequality_bound':unreachable,'zero_step':stationary},'regressions':{'generated_inputs':cases,'returned':returned,'budget_prefixes':prefixes}}


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