#!/usr/bin/env python3
"""Pure teaching models, not a Scheme implementation or a general compiler.
Integer arithmetic is mathematical; Python's own integer bit costs are not hidden.
Checks use explicit exceptions, so normal and -O executions have the same meaning.
"""
import json
from collections import namedtuple

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

# de Bruijn syntax: ('v',i), ('l',body), ('a',function,argument).
def scope(t, depth):
    if t[0]=='v': return type(t[1]) is int and 0<=t[1]<depth
    if t[0]=='l': return scope(t[1],depth+1)
    return scope(t[1],depth) and scope(t[2],depth)

def shift(t, delta, cutoff=0):
    tag=t[0]
    if tag=='v':
        i=t[1]+(delta if t[1]>=cutoff else 0)
        if i<0: raise ValueError('negative index after shift')
        return ('v',i)
    if tag=='l': return ('l',shift(t[1],delta,cutoff+1))
    return ('a',shift(t[1],delta,cutoff),shift(t[2],delta,cutoff))

def subst_index(t, index, replacement):
    # Carry depth; shift the original replacement only at actual replacement sites.
    def walk(t, depth):
        if t[0]=='v': return shift(replacement,depth) if t[1]==index+depth else t
        if t[0]=='l': return ('l',walk(t[1],depth+1))
        return ('a',walk(t[1],depth),walk(t[2],depth))
    return walk(t,0)

def beta(body, argument, external_count):
    if not scope(body,external_count+1) or not scope(argument,external_count):
        raise ValueError('ill-scoped beta operands')
    result=shift(subst_index(body,0,shift(argument,1)),-1)
    check(scope(result,external_count),'beta lost scope')
    return result

# Named syntax: ('v',name), ('l',name,body), ('a',function,argument).
def to_db(t, context=()):
    if t[0]=='v':
        try:return ('v',context.index(t[1]))
        except ValueError:raise ValueError('unlisted free variable: '+t[1])
    if t[0]=='l':return ('l',to_db(t[2],(t[1],)+context))
    return ('a',to_db(t[1],context),to_db(t[2],context))

# Locally nameless: b for bound index, f for free atom, l/a as above.
def to_ln(t, bound=()):
    if t[0]=='v':return ('b',bound.index(t[1])) if t[1] in bound else ('f',t[1])
    if t[0]=='l':return ('l',to_ln(t[2],(t[1],)+bound))
    return ('a',to_ln(t[1],bound),to_ln(t[2],bound))

def lc(t, depth=0):
    if t[0]=='f':return True
    if t[0]=='b':return type(t[1]) is int and 0<=t[1]<depth
    if t[0]=='l':return lc(t[1],depth+1)
    return lc(t[1],depth) and lc(t[2],depth)

def open_term(t, u, depth=0):
    if t[0]=='b':return u if t[1]==depth else t
    if t[0]=='f':return t
    if t[0]=='l':return ('l',open_term(t[1],u,depth+1))
    return ('a',open_term(t[1],u,depth),open_term(t[2],u,depth))

def close_term(t, atom, depth=0):
    if t[0]=='f':return ('b',depth) if t[1]==atom else t
    if t[0]=='b':return t
    if t[0]=='l':return ('l',close_term(t[1],atom,depth+1))
    return ('a',close_term(t[1],atom,depth),close_term(t[2],atom,depth))

def ln_beta(body,u):
    if not lc(body,1) or not lc(u):raise ValueError('opening requires a body and a locally closed argument')
    result=open_term(body,u)
    check(lc(result),'opening lost local closure')
    return result

# One closed function family. Compose(f,g) means x -> f(g(x)).
def source_function(f):
    if f[0]=='Add':return lambda x:x+f[1]
    if f[0]=='Mul':return lambda x:x*f[1]
    a,b=source_function(f[1]),source_function(f[2])
    return lambda x:a(b(x))

def apply_tag(f,x):
    if f[0]=='Add':return x+f[1]
    if f[0]=='Mul':return x*f[1]
    if f[0]=='Compose':return apply_tag(f[1],apply_tag(f[2],x))
    raise ValueError('unknown function tag')

def twice(f,x):return apply_tag(f,apply_tag(f,x))

# Every user-level CPS transfer is delayed, including continuation invocation.
Done=namedtuple('Done','value')
More=namedtuple('More','thunk')

def sum_cps_bounced(n,k):
    if n<0:raise ValueError('nonnegative input required')
    if n==0:return More(lambda:k(0))
    return More(lambda:sum_cps_bounced(n-1,lambda v:More(lambda:k(n+v))))

def run_bounces(job):
    calls=0
    while isinstance(job,More):
        job=job.thunk();calls+=1
    if not isinstance(job,Done):raise ValueError('job must return Done or More')
    return job.value,calls

# Equivalent first-order jobs: Enter(n,K), Resume(v,K), Halt(v).
# Frames are linked immutable ('Plus',n,rest), not copied lists.
def sum_jobs(n):
    if n<0:raise ValueError('nonnegative input required')
    job=('Enter',n,None);trace=[];steps=0;peak=0;depth=0
    while job[0]!='Halt':
        if n<=4:trace.append([job[0],job[1],depth])
        tag,x,k=job
        if tag=='Enter':
            if x==0:job=('Resume',0,k)
            else:
                job=('Enter',x-1,('Plus',x,k));depth+=1;peak=max(peak,depth)
        elif k is None:job=('Halt',x)
        else:job=('Resume',x+k[1],k[2]);depth-=1
        steps+=1
    return {'value':job[1],'job_steps':steps,'peak_frames':peak,'trace':trace}

# Small CEK extension. The top frame is the head of a persistent linked stack.
Closure=namedtuple('Closure','name body env')
Full=namedtuple('Full','stack')
Part=namedtuple('Part','frames')

def num(n):return ('num',n)
def var(x):return ('var',x)
def lam(x,e):return ('lam',x,e)
def app(f,x):return ('app',f,x)
def add(x,y):return ('add',x,y)
def cc(f):return ('cc',f)
def reset(e):return ('reset',e)
def shift_control(k,e):return ('shift',k,e)

def frames(k):
    out=[]
    while k is not None:out.append(k[0][0]);k=k[1]
    return out

def evaluate(term, env=None, initial_stack=None, max_steps=100000, record_events=True):
    control=term;env={} if env is None else env;stack=initial_stack
    value_mode=False;steps=0;events=[]
    def apply(f,v,k):
        if isinstance(f,Closure):return f.body,{**f.env,f.name:v},k,False
        if isinstance(f,Full):
            if record_events: events.append({'event':'invoke-full','discard':frames(k),'restore':frames(f.stack)})
            return v,{},f.stack,True
        if isinstance(f,Part):
            if record_events: events.append({'event':'invoke-delimited','saved':len(f.frames),'caller':frames(k)})
            new=(('mark',),k)
            for frame in reversed(f.frames):new=(frame,new)
            return v,{},new,True
        raise ValueError('application of a non-function')
    while True:
        if steps>=max_steps:raise RuntimeError('step bound reached; not a proof of divergence')
        steps+=1
        if not value_mode:
            tag=control[0]
            if tag=='num':control=control[1];value_mode=True
            elif tag=='var':control=env[control[1]];value_mode=True
            elif tag=='lam':control=Closure(control[1],control[2],env);value_mode=True
            elif tag in ('add','app'):
                stack=((('add-left' if tag=='add' else 'argument'),control[2],env),stack);control=control[1]
            elif tag=='cc':stack=(('capture',),stack);control=control[1]
            elif tag=='reset':stack=(('mark',),stack);control=control[1]
            elif tag=='shift':
                part=[];rest=stack
                while rest is not None and rest[0][0]!='mark':part.append(rest[0]);rest=rest[1]
                if rest is None:raise ValueError('shift without an enclosing reset')
                if record_events: events.append({'event':'capture-delimited','saved':[f[0] for f in part]})
                env={**env,control[1]:Part(tuple(part))};control=control[2];stack=rest
            else:raise ValueError('unknown expression')
        else:
            if stack is None:return control,{'steps':steps,'events':events}
            f,stack=stack;tag=f[0]
            if tag=='add-left':stack=(('add-right',control),stack);control=f[1];env=f[2];value_mode=False
            elif tag=='add-right':
                if type(control) is not int or type(f[1]) is not int:raise ValueError('non-integer addition')
                control=f[1]+control
            elif tag=='argument':stack=(('function',control),stack);control=f[1];env=f[2];value_mode=False
            elif tag=='function':control,env,stack,value_mode=apply(f[1],control,stack)
            elif tag=='capture':
                if record_events: events.append({'event':'capture-full','saved':frames(stack)})
                control,env,stack,value_mode=apply(control,Full(stack),stack)
            elif tag=='mark':pass
            else:raise ValueError('unknown frame')

def main():
    named=('a',('l','x',('l','y',('v','x'))),('v','z'))
    encoded=to_db(named,('z',));db=beta(encoded[1][1],encoded[2],1)
    wrong=shift(subst_index(encoded[1][1],0,encoded[2]),-1)
    check(db==('l',('v',1)) and wrong==('l',('v',0)),'capture witness')
    ln=to_ln(named);ln_result=ln_beta(ln[1][1],ln[2]);check(ln_result==('l',('f','z')),'LN beta')
    body=('a',('b',0),('f','z'));opened=open_term(body,('f','a'))
    check(close_term(opened,'a')==body,'fresh inverse')
    collision=close_term(open_term(body,('f','z')),'z')
    check(collision!=body and not lc(('l',('b',1))),'freshness / closure counterexamples')
    nested_argument=('l','w',('v','z'))
    nested_db=beta(encoded[1][1],to_db(nested_argument,('z',)),1)
    nested_ln=ln_beta(ln[1][1],to_ln(nested_argument))
    check(nested_db==('l',('l',('v',2))) and nested_ln==('l',('l',('f','z'))),'argument binder migration')
    try:ln_beta(ln[1][1],('b',0))
    except ValueError:invalid_ln_rejected=True
    else:raise RuntimeError('dangling LN argument accepted')
    f=('Compose',('Mul',3),('Add',2));check(twice(f,1)==33,'function family')
    reversed_twice=twice(('Compose',('Add',2),('Mul',3)),1)
    check(reversed_twice==17,'composition order migration')
    jobs=sum_jobs(3);check(jobs['value']==6 and jobs['job_steps']==8 and jobs['peak_frames']==3,'jobs trace')
    value,bounces=run_bounces(sum_cps_bounced(10000,Done));check((value,bounces)==(50005000,20001),'large trampoline')
    examples={
      'full_6':add(num(1),cc(lam('k',add(num(100),app(var('k'),num(5)))))),
      'delimited_106':reset(add(num(1),shift_control('k',add(num(100),app(var('k'),num(5)))))),
      'multi_32':reset(add(num(1),shift_control('k',add(app(var('k'),num(10)),app(var('k'),num(20)))))),
      'nested_110':reset(add(num(100),reset(add(num(1),shift_control('k',num(10)))))),
      'removed_inner_10':reset(add(num(100),add(num(1),shift_control('k',num(10))))),
      'normal_full_return_106':add(num(1),cc(lam('k',num(105)))),
      'reinstalled_110':reset(add(shift_control('k',add(num(100),app(var('k'),num(1)))),shift_control('j',num(10)))),
      'nested_exit_6':add(num(1),cc(lam('exit',add(num(10),cc(lam('inner',app(var('exit'),num(5)))))))),
      'nested_inner_16':add(num(1),cc(lam('exit',add(num(10),cc(lam('inner',app(var('inner'),num(5))))))))}
    outcomes={}
    for key,t in examples.items():
        v,trace=evaluate(t);check(v==int(key.rsplit('_',1)[1]),'control '+key);outcomes[key]={'value':v,**trace}
    saved,unused=evaluate(cc(lam('k',var('k'))));check(isinstance(saved,Full),'first-class return')
    repeated=[evaluate(app(var('saved'),num(i)),{'saved':saved})[0] for i in [7,9]]
    check(repeated==[7,9],'immutable full continuation reuse')
    try:evaluate(shift_control('k',num(1)))
    except ValueError as e:unhandled=str(e)
    else:raise RuntimeError('unhandled shift accepted')
    print(json.dumps({'result':'PASS','binding':{'encoded':encoded,'beta':db,'missing_initial_shift':wrong,'locally_nameless':ln_result,'fresh_inverse':True,'nonfresh_close':collision,'nested_argument_db':nested_db,'nested_argument_ln':nested_ln,'invalid_LN_argument_rejected':invalid_ln_rejected},'defunctionalization':{'function':f,'at_1':apply_tag(f,1),'twice_at_1':twice(f,1),'reversed_twice_at_1':reversed_twice},'trampoline':{'small':jobs,'n':10000,'sum':value,'bounce_calls':bounces},'control':outcomes,'reused_full':repeated,'unhandled_shift':unhandled},ensure_ascii=False,indent=2))

if __name__=='__main__':main()
