#!/usr/bin/env python3
"""Exact integer, sequential-event IR for unswitching and counted unrolling.
Python 3.10+, standard library only. No files, network, or subprocesses.
Expressions: int, variable, or (op, left, right); quot/rem require literal k>0.
Statements: set, emit, if, repeat. A repeat captures its count on entry.
The public execution budget counts statement/control visits, not source trips.
"""
from copy import deepcopy
from collections import Counter
import json
import random

OPS = {'add', 'sub', 'mul', 'lt', 'le', 'gt', 'ge', 'eq', 'ne', 'quot', 'rem'}


def need(condition, message):
    if not condition:
        raise ValueError(message)


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


def reads(e):
    if integer(e):
        return set()
    if isinstance(e, str) and e:
        return {e}
    need(isinstance(e, tuple) and len(e) == 3 and e[0] in OPS, 'invalid expression')
    if e[0] in {'quot', 'rem'}:
        need(integer(e[2]) and e[2] > 0, 'quot/rem need a positive literal divisor')
    return reads(e[1]) | reads(e[2])


def check_block(block, defined, inside_loop=False):
    need(isinstance(block, tuple), 'block must be a tuple')
    known = set(defined)
    for s in block:
        need(isinstance(s, tuple) and len(s) >= 1, 'invalid statement')
        tag = s[0]
        if tag == 'set':
            need(len(s) == 3 and isinstance(s[1], str) and bool(s[1]), 'invalid assignment')
            need(reads(s[2]) <= known, 'read before definite assignment')
            known.add(s[1])
        elif tag == 'emit':
            need(len(s) == 3 and isinstance(s[1], str) and bool(s[1]), 'invalid event')
            need(reads(s[2]) <= known, 'event read before definite assignment')
        elif tag == 'if':
            need(len(s) == 4 and reads(s[1]) <= known, 'invalid conditional or read')
            a = check_block(s[2], known, inside_loop)
            b = check_block(s[3], known, inside_loop)
            known = a & b
        elif tag == 'repeat':
            need(len(s) == 3 and not inside_loop, 'nested loops are outside the model')
            need(reads(s[1]) <= known, 'count read before definite assignment')
            check_block(s[2], known, True)  # zero trips add no definite definitions
        else:
            raise ValueError('unsupported statement')
    return known


def validate(p):
    need(isinstance(p, dict) and set(p) == {'params', 'init', 'code', 'returns'}, 'program fields')
    need(isinstance(p['params'], tuple) and all(isinstance(x, str) and x for x in p['params']), 'parameters')
    need(len(set(p['params'])) == len(p['params']), 'duplicate parameter')
    need(isinstance(p['init'], tuple) and all(isinstance(s, tuple) and s and s[0] == 'set' for s in p['init']), 'initialization only assigns')
    known = check_block(p['init'], set(p['params']))
    known = check_block(p['code'], known)
    need(isinstance(p['returns'], tuple) and all(reads(e) <= known for e in p['returns']), 'return before definite assignment')
    return p


def evaluate(e, env, counts):
    if integer(e):
        return e
    if isinstance(e, str):
        return env[e]
    op, a, b = e
    x, y = evaluate(a, env, counts), evaluate(b, env, counts)
    counts[op] += 1
    if op == 'add': return x + y
    if op == 'sub': return x - y
    if op == 'mul': return x * y
    if op == 'quot': return x // y
    if op == 'rem': return x % y
    return int({'lt': x < y, 'le': x <= y, 'gt': x > y, 'ge': x >= y, 'eq': x == y, 'ne': x != y}[op])


class BudgetStop(Exception):
    pass


def execute(p, inputs, budget=1000000):
    validate(p)
    need(set(inputs) == set(p['params']) and all(integer(x) for x in inputs.values()), 'exact integer inputs')
    need(integer(budget) and budget >= 0, 'nonnegative execution budget')
    env, events, counts = dict(inputs), [], Counter()
    left = budget

    def tick():
        nonlocal left
        if left == 0:
            raise BudgetStop()
        left -= 1

    def block(code):
        for s in code:
            tick()
            tag = s[0]
            if tag == 'set':
                env[s[1]] = evaluate(s[2], env, counts)
                counts['assignments'] += 1
            elif tag == 'emit':
                events.append((s[1], evaluate(s[2], env, counts)))
                counts['events'] += 1
            elif tag == 'if':
                counts['if_tests'] += 1
                block(s[2] if evaluate(s[1], env, counts) != 0 else s[3])
            else:
                n = evaluate(s[1], env, counts)
                need(n >= 0, 'negative runtime trip count')
                for j in range(n + 1):
                    tick()
                    counts['loop_tests'] += 1
                    if j == n:
                        break
                    block(s[2])
    try:
        block(p['init'])
        block(p['code'])
        result = tuple(evaluate(e, env, counts) for e in p['returns'])
        status = 'returned'
    except BudgetStop:
        result, status = None, 'budget_exhausted'
    return {'status': status, 'result': result, 'events': events, 'environment': env,
            'counts': dict(sorted(counts.items())), 'visits': budget - left}


def writes(block):
    found = set()
    for s in block:
        if s[0] == 'set': found.add(s[1])
        elif s[0] == 'if': found |= writes(s[2]) | writes(s[3])
        elif s[0] == 'repeat': found |= writes(s[2])
    return found


def names(p):
    # Validation guarantees every read belongs to a parameter or a destination.
    return set(p['params']) | writes(p['init']) | writes(p['code'])


def fresh(used, cursor):
    while '__trip_' + str(cursor[0]) in used:
        cursor[0] += 1
    name = '__trip_' + str(cursor[0])
    cursor[0] += 1
    used.add(name)
    return name


def statement_count(block):
    total = 0
    for s in block:
        total += 1
        if s[0] == 'if': total += statement_count(s[2]) + statement_count(s[3])
        elif s[0] == 'repeat': total += statement_count(s[2])
    return total


def at_path(body, path):
    # A path is (statement_index, branch_index, statement_index, ...).
    # branch_index is 2 or 3, selecting that if's true or false block.
    need(isinstance(path, tuple) and len(path) % 2 == 1, 'odd nonempty conditional path')
    block = body
    for pos in range(0, len(path), 2):
        i = path[pos]
        need(integer(i) and 0 <= i < len(block), 'path index')
        s = block[i]
        need(s[0] == 'if', 'path must select conditionals')
        if pos + 1 == len(path): return s
        branch = path[pos + 1]
        need(integer(branch) and branch in (2, 3), 'path branch')
        block = s[branch]


def specialize(body, path, take_true):
    i = path[0]
    if len(path) == 1:
        return deepcopy(body[:i] + body[i][2 if take_true else 3] + body[i + 1:])
    s = list(deepcopy(body[i])); branch = path[1]
    s[branch] = specialize(body[i][branch], path[2:], take_true)
    return deepcopy(body[:i]) + (tuple(s),) + deepcopy(body[i + 1:])


def unswitch(p, path, max_statements=10000):
    validate(p)
    need(integer(max_statements) and max_statements >= 0, 'nonnegative code budget')
    need(len(p['code']) == 1 and p['code'][0][0] == 'repeat', 'unswitch expects one top-level repeat')
    _, count, body = p['code'][0]
    chosen = at_path(body, path)
    available = check_block(p['init'], set(p['params']))
    if not reads(chosen[1]) <= available:
        return {'status': 'refused', 'reason': 'predicate is not initialized at loop entry', 'program': None}
    if reads(chosen[1]) & writes(body):
        return {'status': 'refused', 'reason': 'predicate reads a variable written in the loop', 'program': None}
    b = statement_count(body)
    needed = 2 + 2 * b - statement_count(chosen[2]) - statement_count(chosen[3])
    if needed > max_statements:
        return {'status': 'size_limit', 'needed': needed, 'program': None}
    used = names(p); cursor = [0]; n = fresh(used, cursor)
    target = deepcopy(p)
    target['code'] = (('set', n, deepcopy(count)),
        ('if', deepcopy(chosen[1]), (('repeat', n, specialize(body, path, True)),),
         (('repeat', n, specialize(body, path, False)),)))
    validate(target)
    need(statement_count(target['code']) == needed, 'internal size mismatch')
    return {'status': 'converted', 'program': target, 'captured_count': n, 'statements': needed}


def unroll(p, factor, max_statements=10000):
    validate(p)
    need(integer(factor) and factor >= 1, 'positive literal unrolling factor')
    need(integer(max_statements) and max_statements >= 0, 'nonnegative code budget')

    def planned(code):
        result = 0
        for s in code:
            if s[0] == 'repeat': result += 3 + (factor + 1) * statement_count(s[2])
            elif s[0] == 'if': result += 1 + planned(s[2]) + planned(s[3])
            else: result += 1
        return result
    needed = planned(p['code'])
    if needed > max_statements:
        return {'status': 'size_limit', 'needed': needed, 'program': None}
    used, captured, cursor = names(p), [], [0]

    def transform(code):
        out = []
        for s in code:
            if s[0] == 'repeat':
                n = fresh(used, cursor); captured.append(n)
                # Empty bodies need no factor-sized iteration during compilation.
                copies = tuple(deepcopy(x) for _ in range(factor) for x in s[2]) if s[2] else ()
                out.extend((('set', n, deepcopy(s[1])),
                    ('repeat', ('quot', n, factor), copies),
                    ('repeat', ('rem', n, factor), deepcopy(s[2]))))
            elif s[0] == 'if': out.append(('if', deepcopy(s[1]), transform(s[2]), transform(s[3])))
            else: out.append(deepcopy(s))
        return tuple(out)
    target = deepcopy(p); target['code'] = transform(p['code'])
    validate(target)
    need(statement_count(target['code']) == needed, 'internal size mismatch')
    return {'status': 'converted', 'program': target, 'captured_counts': captured, 'statements': needed}


def example():
    return {'params': ('n', 'mode', 'seed', 'start'),
        'init': (('set', 's', 'seed'), ('set', 'i', 'start')),
        'code': (('repeat', 'n', (
            ('emit', 'before', 'i'),
            ('if', ('gt', 'mode', 0), (('set', 's', ('add', ('mul', 2, 's'), 'i')),),
             (('set', 's', ('sub', 's', 'i')),)),
            ('emit', 'after', 's'), ('set', 'i', ('add', 'i', 1)))),),
        'returns': ('s', 'i')}


def observe(run):
    return run['status'], run['result'], run['events']


def compare(p, inputs, factor=4):
    u = unswitch(p, (1,)); need(u['status'] == 'converted', 'example unswitch refused')
    a = unroll(p, factor); b = unroll(u['program'], factor)
    variants = [p, u['program'], a['program'], b['program']]
    runs = [execute(v, inputs) for v in variants]
    need(all(r['status'] == 'returned' and observe(r) == observe(runs[0]) for r in runs), 'observation mismatch')
    return {'runs': runs, 'statement_nodes': [statement_count(v['code']) for v in variants]}


def self_test():
    p = example(); rng = random.Random(2307)
    for _ in range(1000):
        inputs = {'n': rng.randrange(0, 28), 'mode': rng.randrange(-2, 3),
                  'seed': rng.randrange(-5, 6), 'start': rng.randrange(-4, 5)}
        compare(p, inputs, rng.randrange(1, 9))
    main = compare(p, {'n': 10, 'mode': 1, 'seed': 1, 'start': 2})
    zero = compare(p, {'n': 0, 'mode': 1, 'seed': 1, 'start': 2})
    negative = compare(p, {'n': 10, 'mode': 0, 'seed': 1, 'start': 2})
    changed = deepcopy(p)
    body = changed['code'][0][2] + (('set', 'mode', ('sub', 'mode', 1)),)
    changed['code'] = (('repeat', 'n', body),)
    need(unswitch(changed, (1,))['status'] == 'refused', 'changing predicate accepted')
    # The original count is captured even if the body changes its source variable.
    mutable = deepcopy(p)
    mutable['code'] = (('repeat', 'n', p['code'][0][2] + (('set', 'n', 0),)),)
    compare(mutable, {'n': 10, 'mode': 1, 'seed': 1, 'start': 2})
    huge = unroll(p, 10**9, max_statements=10000)
    need(huge['status'] == 'size_limit' and huge['program'] is None, 'size budget not enforced')
    return {'status': 'PASS', 'inputs': 1000, 'main': main, 'zero': zero,
            'negative_result': negative['runs'][0]['result'], 'changing_predicate': unswitch(changed, (1,))['reason'],
            'large_factor': huge, 'scope': 'finite pure-total counted loops; exact ordered events, not elapsed-time claims'}


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