#!/usr/bin/env python3
"""Finite rectangular integer-array loops; no packages, writes only JSON to stdout.
Locations are (storage object, offset), so differently named views may alias.
IR is a finite For chain ending in emit; emit executes the fixed source body.
The certificate is recomputed from the actual IR, never from a claimed schedule.
Explicit require checks work with and without python -O. Not a cache benchmark.
"""
from itertools import product, permutations
from collections import Counter
import json


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


def integer(x):
    return type(x) is int


def seq(x):
    return type(x) in (tuple, list)


def aff(*coefficients):
    return tuple(coefficients)


def load(view, *indices):
    return ('load', view, tuple(indices))


def const(x):
    return ('const', x)


def add(a, b):
    return ('add', a, b)


def mul(a, b):
    return ('mul', a, b)


def references(expr):
    stack, out = [expr], []
    while stack:
        e = stack.pop()
        require(seq(e) and e, 'expression tuple')
        op = e[0]
        if op == 'const':
            require(len(e) == 2 and integer(e[1]), 'integer constant')
        elif op == 'load':
            require(len(e) == 3, 'load arity')
            out.append((e[1], e[2]))
        else:
            require(op in ('add', 'mul') and len(e) == 3, 'expression opcode')
            stack.extend((e[2], e[1]))
    return out


def check_kernel(p):
    require(type(p) is dict and set(p) == {'dims', 'objects', 'views', 'body'}, 'kernel fields')
    dims, objects, views, body = (p[k] for k in ('dims', 'objects', 'views', 'body'))
    require(seq(dims) and len(dims) >= 1 and all(integer(n) and n >= 0 for n in dims), 'dimensions')
    require(type(objects) is dict and type(views) is dict, 'storage tables')
    for name, n in objects.items():
        require(type(name) is str and name and integer(n) and n >= 0, 'storage object')
    for name, v in views.items():
        require(type(name) is str and name and seq(v) and len(v) == 4, 'view')
        obj, base, shape, strides = v
        require(type(obj) is str and obj in objects and integer(base), 'view storage')
        require(seq(shape) and seq(strides) and len(shape) == len(strides) >= 1, 'view rank')
        require(all(integer(n) and n >= 0 for n in shape) and all(integer(s) for s in strides), 'view scalars')
    require(seq(body) and len(body) >= 1, 'nonempty body')
    all_refs = []
    for stmt in body:
        require(seq(stmt) and len(stmt) == 2, 'assignment')
        dst, expr = stmt
        require(seq(dst) and len(dst) == 2, 'destination')
        refs = [dst] + references(expr)
        for name, indices in refs:
            require(type(name) is str and name in views and seq(indices), 'reference view')
            require(len(indices) == len(views[name][2]), 'reference rank')
            require(all(seq(a) and len(a) == len(dims) + 1 and all(integer(c) for c in a) for a in indices), 'affine index')
        all_refs.append(refs)
    return all_refs


def location(p, ref, point):
    name, indices = ref
    obj, base, shape, strides = p['views'][name]
    offset = base
    for a, n, stride in zip(indices, shape, strides):
        j = a[-1] + sum(c * i for c, i in zip(a[:-1], point))
        require(0 <= j < n, 'view index out of bounds')
        offset += j * stride
    require(0 <= offset < p['objects'][obj], 'storage offset out of bounds')
    return obj, offset


def source_instances(p):
    # Odometer avoids materializing range(n) pools, especially for empty domains.
    dims = p['dims']
    if any(n == 0 for n in dims):
        return []
    point, out = [0] * len(dims), []
    while True:
        frozen = tuple(point)
        out.extend((frozen, s) for s in range(len(p['body'])))
        k = len(dims) - 1
        while k >= 0:
            point[k] += 1
            if point[k] < dims[k]: break
            point[k] = 0; k -= 1
        if k < 0: return out


def analyze(p):
    refs = check_kernel(p)
    ids = source_instances(p)
    effects = []
    for point, s in ids:
        r = refs[s]
        write = location(p, r[0], point)
        reads = frozenset(location(p, x, point) for x in r[1:])
        effects.append((write, reads))
    edges = []
    for a in range(len(ids)):
        wa, ra = effects[a]
        for b in range(a + 1, len(ids)):
            wb, rb = effects[b]
            if wa in rb:
                edges.append((a, b, 'RAW', wa))
            if wb in ra:
                edges.append((a, b, 'WAR', wb))
            if wa == wb:
                edges.append((a, b, 'WAW', wa))
    return {'ids': ids, 'effects': effects, 'edges': edges}


def lit(n):
    return ('lit', n)


def bound_check(e, scope):
    stack = [e]
    while stack:
        x = stack.pop()
        require(seq(x) and x, 'bound tuple')
        if x[0] == 'lit':
            require(len(x) == 2 and integer(x[1]), 'bound integer')
        elif x[0] == 'var':
            require(len(x) == 2 and type(x[1]) is str and x[1] in scope, 'bound scope')
        else:
            require(x[0] in ('plus', 'min') and len(x) == 3, 'bound opcode')
            stack.extend((x[1], x[2]))


def bound_eval(e, env):
    stack, values = [(e, False)], []
    while stack:
        x, done = stack.pop()
        if x[0] == 'lit': values.append(x[1])
        elif x[0] == 'var': values.append(env[x[1]])
        elif done:
            b, a = values.pop(), values.pop()
            values.append(a + b if x[0] == 'plus' else min(a, b))
        else: stack.extend(((x, True), (x[2], False), (x[1], False)))
    return values[0]


def check_ir(p, ir):
    node, scope = ir, set()
    while seq(node) and node and node[0] == 'for':
        require(len(node) == 6, 'for arity')
        _, name, lo, hi, step, body = node
        require(type(name) is str and name not in scope, 'fresh loop binder')
        require(integer(step) and step > 0, 'positive step')
        bound_check(lo, scope); bound_check(hi, scope)
        scope.add(name); node = body
    require(seq(node) and tuple(node) == ('emit',), 'IR leaf')
    require(all('i' + str(k) in scope for k in range(len(p['dims']))), 'all source coordinates bound')


def walk_ir(p, ir, max_instances=None):
    """Return actual statement-instance sequence and number of stack dispatches K."""
    check_ir(p, ir)
    require(max_instances is None or (integer(max_instances) and max_instances >= 0), 'instance budget')
    env, out, stack, dispatches = {}, [], [('node', ir)], 0
    while stack:
        tag, x = stack.pop(); dispatches += 1
        if tag == 'next':
            node, iterator = x
            try: value = next(iterator)
            except StopIteration:
                env.pop(node[1], None)
                continue
            env[node[1]] = value
            stack.append(('next', (node, iterator)))
            stack.append(('node', node[5]))
        elif x[0] == 'emit':
            require(max_instances is None or len(out) + len(p['body']) <= max_instances, 'too many emitted instances')
            point = tuple(env['i' + str(k)] for k in range(len(p['dims'])))
            require(all(0 <= i < n for i, n in zip(point, p['dims'])), 'emitted point outside domain')
            out.extend((point, s) for s in range(len(p['body'])))
        else:
            lo, hi = bound_eval(x[2], env), bound_eval(x[3], env)
            stack.append(('next', (x, iter(range(lo, hi, x[4])))))
    return out, dispatches


def interchange(p, order):
    check_kernel(p)
    d = len(p['dims'])
    require(seq(order) and len(order) == d and all(integer(k) for k in order) and set(order) == set(range(d)), 'dimension permutation')
    node = ('emit',)
    for k in reversed(order):
        node = ('for', 'i' + str(k), lit(0), lit(p['dims'][k]), 1, node)
    return node


def tile(p, sizes):
    check_kernel(p)
    d = len(p['dims'])
    require(seq(sizes) and len(sizes) == d and all(integer(b) and b > 0 for b in sizes), 'positive tile sizes')
    node = ('emit',)
    for k in reversed(range(d)):
        t = ('var', 't' + str(k))
        node = ('for', 'i' + str(k), t, ('min', ('plus', t, lit(sizes[k])), lit(p['dims'][k])), 1, node)
    for k in reversed(range(d)):
        node = ('for', 't' + str(k), lit(0), lit(p['dims'][k]), sizes[k], node)
    return node


def validate(p, ir):
    """Recompute effects/edges; acceptance is sufficient for every integer memory."""
    graph = analyze(p)
    target, dispatches = walk_ir(p, ir, len(graph['ids']))
    expected = set(graph['ids'])
    require(len(target) == len(expected) and len(set(target)) == len(target) and set(target) == expected, 'instance coverage/uniqueness')
    rank = {x: i for i, x in enumerate(target)}
    ranks = [rank[x] for x in graph['ids']]
    for a, b, kind, loc in graph['edges']:
        x, y = graph['ids'][a], graph['ids'][b]
        if ranks[a] >= ranks[b]:
            return {'accepted': False, 'edge': (x, y, kind, loc), 'target_ranks': (ranks[a], ranks[b]), 'dispatches': dispatches}
    return {'accepted': True, 'instances': len(target), 'edges': len(graph['edges']), 'dispatches': dispatches,
            'source_to_target': ranks}


def eval_expr(p, e, point, mem):
    stack, values = [(e, False)], []
    while stack:
        x, done = stack.pop()
        if x[0] == 'const': values.append(x[1])
        elif x[0] == 'load':
            obj, index = location(p, (x[1], x[2]), point)
            values.append(mem[obj][index])
        elif done:
            b, a = values.pop(), values.pop()
            values.append(a + b if x[0] == 'add' else a * b)
        else: stack.extend(((x, True), (x[2], False), (x[1], False)))
    return values[0]


def run(p, ir, memory):
    """Execute actual IR even if dependence-invalid, for explicit counterexamples.
    Require full source-instance coverage; no memory is touched until checks finish.
    The integer memory is copied, so the caller's input remains unchanged.
    """
    check_kernel(p)
    source = source_instances(p)
    ids, dispatches = walk_ir(p, ir, len(source))
    require(len(ids) == len(source) and len(set(ids)) == len(ids) and set(ids) == set(source), 'execution coverage')
    require(type(memory) is dict and set(memory) == set(p['objects']), 'memory objects')
    for obj, n in p['objects'].items():
        require(seq(memory[obj]) and len(memory[obj]) == n and all(integer(x) for x in memory[obj]), 'integer memory')
    # Preflight all accessed locations without constructing quadratic edges.
    refs = check_kernel(p)
    for point, s in source:
        for ref in refs[s]: location(p, ref, point)
    mem, trace = {o: list(xs) for o, xs in memory.items()}, []
    for point, s in ids:
        dst, expr = p['body'][s]
        obj, index = location(p, dst, point)
        value = eval_expr(p, expr, point, mem)
        mem[obj][index] = value
        trace.append((point, s, (obj, index), value))
    return mem, trace, dispatches


def format_ir(node, indent=0):
    def f(e):
        if e[0] == 'lit': return str(e[1])
        if e[0] == 'var': return e[1]
        return '(' + f(e[1]) + '+' + f(e[2]) + ')' if e[0] == 'plus' else 'min(' + f(e[1]) + ',' + f(e[2]) + ')'
    rows = []
    while node[0] == 'for':
        rows.append('  ' * indent + f'for {node[1]} in range({f(node[2])},{f(node[3])},{node[4]}):')
        indent += 1; node = node[5]
    rows.append('  ' * indent + 'emit source body at (i0,...,id-1)')
    return '\n'.join(rows)


def matrix_case(m=3, n=5, k=2):
    a = load('A', aff(1,0,0,0), aff(0,0,1,0))
    b = load('B', aff(0,0,1,0), aff(0,1,0,0))
    c = load('C', aff(1,0,0,0), aff(0,1,0,0))
    p = {'dims': (m,n,k), 'objects': {'a':m*k,'b':k*n,'c':m*n},
         'views': {'A':('a',0,(m,k),(k,1)), 'B':('b',0,(k,n),(n,1)), 'C':('c',0,(m,n),(n,1))},
         'body': (((c[1],c[2]),add(c,mul(a,b))),)}
    return p, {'a':list(range(1,m*k+1)), 'b':list(range(1,k*n+1)), 'c':[0]*(m*n)}


def backward_case():
    # Coordinates (u,v) represent source indices (i,j)=(u+1,v+1).
    dst = ('A',(aff(1,0,1), aff(0,1,1)))
    e = add(load('A',aff(1,0,0),aff(0,1,2)),const(1))
    p = {'dims':(2,3),'objects':{'a':15},'views':{'A':('a',0,(3,5),(5,1))},'body':((dst,e),)}
    return p, {'a':[0]*15}


def main():
    p, mem = matrix_case(); src = interchange(p,(0,1,2)); swapped = interchange(p,(0,2,1)); blocked = tile(p,(2,3,2))
    variants = {}
    baseline = run(p,src,mem)[0]
    for name, ir in [('source',src),('interchange',swapped),('tile',blocked)]:
        cert = validate(p,ir); result, trace, _ = run(p,ir,mem)
        require(cert['accepted'] and result == baseline, 'matrix transformation')
        variants[name] = {'ir':format_ir(ir),'certificate':cert,'C':[result['c'][i:i+5] for i in range(0,15,5)],'updates':len(trace)}
    groups = Counter(tuple(pt[z]//b for z,b in enumerate((2,3,2))) for pt,_ in walk_ir(p,blocked)[0])
    bp,bm = backward_case(); bs = interchange(bp,(0,1)); bswap=interchange(bp,(1,0)); bt=tile(bp,(2,2))
    bad = {}
    for name, ir in [('source',bs),('interchange',bswap),('tile',bt)]:
        state,trace,_=run(bp,ir,bm)
        bad[name]={'certificate':validate(bp,ir),'written':[state['a'][i*5+j] for i in (1,2) for j in (1,2,3)],'order':[x[0] for x in trace]}
    require(bad['source']['written']==[1,1,1,2,2,1], 'source stencil')
    require(not bad['interchange']['certificate']['accepted'] and not bad['tile']['certificate']['accepted'], 'reject reversed edges')
    tests = 0
    for m,n,k in product(range(4),range(5),range(4)):
        q, qm=matrix_case(m,n,k); base=run(q,interchange(q,(0,1,2)),qm)[0]
        for order in permutations(range(3)):
            ir=interchange(q,order); require(validate(q,ir)['accepted'] and run(q,ir,qm)[0]==base,'all permutations');tests+=1
        for sizes in product((1,2,5),repeat=3):
            ir=tile(q,sizes); require(validate(q,ir)['accepted'] and run(q,ir,qm)[0]==base,'fringe tiling');tests+=1
    alias, am = matrix_case()
    alias['views']['A'] = ('c',0,(3,2),(2,1))
    am['c'] = list(range(1,16))
    alias_result = {}
    for name, ir in [('source',src),('interchange',swapped),('tile',blocked)]:
        alias_result[name] = {'certificate':validate(alias,ir),'C':run(alias,ir,am)[0]['c']}
    require(not alias_result['interchange']['certificate']['accepted'] and not alias_result['tile']['certificate']['accepted'], 'new alias edges')
    nonzero = dict(mem); nonzero['c'] = list(range(1,16))
    for ir in (src,swapped,blocked):
        require(run(p,ir,nonzero)[0]['c'] == [x+y for x,y in zip(baseline['c'],nonzero['c'])], 'preserve initial accumulator')
    rejected = []
    for name, action in [('boolean block',lambda:tile(p,(True,3,2))),('zero block',lambda:tile(p,(0,3,2))),('repeat dimension',lambda:interchange(p,(0,0,2))),('boolean dimension',lambda:interchange(p,(False,1,2))),('unbound limit',lambda:walk_ir(p,('for','i0',lit(0),('var','i0'),1,('emit',)))),('missing points',lambda:validate(p,('for','i0',lit(0),lit(0),1,src[5]))),('extra points',lambda:validate(p,('for','repeat',lit(0),lit(2),1,src)))]:
        try: action()
        except ValueError: rejected.append(name)
        else: raise ValueError('failed rejection '+name)
    print(json.dumps({'matrix':variants,'tile_counts':[[key,v] for key,v in groups.items()], 'backward':bad,'alias':alias_result,'matrix_parameter_transforms':tests,'rejected':rejected},ensure_ascii=False,indent=2))


if __name__ == '__main__': main()
