#!/usr/bin/env python3
"""Executable two-level IMP examples, not a production security boundary.

All names initialized, unbounded integers, no heap/I/O/exceptions/concurrency.
One final sink. A fuel limit reports UNKNOWN, never divergence or a proof.
Run normally or with -O; every validation/test uses explicit checks.
"""
import json
from itertools import product

L, H, P = 0, 1, 2
LABEL = ('L', 'H', 'P')
OPS = {'+', '-', '==', '<', '<='}

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

def lit(n): return ('lit', n)
def var(x): return ('var', x)
def op(f, a, b): return (f, a, b)
def release(site, e): return ('release', site, e)
def assign(x, e): return ('set', x, e)
def seq(*cs): return ('seq', cs)
def iff(e, a, b=('skip',)): return ('if', e, a, b)
def loop(e, c): return ('while', e, c)
SKIP = ('skip',)

def expression_info(e, names, policy=None, collect=True):
    """Postorder validation returns (free-name mask, release-source mask).
    Expressions are trees. No nested release; exact approved site/expression.
    """
    todo, values = [(e, False)], []
    while todo:
        node, done = todo.pop()
        if not isinstance(node, tuple) or not node:
            raise ValueError('expression must be a nonempty tuple')
        k = node[0]
        if k == 'lit':
            if len(node) != 2 or type(node[1]) is not int: raise ValueError('integer literal')
            values.append((0, 0, False))
        elif k == 'var':
            if len(node) != 2 or node[1] not in names: raise ValueError('unknown name')
            values.append(((1 << names[node[1]]) if collect else 0, 0, False))
        elif k in OPS:
            if len(node) != 3: raise ValueError('binary expression')
            if not done: todo.extend([(node, True), (node[2], False), (node[1], False)])
            else:
                b, a = values.pop(), values.pop()
                values.append((a[0] | b[0], a[1] | b[1], a[2] or b[2]))
        elif k == 'release':
            if len(node) != 3 or policy is None or node[1] not in policy or policy[node[1]] != node[2]:
                raise ValueError('unapproved release site/expression')
            if not done: todo.extend([(node, True), (node[2], False)])
            else:
                a = values.pop()
                if a[2]: raise ValueError('nested release')
                values.append((a[0], a[0], True))
        else: raise ValueError('unknown expression')
    if len(values) != 1: raise ValueError('expression stack')
    return values[0][:2]

def validate(c, gamma, policy=None):
    if not gamma or any(v not in (L, H) for v in gamma.values()): raise ValueError('L/H initial map')
    names = {x:i for i,x in enumerate(gamma)}
    if policy is not None:
        for site, e in policy.items():
            if not isinstance(site, str): raise ValueError('release site must be text')
            expression_info(e, names, collect=False)  # every approved expression is pure and initialized
    todo = [c]
    while todo:
        n = todo.pop()
        if not isinstance(n, tuple) or not n: raise ValueError('command tuple')
        k = n[0]
        if k == 'skip' and len(n) == 1: pass
        elif k == 'set' and len(n) == 3:
            if n[1] not in names: raise ValueError('unknown assignment name')
            expression_info(n[2], names, policy, collect=False)
        elif k == 'seq' and len(n) == 2 and isinstance(n[1], tuple): todo.extend(n[1])
        elif k == 'if' and len(n) == 4:
            expression_info(n[1], names, policy, collect=False); todo.extend(n[2:])
        elif k == 'while' and len(n) == 3:
            expression_info(n[1], names, policy, collect=False); todo.append(n[2])
        else: raise ValueError('malformed command')
    return names

def eval_expr(e, state, labels, taint=False):
    """One explicit-stack postorder visit; returns value, label, visits.
    Release is an identity on values, resets its result label to L. Only raw
    execution with an explicitly supplied policy accepts its syntax.
    """
    todo, vals, visits = [(e, False)], [], 0
    while todo:
        n, done = todo.pop(); k = n[0]
        if not done: visits += 1
        if k == 'lit': vals.append((n[1], 0))
        elif k == 'var': vals.append((state[n[1]], labels[n[1]]))
        elif k == 'release':
            if not done: todo.extend([(n, True), (n[2], False)])
            else:
                v, _ = vals.pop(); vals.append((v, L))
        elif not done: todo.extend([(n, True), (n[2], False), (n[1], False)])
        else:
            (b, lb), (a, la) = vals.pop(), vals.pop()
            if k == '+': v = a+b
            elif k == '-': v = a-b
            elif k == '==': v = int(a == b)
            elif k == '<': v = int(a < b)
            else: v = int(a <= b)
            vals.append((v, (la | lb) if taint else max(la, lb)))
    v, label = vals.pop()
    return v, label, visits

def static_check(c, gamma, policy=None):
    """Syntax-directed pc typing; optional U/D anti-laundering effects.
    name sets are integer bitsets; seq summarized left to right.
    """
    names = validate(c, gamma, policy)
    events = []
    def expr(e):
        # Expression label uses fixed Gamma, except approved release result L.
        vals, todo = [], [(e, False)]
        while todo:
            n, done = todo.pop(); k = n[0]
            if k == 'lit': vals.append((L, 0, 0))
            elif k == 'var': vals.append((gamma[n[1]], 0, 1 << names[n[1]]))
            elif not done:
                todo.append((n, True))
                if k == 'release': todo.append((n[2], False))
                else: todo.extend([(n[2], False), (n[1], False)])
            elif k == 'release':
                a = vals.pop(); vals.append((L, a[2], a[2]))
            else:
                b, a = vals.pop(), vals.pop(); vals.append((max(a[0], b[0]), a[1]|b[1], a[2]|b[2]))
        label, d, _ = vals.pop(); return label, d
    # Explicit postorder command traversal; pc is restored by stored frames.
    todo, vals = [(c, L, False)], []
    while todo:
        n, pc, done = todo.pop(); k = n[0]
        if k == 'skip': vals.append((True, 0, 0, None))
        elif k == 'set':
            lab, d = expr(n[2]); ok = max(pc, lab) <= gamma[n[1]]
            reason = None if ok else 'assignment '+n[1]+': pc/join exceeds fixed label'
            vals.append((ok, 1 << names[n[1]], d, reason))
            events.append({'rule':'set', 'name':n[1], 'pc':LABEL[pc], 'rhs':LABEL[lab], 'accepted':ok})
        elif not done:
            if k == 'seq':
                todo.append((n, pc, True)); todo.extend((q, pc, False) for q in reversed(n[1]))
            else:
                lab, _ = expr(n[1]); npc = max(pc, lab)
                todo.append((n, pc, True)); todo.extend((q, npc, False) for q in reversed(n[2:]))
        elif k == 'seq':
            cnt = len(n[1]); rs = vals[-cnt:] if cnt else []
            if cnt: del vals[-cnt:]
            ok, u, d, reason = True, 0, 0, None
            for child_ok, cu, cd, why in rs:
                bad = policy is not None and bool(u & cd)
                if reason is None:
                    reason = why if not child_ok else ('sequence modifies a later release source' if bad else None)
                ok = ok and child_ok and not bad; u |= cu; d |= cd
            vals.append((ok, u, d, reason))
        elif k == 'if':
            b, a = vals.pop(), vals.pop(); _, d = expr(n[1])
            vals.append((a[0] and b[0], a[1]|b[1], d|a[2]|b[2], a[3] or b[3]))
        else:
            a = vals.pop(); _, d = expr(n[1]); d |= a[2]
            bad = policy is not None and bool(a[1] & d)
            vals.append((a[0] and not bad, a[1], d, a[3] or ('loop modifies a repeated release source' if bad else None)))
    ok, u, d, reason = vals.pop()
    return {'accepted':ok, 'U':[x for x in names if u >> names[x] & 1], 'D':[x for x in names if d >> names[x] & 1], 'reason':reason, 'assignments':events}

def run(c, initial, gamma, sink, mode='raw', fuel=10000, policy=None, sources=None, allowed_sources=None, trace=False):
    """Single final output. Internal state/trace returned for instructor inspection;
    never publish these as the program's low observation. Modes naive and taint
    deliberately do not enforce implicit-flow noninterference.
    """
    if mode not in ('raw','taint','naive','nsu','pu'): raise ValueError('mode')
    if policy is not None and mode != 'raw': raise ValueError('release only in raw policy lane')
    names = validate(c, gamma, policy)
    if set(initial) != set(gamma) or any(type(v) is not int for v in initial.values()): raise ValueError('initialized integer store')
    if sink not in gamma or gamma[sink] != L: raise ValueError('fixed L sink')
    if type(fuel) is not int or fuel < 0: raise ValueError('nonnegative fuel')
    state = dict(initial)
    if mode == 'taint':
        labels = dict(sources) if sources is not None else {x:(1 << names[x] if gamma[x] == H else 0) for x in names}
        if set(labels) != set(names) or any(type(v) is not int or v < 0 for v in labels.values()): raise ValueError('source bit masks')
    else: labels = dict(gamma)
    todo, events, steps, visits, dispatches = [(c, L)], [], 0, 0, 0
    def result(status, reason=None):
        tag = labels[sink]
        return {'status':status, 'value':state[sink] if status == 'OK' else None, 'sink_label':None if mode == 'raw' else (tag if mode == 'taint' else LABEL[tag]), 'steps':steps, 'dispatches':dispatches, 'expression_visits':visits, 'reason':reason, 'state':state, 'labels':labels, 'events':events}
    while todo:
        n, pc = todo.pop(); k = n[0]; dispatches += 1
        if k == 'seq': todo.extend((q, pc) for q in reversed(n[1])); continue
        if steps == fuel: return result('UNKNOWN', 'semantic step budget exhausted')
        steps += 1
        if k == 'skip':
            if trace: events.append({'step':steps,'op':'skip','pc':LABEL[pc]})
            continue
        e = n[2] if k == 'set' else n[1]
        value, tag, count = eval_expr(e, state, labels, mode == 'taint'); visits += count
        if k == 'set':
            x, old = n[1], labels[n[1]]
            if mode == 'nsu' and pc > old: return result('REJECT', 'NSU sensitive upgrade at '+x)
            if mode == 'pu' and pc == H and old != H: new = P
            elif mode in ('naive','nsu','pu'): new = max(pc, tag)
            else: new = tag
            state[x], labels[x] = value, new
            if trace: events.append({'step':steps,'op':'set','name':x,'value':value,'pc':LABEL[pc],'old':old,'new':new})
        else:
            if mode in ('naive','nsu','pu') and tag == P: return result('REJECT', 'partially leaked branch guard')
            npc = max(pc, tag) if mode in ('naive','nsu','pu') else L
            if trace: events.append({'step':steps,'op':k,'taken':bool(value),'pc':LABEL[pc],'guard_label':tag})
            if k == 'if': todo.append((n[2] if value else n[3], npc))
            elif value: todo.extend([(n, pc), (n[2], npc)])
    if mode in ('naive','nsu','pu') and labels[sink] != L: return result('REJECT', 'final sink is not L')
    if mode == 'taint' and allowed_sources is not None and labels[sink] & ~allowed_sources: return result('REJECT', 'forbidden source at sink')
    return result('OK')

def sme_low(c, initial, gamma, sink, default=0, fuel=10000, trace=False):
    """The low-public interface: no evaluation or waiting for the high copy."""
    if type(default) is not int: raise ValueError('integer default')
    if set(initial) != set(gamma) or any(type(v) is not int for v in initial.values()): raise ValueError('initialized integer store')
    projected = {x:(initial[x] if gamma[x] == L else default) for x in gamma}
    return run(c, projected, gamma, sink, fuel=fuel, trace=trace)

def summarize(r):
    return {k:r[k] for k in ('status','value','sink_label','steps','dispatches','expression_visits','reason')}

def examples():
    h,y,z = var('h'),var('y'),var('z')
    gamma = {'h':H, 'y':L, 'z':L}
    programs = {
        'explicit':assign('z',h),
        'implicit':iff(h,assign('z',lit(1)),assign('z',lit(0))),
        'same_branches':iff(h,assign('z',lit(0)),assign('z',lit(0))),
        'launder':seq(assign('y',lit(1)),assign('z',lit(1)),iff(h,assign('y',lit(0))),iff(y,assign('z',lit(0)))),
        'public_overwrite':seq(iff(h,assign('z',lit(1))),assign('z',lit(0))),
        'cancellation':assign('z',op('-',h,h)),
    }
    out = {}
    for name,c in programs.items():
        records = {'static':static_check(c,gamma)['accepted'], 'runs':{}}
        for b in (0,1):
            initial={'h':b,'y':0,'z':0}
            records['runs'][str(b)]={m:summarize(run(c,initial,gamma,'z',m)) for m in ('raw','taint','naive','nsu','pu')}
            records['runs'][str(b)]['sme_low']=summarize(sme_low(c,initial,gamma,'z'))
        out[name]=records
    check(out['launder']['runs']['0']['naive']['value']==0 and out['launder']['runs']['1']['naive']['value']==1,'naive implicit leak')
    check(out['public_overwrite']['runs']['1']['nsu']['status']=='REJECT' and out['public_overwrite']['runs']['1']['pu']['value']==0,'PU overwrite')
    check(out['cancellation']['runs']['1']['taint']['sink_label'] != 0,'syntactic data dependence')
    check(not out['same_branches']['static'],'conservative static check')
    return programs,gamma,out

def release_examples():
    gamma={'a':H,'b':H,'out':L,'n':L}
    total=op('+',var('a'),var('b')); policy={'sum':total}
    emit=assign('out',release('sum',total))
    good=emit; bad=seq(assign('b',var('a')),emit); after=seq(emit,assign('b',var('a')))
    repeated=loop(var('n'),seq(emit,assign('b',var('a')),assign('n',op('-',var('n'),lit(1)))))
    okloop=loop(var('n'),seq(emit,assign('n',op('-',var('n'),lit(1)))))
    output={}
    for name,c in [('good',good),('laundered',bad),('mutate_after',after),('repeat_bad',repeated),('repeat_good',okloop)]:
        typed=static_check(c,gamma,policy); pairs=[]
        for a,b in [(2,5),(3,4)]:
            pairs.append(summarize(run(c,{'a':a,'b':b,'out':0,'n':2},gamma,'out',policy=policy)))
        output[name]={'typing':typed,'runs':pairs}
    check([r['value'] for r in output['good']['runs']]==[7,7],'sum release')
    check([r['value'] for r in output['laundered']['runs']]==[4,6] and not output['laundered']['typing']['accepted'],'anti-laundering')
    check(output['mutate_after']['typing']['accepted'] and output['repeat_good']['typing']['accepted'] and not output['repeat_bad']['typing']['accepted'],'directional and loop effects')
    return output

def termination_examples():
    g={'h':H,'z':L}; initial={'h':1,'z':0}
    default_diverges=seq(loop(op('==',var('h'),lit(0)),SKIP),assign('z',lit(0)))
    secret_delay=seq(loop(var('h'),assign('h',op('-',var('h'),lit(1)))),assign('z',lit(0)))
    out={'default_diverges':{'static':static_check(default_diverges,g)['accepted'],'real':summarize(run(default_diverges,initial,g,'z',fuel=20)),'low':summarize(sme_low(default_diverges,initial,g,'z',fuel=20))}, 'delay':[]}
    for h in (0,1,3):
        s={'h':h,'z':0}; a=sme_low(secret_delay,s,g,'z');b=run(secret_delay,s,g,'z')
        out['delay'].append({'h':h,'low_steps':a['steps'],'high_steps':b['steps'],'public_value':a['value'],'both_work':a['steps']+b['steps']})
    check(out['default_diverges']['real']['status']=='OK' and out['default_diverges']['low']['status']=='UNKNOWN','fuel is not divergence')
    check([a['low_steps'] for a in out['delay']]==[2,2,2],'public interface independent of high work')
    return out

def lineage_example():
    g={'a':H,'b':H,'t':L,'u':L,'out':L};initial={'a':3,'b':4,'t':0,'u':0,'out':0};sources={'a':1,'b':2,'t':0,'u':0,'out':0}
    c=seq(assign('t',op('+',var('a'),var('b'))),assign('u',op('-',var('t'),var('a'))),assign('t',lit(0)),assign('out',var('u')))
    narrow=seq(assign('t',op('+',var('a'),var('b'))),assign('u',var('b')),assign('t',lit(0)),assign('out',var('u')))
    a=run(c,initial,g,'out','taint',sources=sources,trace=True);b=run(narrow,initial,g,'out','taint',sources=sources,allowed_sources=2)
    rejected=run(c,initial,g,'out','taint',sources=sources,allowed_sources=2)
    check(a['value']==4 and a['labels']['out']==3 and b['value']==4 and b['labels']['out']==2 and rejected['status']=='REJECT','two-source replacement')
    return {'original':a,'B_only_sink':summarize(rejected),'direct_b_migration':summarize(b)}

def finite_pair_checks():
    # A bounded family test, not a theorem for all integer programs.
    g={'h':H,'a':H,'y':L,'z':L}
    terms=[lit(0),lit(1),var('h'),var('a'),var('y'),op('-',var('h'),var('h'))]
    atoms=[assign(x,e) for x in g for e in terms]
    programs=atoms+[iff(var('h'),a,b) for a in atoms for b in atoms]+[seq(a,b) for a in atoms for b in atoms]
    initial=[{'h':h,'a':a,'y':y,'z':0} for h,a,y in product((-1,0,1),repeat=3)]
    typed, pair_count, accepted_pairs=0,0,{'nsu':0,'pu':0}
    for c in programs:
        accepted=static_check(c,g)['accepted']; typed+=accepted
        by_low={}
        for s in initial:
            rr=run(c,s,g,'z'); group=by_low.setdefault((s['y'],s['z']),[])
            group.append({m:run(c,s,g,'z',m) for m in ('nsu','pu')} | {'raw':rr,'low':sme_low(c,s,g,'z')})
        for rs in by_low.values():
            for i in range(len(rs)):
                for j in range(i):
                    a,b=rs[i],rs[j]; pair_count+=1
                    if accepted: check(a['raw']['value']==b['raw']['value'],'static bounded pair')
                    check(a['low']['value']==b['low']['value'] and a['low']['steps']==b['low']['steps'],'SME bounded pair')
                    for m in ('nsu','pu'):
                        if a[m]['status']==b[m]['status']=='OK':
                            accepted_pairs[m]+=1; check(a[m]['value']==b[m]['value'],'monitor bounded pair')
    # Semantic loops with pc restore, unknown, and empty sequence regressions.
    check(run(seq(),{'h':0,'a':0,'y':0,'z':0},g,'z',fuel=0)['status']=='OK','empty command constant case')
    check(run(SKIP,{'h':0,'a':0,'y':0,'z':0},g,'z',fuel=0)['status']=='UNKNOWN','zero fuel')
    return {'programs':len(programs),'inputs_per_program':len(initial),'static_accepted':typed,'pairs':pair_count,'monitor_accepted_pairs':accepted_pairs}

def main():
    _,_,ex=examples()
    output={'model':'two-level pure initialized IMP; one final low sink; fuel exhaustion UNKNOWN','examples':ex,'lineage':lineage_example(),'release':release_examples(),'termination':termination_examples(),'finite_regression':finite_pair_checks()}
    print(json.dumps(output,ensure_ascii=False,indent=2))

if __name__ == '__main__': main()
