#!/usr/bin/env python3
"""A21. Offline historical-band matching and exact fixed gamma-cost LZ77 parsing.
Standard library, stdout only. This is NOT DEFLATE and not a standalone file format.
Indices are zero-based; half-open ranges. Checks survive python -O.
"""
from itertools import product
import json
import random


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


def integer(x, low=0):
    return type(x) is int and x >= low


def parameters(data, window, cap):
    need(type(data) is bytes, 'input must be immutable bytes')
    need(integer(window) and integer(cap), 'nonnegative integer window and match cap')


def counting_order(order, key, classes):
    counts = [0] * classes
    for x in order:
        counts[key[x]] += 1
    starts, total = [], 0
    for c in counts:
        starts.append(total)
        total += c
    out = [0] * len(order)
    for x in order:
        k = key[x]
        out[starts[k]] = x
        starts[k] += 1
    return out


def suffix_array(data):
    # Remap ALL bytes to1..256; unique virtual sentinel0 is not an input byte.
    a = [b + 1 for b in data] + [0]
    N = len(a)
    order = counting_order(list(range(N)), a, 257)
    rank = [0] * N
    for j in range(1, N):
        rank[order[j]] = rank[order[j-1]] + (a[order[j]] != a[order[j-1]])
    k = 1
    while k < N:
        order = counting_order([(p-k) % N for p in order], rank, rank[order[-1]]+1)
        new = [0] * N
        for j in range(1, N):
            u, v = order[j-1], order[j]
            new[v] = new[u] + ((rank[u], rank[(u+k) % N]) != (rank[v], rank[(v+k) % N]))
        rank = new
        k *= 2
    need(order[0] == len(data), 'virtual sentinel not first')
    return order[1:]


class MinTree:
    """Minimum (value,index), with leftmost index on equal values."""
    def __init__(self, values, infinity):
        self.n = len(values)
        self.size = 1
        while self.size < self.n:
            self.size *= 2
        self.empty = (infinity, self.n)
        self.tree = [self.empty] * (2*self.size)
        for i, v in enumerate(values):
            self.tree[self.size+i] = (v, i)
        for i in range(self.size-1, 0, -1):
            self.tree[i] = min(self.tree[2*i], self.tree[2*i+1])

    def assign(self, i, value):
        p = self.size+i
        self.tree[p] = (value, i)
        p //= 2
        while p:
            self.tree[p] = min(self.tree[2*p], self.tree[2*p+1])
            p //= 2

    def minimum(self, left, right):
        need(0 <= left <= right <= self.n, 'range bounds')
        a, b, answer = left+self.size, right+self.size, self.empty
        while a < b:
            if a & 1:
                answer = min(answer, self.tree[a]); a += 1
            if b & 1:
                b -= 1; answer = min(answer, self.tree[b])
            a //= 2; b //= 2
        return answer


class Fenwick:
    def __init__(self, n):
        self.n = n
        self.tree = [0] * (n+1)

    def add(self, i, delta):
        i += 1
        while i <= self.n:
            self.tree[i] += delta
            i += i & -i

    def prefix(self, end):
        answer = 0
        while end:
            answer += self.tree[end]
            end -= end & -end
        return answer

    def kth(self, k):
        need(1 <= k <= self.prefix(self.n), 'absent order statistic')
        p = 0
        step = 1 << self.n.bit_length()
        while step:
            q = p+step
            if q <= self.n and self.tree[q] < k:
                p = q
                k -= self.tree[q]
            step //= 2
        return p  # Zero-based rank containing the original k-th1.


class SuffixIndex:
    def __init__(self, data):
        need(type(data) is bytes, 'input must be immutable bytes')
        self.data = data
        self.n = len(data)
        self.sa = suffix_array(data)
        self.rank = [0] * self.n
        for r, p in enumerate(self.sa):
            self.rank[p] = r
        self.lcp = [0] * self.n
        h = 0
        for p in range(self.n):
            r = self.rank[p]
            if r == 0:
                h = 0
                continue
            q = self.sa[r-1]
            while p+h < self.n and q+h < self.n and data[p+h] == data[q+h]:
                h += 1
            self.lcp[r] = h
            h = max(0, h-1)
        self.rmq = MinTree(self.lcp, self.n+1)

    def common_prefix(self, p, q):
        need(integer(p) and integer(q) and p < self.n and q < self.n, 'suffix position')
        if p == q:
            return self.n-p
        a, b = sorted((self.rank[p], self.rank[q]))
        return self.rmq.minimum(a+1, b+1)[0]

    def band(self, lo, hi, cap, trace=False):
        need(integer(lo, 1) and integer(hi, lo) and integer(cap), 'historical band parameters')
        need(type(trace) is bool, 'trace flag')
        active = Fenwick(self.n)
        matches, rows = [], []
        for p in range(self.n):
            entering, leaving = p-lo, p-hi-1
            if entering >= 0:
                active.add(self.rank[entering], 1)
            if leaving >= 0:
                active.add(self.rank[leaving], -1)
            r = self.rank[p]
            smaller, total = active.prefix(r), active.prefix(self.n)
            candidate_ranks = []
            if smaller:
                candidate_ranks.append(active.kth(smaller))
            if smaller < total:
                candidate_ranks.append(active.kth(smaller+1))
            best, source, candidates = 0, None, []
            for s in candidate_ranks:
                q = self.sa[s]
                length = min(cap, self.n-p, self.common_prefix(p, q))
                if trace:
                    candidates.append({'source': q, 'rank': s, 'length': length})
                if length > best:  # Equal lengths retain the predecessor, if present.
                    best, source = length, q
            matches.append((best, source))
            if trace:
                # Full snapshots are diagnostics, not O(n)-space core state.
                rows.append({'position': p, 'rank': r, 'enter': entering if entering >= 0 else None,
                             'expire': leaving if leaving >= 0 else None,
                             'active_sources_in_rank_order': [self.sa[active.kth(k)] for k in range(1, total+1)],
                             'candidates': candidates, 'length': best, 'source': source})
        return {'matches': matches, 'trace': rows}


def greedy_parse(data, window, cap):
    parameters(data, window, cap)
    index = SuffixIndex(data)
    matches = index.band(1, window, cap)['matches'] if window else [(0, None)]*len(data)
    tokens, p = [], 0
    while p < len(data):
        length, q = matches[p]
        if length >= 3:
            tokens.append(('M', length, p-q)); p += length
        else:
            tokens.append(('L', data[p])); p += 1
    return tokens


def gamma_bits(z):
    need(integer(z, 1), 'positive gamma integer')
    b = bin(z)[2:]
    return '0'*(len(b)-1)+b


def match_cost(length, distance):
    need(integer(length, 3) and integer(distance, 1), 'match cost domain')
    return 3 + 2*((length-2).bit_length()-1) + 2*(distance.bit_length()-1)


def optimal_parse(data, window, cap, trace=False):
    parameters(data, window, cap)
    need(type(trace) is bool, 'trace flag')
    n = len(data)
    index = SuffixIndex(data)
    distance_bands = []
    lo, k = 1, 0
    while lo <= min(window, n-1):
        hi = min(2*lo-1, window, n-1)
        distance_bands.append((k, lo, hi, index.band(lo, hi, cap)['matches']))
        lo *= 2; k += 1
    infinity = 9*n+1
    dp = [infinity]*n + [0]
    choices, rows = [None]*n, []
    future = MinTree(dp, infinity)
    for p in range(n-1, -1, -1):
        best, choice = 9+dp[p+1], ('L', data[p])
        for k, lo, hi, matches in distance_bands:
            M, q = matches[p]
            j, lower = 0, 3
            while lower <= M:
                upper = min((1 << (j+1))+1, M)
                remaining, end = future.minimum(p+lower, p+upper+1)
                cost = 3+2*k+2*j+remaining
                if trace:
                    rows.append({'position': p, 'distance_band': [lo, hi], 'witness': q,
                                 'max_length': M, 'length_interval': [lower, upper],
                                 'destination_range': [p+lower, p+upper+1], 'best_destination': end,
                                 'suffix_bits': remaining, 'token_bits': 3+2*k+2*j, 'total_bits': cost})
                if cost < best:
                    best, choice = cost, ('M', end-p, p-q)
                j += 1; lower = (1 << j)+2
        dp[p], choices[p] = best, choice
        future.assign(p, best)
    tokens, p = [], 0
    while p < n:
        token = choices[p]
        tokens.append(token)
        p += 1 if token[0] == 'L' else token[1]
    return {'cost': dp[0], 'dp': dp, 'tokens': tokens, 'decisions': choices,
            'distance_bands': [{'exponent': k, 'range': [lo, hi], 'matches': ms} for k, lo, hi, ms in distance_bands],
            'trace': rows}


def encode_tokens(tokens, size, window, cap):
    need(all(integer(v) for v in (size, window, cap)), 'encoding bounds')
    parts, p = [], 0
    for t in tokens:
        need(type(t) in (tuple, list), 'token shape')
        if len(t) == 2 and t[0] == 'L':
            need(integer(t[1]) and t[1] < 256 and p < size, 'literal domain/length')
            parts.append('0'+f'{t[1]:08b}'); p += 1
        else:
            need(len(t) == 3 and t[0] == 'M', 'match shape')
            length, d = t[1:]
            need(integer(length, 3) and length <= min(cap, size-p), 'match length')
            need(integer(d, 1) and d <= min(window, p), 'match distance')
            parts.append('1'+gamma_bits(length-2)+gamma_bits(d)); p += length
    need(p == size, 'token output length')
    bits = ''.join(parts)
    padded = bits + '0'*((-len(bits)) % 8)
    return {'bits': bits, 'valid_bits': len(bits), 'payload_hex': bytes(int(padded[i:i+8], 2) for i in range(0, len(padded), 8)).hex()}


def decode_payload(payload, valid_bits, size, window, cap, max_output=100000, max_bits=1000000, trace=False):
    need(type(payload) is bytes, 'payload bytes')
    need(all(integer(v) for v in (valid_bits, size, window, cap, max_output, max_bits)), 'decode bounds')
    need(type(trace) is bool, 'trace flag')
    need(size <= max_output and valid_bits <= max_bits, 'decode budget')
    need(len(payload) == (valid_bits+7)//8, 'physical payload length')
    bits = ''.join(f'{b:08b}' for b in payload)
    need('1' not in bits[valid_bits:], 'nonzero padding')
    bits = bits[:valid_bits]
    cursor, out, copies = 0, bytearray(), []
    def gamma(maximum):
        nonlocal cursor
        need(maximum >= 1, 'gamma value has empty domain')
        z = 0
        while cursor < len(bits) and bits[cursor] == '0':
            z += 1; cursor += 1
            need(z < maximum.bit_length(), 'gamma value beyond allowed bound')
        need(cursor+z < len(bits), 'truncated gamma')
        v = int(bits[cursor:cursor+z+1], 2); cursor += z+1
        need(v <= maximum, 'gamma value beyond allowed bound')
        return v
    while len(out) < size:
        need(cursor < len(bits), 'truncated token')
        flag = bits[cursor]; cursor += 1
        if flag == '0':
            need(cursor+8 <= len(bits), 'truncated literal')
            out.append(int(bits[cursor:cursor+8], 2)); cursor += 8
        else:
            length = gamma(min(cap, size-len(out))-2)+2
            d = gamma(min(window, len(out)))
            for _ in range(length):
                source = len(out)-d
                value = out[source]
                if trace:
                    copies.append({'source': source, 'destination': len(out), 'byte': value})
                out.append(value)
    need(cursor == len(bits), 'trailing valid bits')
    return bytes(out), copies


def naive_length(data, p, q, cap):
    k = 0
    while k < cap and p+k < len(data) and data[p+k] == data[q+k]:
        k += 1
    return k


def brute_cost(data, window, cap):
    # Forward shortest paths over every real distance and every legal length.
    n = len(data); cost = [10**9]*(n+1); cost[0] = 0
    for p in range(n):
        cost[p+1] = min(cost[p+1], cost[p]+9)
        for d in range(1, min(window, p)+1):
            M = naive_length(data, p, p-d, cap)
            for length in range(3, M+1):
                # Count actual independent integer-code strings rather than use match_cost.
                charge = 1+len(gamma_bits(length-2))+len(gamma_bits(d))
                cost[p+length] = min(cost[p+length], cost[p]+charge)
    return cost[n]


def run():
    counts = {'strings': 0, 'band_positions': 0, 'optimal_configurations': 0, 'random_strings': 0, 'rejections': 0}
    def verify(data, comprehensive):
        index = SuffixIndex(data); n = len(data)
        need(index.sa == sorted(range(n), key=lambda p: data[p:]), 'suffix array oracle')
        for p in range(n):
            for q in range(n):
                need(index.common_prefix(p, q) == naive_length(data, min(p,q), max(p,q), n-max(p,q)), 'LCP oracle')
        bands = [(lo, hi) for lo in range(1, n+2) for hi in range(lo, n+2)] if comprehensive else [(1, 1), (1, 5), (2, 3), (4, 9)]
        for lo, hi in bands:
            got = index.band(lo, hi, n)['matches']
            for p, (length, q) in enumerate(got):
                sources = range(max(0, p-hi), max(0, p-lo+1))
                best = max((naive_length(data, p, s, n) for s in sources), default=0)
                need(length == best, 'historical maximum')
                need((q is None) == (length == 0), 'empty witness convention')
                if q is not None:
                    need(q in sources and naive_length(data, p, q, n) == length, 'historical witness')
                counts['band_positions'] += 1
        for W in (0, 1, 3, 8):
            for cap in (2, 3, 5, 16):
                got = optimal_parse(data, W, cap)
                need(got['cost'] == brute_cost(data, W, cap), 'optimal parsing oracle')
                record = encode_tokens(got['tokens'], n, W, cap)
                need(record['valid_bits'] == got['cost'], 'exact cost encoding')
                need(decode_payload(bytes.fromhex(record['payload_hex']), record['valid_bits'], n, W, cap)[0] == data, 'optimal roundtrip')
                greedy = greedy_parse(data, W, cap); g = encode_tokens(greedy, n, W, cap)
                need(got['cost'] <= g['valid_bits'] and decode_payload(bytes.fromhex(g['payload_hex']), g['valid_bits'], n, W, cap)[0] == data, 'greedy roundtrip/order')
                counts['optimal_configurations'] += 1
    for n in range(9):
        for letters in product(b'AB', repeat=n):
            verify(bytes(letters), True); counts['strings'] += 1
    rng = random.Random(21021)
    for _ in range(80):
        verify(bytes(rng.randrange(4) for _ in range(rng.randrange(1,30))), False); counts['random_strings'] += 1
    for data in (bytes(range(256)), b'\x00\xff'*32, b'A'*128):
        index = SuffixIndex(data)
        need(index.sa == sorted(range(len(data)), key=lambda p: data[p:]), 'byte/sentinel regression')
    sample = b'AAAAAAA'; optimal = optimal_parse(sample, 8, 18, True)
    greedy = greedy_parse(sample, 8, 18); encoded = encode_tokens(optimal['tokens'], 7, 8, 18); greedy_encoded = encode_tokens(greedy, 7, 8, 18)
    need(optimal['cost'] == 15 and greedy_encoded['valid_bits'] == 16 and optimal['tokens'] == [('L',65),('M',3,1),('M',3,1)], 'seven-A cost boundary')
    tie = SuffixIndex(b'ABAABABA'); wide = tie.band(1,5,8,True); cheap = tie.band(2,3,8,True)
    need(wide['matches'][5] == (3,0) and cheap['matches'][5] == (3,3), 'two neighbors not cheapest copy')
    overlap = greedy_parse(b'ABABABABA', 2, 7); o = encode_tokens(overlap, 9, 2, 7)
    recovered, copies = decode_payload(bytes.fromhex(o['payload_hex']), o['valid_bits'], 9, 2, 7, trace=True)
    need(overlap == [('L',65),('L',66),('M',7,2)] and recovered == b'ABABABABA', 'overlap length above distance')
    def rejected(f):
        try:
            f()
        except ValueError:
            counts['rejections'] += 1
        else:
            raise RuntimeError('invalid input accepted')
    for f in (lambda: optimal_parse(b'',True,3),lambda: optimal_parse(b'',1,-1),lambda: SuffixIndex(bytearray()),lambda: tie.band(0,3,8),lambda: tie.band(3,2,8),lambda: tie.common_prefix(0,8),lambda: encode_tokens([('M',3,1)],3,8,8),lambda: encode_tokens([('L',True)],1,8,8),lambda: decode_payload(b'',0,0,1,3,max_output=-1),lambda: decode_payload(b'',0,True,1,3),lambda: decode_payload(bytes.fromhex(encoded['payload_hex']),15,7,8,18,max_output=6),lambda: decode_payload(bytes.fromhex(encoded['payload_hex']),15,7,8,18,max_bits=14),lambda: decode_payload(bytes.fromhex(encoded['payload_hex'])+b'\0',15,7,8,18),lambda: decode_payload(bytes.fromhex('20ff'),15,7,8,18),lambda: decode_payload(b'\x20',9,1,1,3)):
        rejected(f)
    # Actual malformed bit witnesses, independent of the validating token encoder.
    literal_A, literal_B = '001000001', '001000010'
    malformed = [
        ('truncated-length-gamma', literal_A+'10', 7, 8, 18, 'truncated gamma'),
        ('length-over-cap', literal_A+'1'+'011'+'1', 7, 8, 4, 'gamma value beyond allowed bound'),
        ('distance-over-window', literal_A+literal_B+'11'+'011', 5, 2, 3, 'gamma value beyond allowed bound'),
        ('match-before-output', '111', 3, 8, 3, 'gamma value has empty domain'),
        ('length-over-remaining-output', literal_A+'1'+'011'+'1', 5, 8, 18, 'gamma value beyond allowed bound'),
        ('trailing-valid-bit', literal_A+'0', 1, 8, 18, 'trailing valid bits'),
        ('truncated-distance-gamma', literal_A+literal_B+'11'+'0', 5, 2, 3, 'truncated gamma'),
        ('oversized-gamma-zero-prefix', literal_A+'1'+'000', 7, 8, 18, 'gamma value beyond allowed bound'),
    ]
    refusal_witnesses = []
    for label, bits, size, window, cap, message in malformed:
        padded = bits+'0'*((-len(bits)) % 8)
        payload = bytes(int(padded[i:i+8], 2) for i in range(0,len(padded),8))
        try:
            decode_payload(payload, len(bits), size, window, cap)
        except ValueError as error:
            need(str(error) == message, 'unexpected malformed-stream refusal')
            refusal_witnesses.append({'case':label,'bits':bits,'payload_hex':payload.hex(),
                                     'valid_bits':len(bits),'size':size,'window':window,'cap':cap,'error':str(error)})
            counts['rejections'] += 1
        else:
            raise RuntimeError('malformed bit witness accepted')
    return {'status':'PASS','counts':counts,'refusal_witnesses':refusal_witnesses,'historical_example':{'text':'ABAABABA','suffix_array':tie.sa,'lcp':tie.lcp,'wide_band':wide,'distance_2_to_3':cheap,'at_position_5_costs':{'distance5':match_cost(3,5),'distance2':match_cost(3,2)}},'seven_A':{'optimal':optimal,'optimal_payload':encoded,'greedy_tokens':greedy,'greedy_payload':greedy_encoded},'overlap':{'tokens':overlap,'payload':o,'copy_trace':copies},'scope':'Author cross-check, not independent review; payload requires external valid-bit count, size, window and cap; not DEFLATE.'}


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