#!/usr/bin/env python3
"""Finite rANS and real package-merge, Python 3.10+, standard library only.
No assert-dependent checks; default mode is read-only. --write DIR emits fixtures.
Teaching RAN1 is deliberately not an existing compression standard.
"""
from __future__ import annotations
import argparse
from dataclasses import dataclass
from itertools import product
import heapq
import json
from pathlib import Path
import struct

U32 = (1 << 32) - 1
HEADER = struct.Struct('>4sBBHHIHI')  # 20 bytes; all unsigned, big-endian
DEFAULT_MODEL = ((65, 2), (66, 1), (67, 1))


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


def equal(actual, expected, label):
    if actual != expected:
        raise ValueError(f'{label}: {actual!r} != {expected!r}')


def uint(value, maximum, label, minimum=0):
    check(type(value) is int and minimum <= value <= maximum, label)


def reject(fn, label):
    try:
        fn()
    except ValueError:
        return
    raise ValueError('accepted invalid case: ' + label)


@dataclass(frozen=True)
class Model:
    entries: tuple
    base: int = 16
    lower: int = 16

    def __post_init__(self):
        check(type(self.entries) is tuple and 1 <= len(self.entries) <= 256, 'model size')
        uint(self.base, 256, 'base', 2)
        check(self.base & (self.base - 1) == 0, 'base must be a power of two')
        uint(self.lower, 1 << 31, 'lower', 1)
        last, total = -1, 0
        for entry in self.entries:
            check(type(entry) is tuple and len(entry) == 2, 'model entry')
            s, f = entry
            uint(s, 255, 'symbol')
            uint(f, 4096, 'frequency', 1)
            check(s > last, 'symbol order/duplicate')
            last, total = s, total + f
        check(total <= 4096 and self.lower % total == 0, 'M must divide L; M <= 4096')
        check(self.base * self.lower <= 1 << 32, '32-bit state window')

    @property
    def total(self):
        return sum(f for _, f in self.entries)

    def tables(self):
        forward, inverse, c = {}, [], 0
        for s, f in self.entries:
            forward[s] = (f, c)
            inverse.extend([(s, f, c)] * f)
            c += f
        return forward, inverse


def rans_encode(data: bytes, model=Model(DEFAULT_MODEL), trace=False):
    check(type(data) is bytes, 'bytes input required')
    check(len(data) <= 1_000_000, 'encoder block budget')
    forward, _ = model.tables()
    M, b, L = model.total, model.base, model.lower
    x, stack, rows = L, [], []
    for pos in range(len(data) - 1, -1, -1):
        s, before = data[pos], x
        check(s in forward, 'symbol missing from model')
        f, c = forward[s]
        threshold = b * (L // M) * f  # compute in >=64 bits in a fixed-width port
        emitted = []
        while x >= threshold:
            digit = x % b
            stack.append(digit)
            emitted.append(digit)
            x //= b
        normalized = x
        x = (x // f) * M + x % f + c
        check(L <= x < b * L, 'encoder invariant')
        if trace:
            rows.append({'index': pos, 'symbol': chr(s), 'before': before,
                         'threshold': threshold, 'extracted': emitted,
                         'normalized': normalized, 'after': x})
    return x, list(reversed(stack)), rows


def rans_decode(x, digits, n, model=Model(DEFAULT_MODEL), trace=False,
                max_output=1_000_000, max_digits=8_000_000):
    uint(max_output, U32, 'output budget')
    uint(max_digits, U32, 'digit budget')
    uint(n, max_output, 'declared output budget')
    check(type(digits) in (list, tuple), 'digit sequence')
    check(len(digits) <= max_digits, 'digit budget exceeded')
    M, b, L = model.total, model.base, model.lower
    uint(x, b * L - 1, 'final state', L)
    for digit in digits:
        uint(digit, b - 1, 'digit range')
    _, inverse = model.tables()
    out, rows, p = bytearray(), [], 0
    for pos in range(n):
        before = x
        s, f, c = inverse[x % M]
        x = f * (x // M) + x % M - c
        raw, consumed = x, []
        while x < L:
            check(p < len(digits), 'truncated digit stream')
            d = digits[p]
            p += 1
            consumed.append(d)
            x = b * x + d
        check(L <= x < b * L, 'decoder invariant')
        out.append(s)
        if trace:
            rows.append({'index': pos, 'symbol': chr(s), 'before': before,
                         'raw': raw, 'consumed': consumed, 'after': x})
    equal(x, L, 'terminal initial-state check')
    equal(p, len(digits), 'trailing digits')
    return bytes(out), rows


def ran1_encode(data, model=Model(DEFAULT_MODEL)):
    check(model.base == 16 and model.lower == 16 and model.total == 4, 'RAN1 fixed parameters')
    x, digits, rows = rans_encode(data, model, trace=True)
    body = bytearray()
    for k in range(0, len(digits), 2):
        body.append((digits[k] << 4) | (digits[k+1] if k+1 < len(digits) else 0))
    header = HEADER.pack(b'RAN1', 4, 2, 16, len(model.entries), len(data), x, len(digits))
    table = bytes(v for pair in model.entries for v in pair)
    return header + table + body, {'state': x, 'digits': digits, 'encode_trace': rows}


def ran1_decode(blob, max_output=1_000_000, max_digits=8_000_000):
    uint(max_output, U32, 'output budget')
    uint(max_digits, U32, 'digit budget')
    check(type(blob) is bytes and len(blob) >= HEADER.size, 'truncated fixed header')
    magic, logb, logm, L, m, n, state, nd = HEADER.unpack_from(blob)
    check((magic, logb, logm, L) == (b'RAN1', 4, 2, 16), 'unsupported RAN1 parameters')
    check(1 <= m <= 4, 'model count')
    check(n <= max_output and nd <= max_digits, 'declared resource budget')
    size = HEADER.size + 2*m + (nd+1)//2
    equal(len(blob), size, 'exact physical file length')
    entries = tuple((blob[20+2*i], blob[21+2*i]) for i in range(m))
    model = Model(entries)
    equal(model.total, 4, 'RAN1 frequency total')
    payload = blob[20+2*m:]
    if nd % 2:
        check(payload[-1] & 15 == 0, 'nonzero alignment nibble')
    # Check budgets and exact physical length before materializing digit list/output.
    digits = [(payload[i//2] >> (4 if i % 2 == 0 else 0)) & 15 for i in range(nd)]
    decoded, rows = rans_decode(state, digits, n, model, trace=True,
                               max_output=max_output, max_digits=max_digits)
    return decoded, {'model': entries, 'state': state, 'digits': digits, 'decode_trace': rows}


@dataclass
class Item:
    weight: int
    tie: int
    serial: int
    leaf: tuple | None = None
    children: tuple | None = None

    @property
    def key(self):
        return self.weight, self.tie, self.serial


def merge(a, b):
    result, i, j = [], 0, 0
    while i < len(a) and j < len(b):
        if a[i].key <= b[j].key:
            result.append(a[i]); i += 1
        else:
            result.append(b[j]); j += 1
    result.extend(a[i:]); result.extend(b[j:])
    return result


def package_merge(weights, limit, trace=False):
    """Return original-symbol-order lengths and an explicit nodeset certificate.
    Leaves are (symbol, level), each with width 2^-level and weight w_symbol.
    This is the O(nL)-space package-tree algorithm, not a DP or exhaustive search.
    """
    check(type(weights) in (list, tuple) and 1 <= len(weights) <= 256, 'symbol budget')
    uint(limit, 32, 'length limit', 1)
    for w in weights:
        uint(w, U32, 'positive integer weight', 1)
    n = len(weights)
    check(n <= (1 << limit), 'infeasible: n > 2^L')
    if n == 1:
        return [1], {'singleton_convention': True, 'cost': weights[0],
                     'nodeset': [[0, 1]], 'layers': [], 'created_nodes': 1}
    order = sorted(range(n), key=lambda i: (weights[i], i))
    ranks = {i: rank+1 for rank, i in enumerate(order)}
    serial, current, layers = 0, [], []
    for level in range(limit, 0, -1):
        fresh = []
        for i in order:
            fresh.append(Item(weights[i], ranks[i], serial, leaf=(i, level)))
            serial += 1
        if current:
            packages = []
            for k in range(0, len(current)-1, 2):
                a, b = current[k:k+2]
                packages.append(Item(a.weight+b.weight, a.tie+b.tie, serial,
                                     children=(a, b)))
                serial += 1
            current = merge(fresh, packages)
        else:
            current = fresh
        if trace:
            layers.append({'level': level, 'width_denominator': 1 << level,
                           'items': [{'weight': z.weight, 'tie': z.tie, 'serial': z.serial,
                                      'leaf': z.leaf,
                                      'children': [v.serial for v in z.children] if z.children else None}
                                     for z in current]})
    count = 2*n-2
    check(len(current) >= count, 'insufficient width')
    selected = current[:count]
    stack, nodeset = selected[:], []
    while stack:
        item = stack.pop()
        if item.leaf is not None:
            nodeset.append(item.leaf)
        else:
            stack.extend(item.children)
    # Linear-time certificate ordering: no comparison sort of the nL coins.
    seen, lengths = [[False]*(limit+1) for _ in range(n)], [0]*n
    for i, level in nodeset:
        check(not seen[i][level], 'unique coin identities')
        seen[i][level] = True
        lengths[i] += 1
    nodeset = [(i, level) for i in range(n) for level in range(1, limit+1)
               if seen[i][level]]
    # Certificate checks are consequences, not repairs or substitutions.
    equal(nodeset, [(i, level) for i, length in enumerate(lengths)
                    for level in range(1, length+1)], 'nodeset prefix closure')
    check(all(1 <= v <= limit for v in lengths), 'bounded positive lengths')
    equal(sum(1 << (limit-v) for v in lengths), 1 << limit, 'complete Kraft equality')
    cost = sum(w*v for w, v in zip(weights, lengths))
    equal(cost, sum(x.weight for x in selected), 'tree/package cost equality')
    return lengths, {'cost': cost, 'nodeset': nodeset, 'layers': layers,
                     'selected_top_serials': [x.serial for x in selected],
                     'created_nodes': serial, 'singleton_convention': False}


def canonical(lengths):
    check(lengths and all(type(v) is int and 1 <= v <= 32 for v in lengths), 'code lengths')
    buckets = [[] for _ in range(max(lengths)+1)]
    for i, length in enumerate(lengths):
        buckets[length].append(i)  # Stable symbol order within each length.
    code, previous, out = 0, 0, {}
    for length, symbols in enumerate(buckets):
        for i in symbols:
            code <<= length - previous
            check(code < 1 << length, 'oversubscribed code lengths')
            out[i] = format(code, f'0{length}b')
            code += 1
            previous = length
    return out


def huffman_lengths(weights):
    heap, serial = [(w, i, i) for i, w in enumerate(weights)], len(weights)
    heapq.heapify(heap)
    while len(heap) > 1:
        a, b = heapq.heappop(heap), heapq.heappop(heap)
        heapq.heappush(heap, (a[0]+b[0], serial, (a[2], b[2])))
        serial += 1
    out, stack = [0]*len(weights), [(heap[0][2], 0)]
    while stack:
        node, depth = stack.pop()
        if isinstance(node, int):
            out[node] = max(1, depth)
        else:
            stack.extend([(node[0], depth+1), (node[1], depth+1)])
    return out


def exhaustive_oracle(weights, limit):
    """Independent feasibility enumeration; never called by package_merge."""
    best, winners, feasible = None, [], 0
    for lengths in product(range(1, limit+1), repeat=len(weights)):
        if sum(1 << (limit-v) for v in lengths) <= 1 << limit:
            feasible += 1
            cost = sum(w*v for w, v in zip(weights, lengths))
            if best is None or cost < best:
                best, winners = cost, [lengths]
            elif cost == best:
                winners.append(lengths)
    return {'cost': best, 'winners': winners, 'feasible_vectors': feasible}


def inverse_enumeration(y, model):
    """Independent one-step inverse by enumerating C_s(x), without inverse formula."""
    M, c, matches = model.total, 0, []
    for s, f in model.entries:
        for x in range(model.base*model.lower):
            if (x//f)*M+x%f+c == y:
                matches.append((s, x))
        c += f
    equal(len(matches), 1, 'enumerated inverse uniqueness')
    return matches[0]


def run_tests():
    message = b'ABCABAA'
    blob, encoded = ran1_encode(message)
    decoded, parsed = ran1_decode(blob)
    equal(decoded, message, 'non-palindrome endpoint')
    count = 0
    for n in range(8):
        for word in product(b'ABC', repeat=n):
            data = bytes(word)
            f, _ = ran1_encode(data)
            equal(ran1_decode(f)[0], data, 'all ternary roundtrips')
            count += 1
    one_steps = 0
    model = Model(DEFAULT_MODEL)
    forward, inv = model.tables()
    for y in range(16, 256):
        s, f, c = inv[y%4]
        equal(inverse_enumeration(y, model), (s, f*(y//4)+y%4-c), 'independent inverse')
        one_steps += 1
    negative = 0
    for end in range(len(blob)):
        reject(lambda end=end: ran1_decode(blob[:end]), f'byte truncation {end}'); negative += 1
    alterations = [blob+b'\0']
    def altered(offset, fmt, value):
        b = bytearray(blob); struct.pack_into(fmt, b, offset, value); return bytes(b)
    alterations += [altered(14, '>H', 15), altered(14, '>H', 256),
                    altered(10, '>I', 6), altered(10, '>I', 8),
                    altered(16, '>I', 1), altered(16, '>I', 3),
                    altered(8, '>H', 5), altered(21, '>B', 0),
                    altered(22, '>B', 65), altered(23, '>B', 2),
                    altered(4, '>B', 8), altered(10, '>I', U32)]
    for k, bad in enumerate(alterations):
        reject(lambda bad=bad: ran1_decode(bad), f'bad framing {k}'); negative += 1
    # Find and corrupt an actual odd-nibble frame, separate from the endpoint.
    odd, _ = ran1_encode(b'AAAAA')
    check(HEADER.unpack_from(odd)[-1] % 2 == 1, 'odd-nibble fixture')
    reject(lambda: ran1_decode(odd[:-1]+bytes([odd[-1]|1])), 'padding'); negative += 1
    x, digits, _ = rans_encode(message)
    reject(lambda: rans_decode(x, digits+[0], len(message)), 'extra digit'); negative += 1
    reject(lambda: rans_decode(x, digits[:-1], len(message)), 'missing digit'); negative += 1
    reject(lambda: rans_decode(17, [], 0), 'wrong empty final state'); negative += 1
    reject(lambda: ran1_decode(blob, max_output=-1), 'negative budget'); negative += 1
    reject(lambda: ran1_decode(blob, max_output=6), 'output cap'); negative += 1
    reject(lambda: ran1_decode(blob, max_digits=0), 'digit cap'); negative += 1
    for data in [b'', b'Z', b'Z'*1000]:
        f, _ = ran1_encode(data, Model(((90, 4),)))
        equal(ran1_decode(f)[0], data, 'singleton f=M uses explicit n')
    # Structural migration: selected positive M=4 models and a byte-base model.
    migrated = 0
    models = [Model(((65,1),(66,3))), Model(((65,3),(66,1))),
              Model(((65,1),(66,1),(67,1),(68,1))),
              Model(((65,2048),(66,1024),(67,1024)), 256, 1 << 23)]
    for m in models:
        alphabet = tuple(s for s, _ in m.entries)
        for n in range(5):
            for word in product(alphabet, repeat=n):
                data = bytes(word)
                x, ds, _ = rans_encode(data, m)
                equal(rans_decode(x, ds, n, m)[0], data, 'structural migration')
                migrated += 1
    block_files = [ran1_encode(part)[0] for part in [b'ABCA', b'BAA']]
    equal(b''.join(ran1_decode(f)[0] for f in block_files), message, 'framed block migration')
    reject(lambda: ran1_decode(b''.join(block_files)), 'unframed concatenation'); negative += 1
    byte_model = models[-1]
    byte_x, byte_digits, _ = rans_encode(message, byte_model)
    weights = [1,1,2,3,5,8]
    lengths, package = package_merge(weights, 3, trace=True)
    oracle = exhaustive_oracle(weights, 3)
    equal(lengths, [3,3,3,3,2,2], 'length-limited endpoint')
    equal(package['cost'], 47, 'length-limited endpoint cost')
    equal(oracle['feasible_vectors'], 22, 'full independent oracle count')
    equal(oracle['winners'], [tuple(lengths)], 'unique oracle winner')
    cross = 0
    for n in range(2, 6):
        for limit in range(1, 4):
            if n > 1 << limit:
                continue
            for ws in product(range(1,4), repeat=n):
                actual, certificate = package_merge(ws, limit)
                expected = exhaustive_oracle(ws, limit)
                equal(certificate['cost'], expected['cost'], 'package-merge/oracle cost')
                check(tuple(actual) in expected['winners'], 'package-merge/oracle membership')
                cross += 1
    migration = []
    for limit in [2,3,4,5]:
        if len(weights) > 1 << limit:
            reject(lambda: package_merge(weights, limit), 'infeasible limit')
            migration.append({'limit': limit, 'feasible': False})
        else:
            ls, cert = package_merge(weights, limit)
            o = exhaustive_oracle(weights, limit)
            equal(cert['cost'], o['cost'], 'limit migration oracle')
            migration.append({'limit': limit, 'feasible': True, 'lengths': ls,
                              'cost': cert['cost'], 'optimal_vectors': len(o['winners'])})
    for ws, limit in [([],3), ([1,0],2), ([1,-1],2), ([True,1],2), ([1,2],0),
                      ([1,2],33), ([1,2],True), ([U32+1,1],2)]:
        reject(lambda ws=ws, limit=limit: package_merge(ws, limit), 'bad optimization input')
    equal(package_merge([7], 3)[0], [1], 'positive singleton convention')
    huge, huge_cert = package_merge([U32]*256, 32)
    equal(huge, [8]*256, 'large equal weights')
    check(huge_cert['cost'] < 1 << 64, '64-bit cost bound')
    repeated, repeated_cert = package_merge([1,1,1], 3)
    ties = exhaustive_oracle([1,1,1],3)
    equal(ties['cost'], repeated_cert['cost'], 'tie case')
    unbounded = huffman_lengths(weights)
    equal(unbounded, [5,5,4,3,2,1], 'Huffman comparison')
    clipped = [min(3,v) for v in unbounded]
    reject(lambda: canonical(clipped), 'clipped overfull code')
    codes = canonical(lengths)
    bits = ''.join(codes[i] for i, w in enumerate(weights) for _ in range(w))
    equal(len(bits), 47, 'explicit weighted message size')
    reverse = {v:k for k,v in codes.items()}
    read, prefix = [], ''
    for bit in bits:
        prefix += bit
        if prefix in reverse:
            read.append(reverse[prefix]); prefix = ''
    equal(prefix, '', 'complete prefix parsing')
    equal(read, [i for i,w in enumerate(weights) for _ in range(w)], 'canonical message recovery')
    return {'rans': {'input': message.decode(), 'file_hex': blob.hex(' '),
                     'file_bytes': len(blob), **encoded, **parsed,
                     'ternary_roundtrips': count, 'independent_one_steps': one_steps,
                     'negative_cases': negative, 'migration_roundtrips': migrated,
                     'odd_nibble_file_hex': odd.hex(' '),
                     'split_block_bytes': [len(f) for f in block_files],
                     'byte_base_endpoint': {'state': byte_x, 'digits': byte_digits}},
            'package_merge': {'weights': weights, 'limit': 3, 'lengths': lengths,
                              'canonical': codes, 'unbounded': unbounded,
                              'unbounded_cost': sum(w*v for w,v in zip(weights,unbounded)),
                              'clipped': clipped, 'clipped_kraft_eighths': sum(1<<(3-v) for v in clipped),
                              'oracle': oracle, 'cross_checked_instances': cross,
                              'migration': migration, 'equal_weight_example': repeated,
                              'equal_weight_optimal_vectors': len(ties['winners']), **package}}


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument('--write', type=Path, help='explicitly emit three teaching fixtures here')
    args = parser.parse_args()
    result = run_tests()
    if args.write:
        args.write.mkdir(parents=True, exist_ok=True)
        (args.write/'foundations-abcabaa.txt').write_bytes(b'ABCABAA')
        (args.write/'foundations-abcabaa.ran1').write_bytes(ran1_encode(b'ABCABAA')[0])
        (args.write/'foundations-bounded-coders-fixture.json').write_text(
            json.dumps(result, ensure_ascii=False, indent=2)+'\n')
    print(json.dumps(result, ensure_ascii=False, indent=2))


if __name__ == '__main__':
    main()
