#!/usr/bin/env python3
"""Exact byte-class Shift-And and Myers-1999 streaming scores, using word arrays.
Python 3.10+. No packages. All assertions used for verification are explicit checks.
"""
import itertools
import json
import random

MAX_PATTERN = 4096
MAX_POSITION = (1 << 63) - 1


def integer(x, lo, hi, name):
    if type(x) is not int or not lo <= x <= hi:
        raise ValueError(name)
    return x


def byte_string(x, name):
    if type(x) is not bytes:
        raise ValueError(name)
    return x


def geometry(m, width):
    integer(width, 1, 63, 'word width')
    q = (m + width - 1) // width
    masks = [(1 << width) - 1] * q
    if q:
        masks[-1] = (1 << (m - (q - 1) * width)) - 1
    return tuple(masks)


def shift_one(words, width, masks, low):
    """Old words are never overwritten before their outgoing bit is read."""
    out = []
    carry = low
    for value, mask in zip(words, masks):
        out.append(((value << 1) | carry) & mask)
        carry = value >> (width - 1)
    return out


def packed_step(pv, mv, eq, width, masks, beta):
    """All operands have m valid bits, stored little-endian in q words."""
    xv, xh, ph, mh = [], [], [], []
    addition_carries = []
    carry = 0
    full = (1 << width) - 1
    for p, v, e, mask in zip(pv, mv, eq, masks):
        total = (e & p) + p + carry
        carry = total >> width
        addition_carries.append(carry)
        h = (((total & full) ^ p) | e) & mask
        xh.append(h)
        xv.append(e | v)
        ph.append((v | ~(h | p)) & mask)
        mh.append(p & h)
    top = (masks[-1] + 1) >> 1
    delta = int(bool(ph[-1] & top)) - int(bool(mh[-1] & top))
    ps = shift_one(ph, width, masks, beta)
    ms = shift_one(mh, width, masks, 0)
    new_p = [(b | ~(x | a)) & mask
             for x, a, b, mask in zip(xv, ps, ms, masks)]
    new_m = [a & x for a, x in zip(ps, xv)]
    detail = {'eq': list(eq), 'xv': xv, 'xh': xh, 'ph': ph, 'mh': mh,
              'add_carry_out': addition_carries, 'ph_shift': ps,
              'mh_shift': ms, 'pv': new_p[:], 'mv': new_m[:], 'delta': delta}
    return new_p, new_m, delta, detail


class ShiftAnd:
    """Each pattern position is a bytes object listing its allowed bytes.
    Empty classes reject every byte. Empty pattern matches all boundaries.
    initial_end is reported separately once, never by an empty feed call.
    """
    def __init__(self, classes, width=63):
        if type(classes) not in (tuple, list) or len(classes) > MAX_PATTERN:
            raise ValueError('class sequence')
        for c in classes:
            byte_string(c, 'class bytes')
        self.m, self.width = len(classes), width
        self.masks = geometry(self.m, width)
        table = [[0] * len(self.masks) for _ in range(256)]
        for i, allowed in enumerate(classes):
            for c in allowed:
                table[c][i // width] |= 1 << (i % width)
        self.table = tuple(tuple(row) for row in table)
        self.state = [0] * len(self.masks)
        self.position = 0
        self.initial_end = 0 if not self.m else None

    def step(self, c):
        integer(c, 0, 255, 'byte')
        if self.position == MAX_POSITION:
            raise ValueError('position overflow')
        if self.m:
            shifted = shift_one(self.state, self.width, self.masks, 1)
            self.state = [a & b for a, b in zip(shifted, self.table[c])]
        self.position += 1
        hit = not self.m or bool(self.state[-1] & ((self.masks[-1] + 1) >> 1))
        return self.position if hit else None

    def feed(self, data):
        byte_string(data, 'text chunk')
        if len(data) > MAX_POSITION - self.position:
            raise ValueError('position overflow')
        out = []
        for c in data:
            end = self.step(c)
            if end is not None:
                out.append(end)
        return out


class Myers:
    """mode='global': d(P,T[:j]); mode='infix': min_s d(P,T[s:j]).
    The empty suffix s=j is allowed. Neither mode returns a start or a script.
    The initial j=0 score is self.score; feed emits only newly read positions.
    """
    def __init__(self, pattern, mode='infix', width=63):
        byte_string(pattern, 'pattern bytes')
        if len(pattern) > MAX_PATTERN:
            raise ValueError('pattern length')
        if type(mode) is not str or mode not in ('global', 'infix'):
            raise ValueError('boundary mode')
        self.m, self.width = len(pattern), width
        self.beta = int(mode == 'global')
        self.masks = geometry(self.m, width)
        table = [[0] * len(self.masks) for _ in range(256)]
        for i, c in enumerate(pattern):
            table[c][i // width] |= 1 << (i % width)
        self.table = tuple(tuple(row) for row in table)
        self.pv, self.mv = list(self.masks), [0] * len(self.masks)
        self.position, self.score = 0, self.m
        self.last = None

    def step(self, c):
        integer(c, 0, 255, 'byte')
        if self.position == MAX_POSITION:
            raise ValueError('position overflow')
        if self.m:
            p, v, delta, detail = packed_step(
                self.pv, self.mv, self.table[c], self.width, self.masks, self.beta)
            self.pv, self.mv = p, v
            self.score += delta
            self.last = detail  # One column only; no accumulated trace.
        else:
            self.score += self.beta
        self.position += 1
        return self.score

    def feed(self, data):
        byte_string(data, 'text chunk')
        if len(data) > MAX_POSITION - self.position:
            raise ValueError('position overflow')
        return [self.step(c) for c in data]

    def column(self):
        """Diagnostic O(m) reconstruction; never called by step/feed."""
        out = [self.beta * self.position]
        for i in range(self.m):
            word, bit = i // self.width, 1 << (i % self.width)
            out.append(out[-1] + bool(self.pv[word] & bit) - bool(self.mv[word] & bit))
        return out


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


def scalar_columns(pattern, text, mode):
    col = list(range(len(pattern) + 1))
    out = [col]
    for j, c in enumerate(text, 1):
        nxt = [j if mode == 'global' else 0]
        for i, a in enumerate(pattern, 1):
            nxt.append(min(nxt[-1] + 1, col[i] + 1, col[i - 1] + (a != c)))
        col = nxt
        out.append(col)
    return out


def literal_ends(classes, text):
    m = len(classes)
    return [j for j in range(m, len(text) + 1)
            if all(text[j-m+i] in classes[i] for i in range(m))]


def strings(length, alphabet=b'ab'):
    return [bytes(x) for n in range(length + 1) for x in itertools.product(alphabet, repeat=n)]


def verify_column(obj, expected):
    require(obj.column() == expected and obj.score == expected[-1], 'full DP column')
    for p, v, mask in zip(obj.pv, obj.mv, obj.masks):
        require(p & v == 0 and (p | v) & ~mask == 0, 'disjoint valid delta masks')


def main():
    counts = {'class_cases': 0, 'class_states': 0, 'edit_cases': 0,
              'edit_columns': 0, 'random_cases': 0, 'chunk_checks': 0,
              'cell_inputs': 0, 'carry_inputs': 0, 'refusals': 0}
    # Exhaust all local cell inputs, and prove arithmetic propagation finitely.
    for v, h, e in itertools.product((-1, 0, 1), (-1, 0, 1), (0, 1)):
        base = min(-e, v, h) + 1
        xv, xh = bool(e or v == -1), bool(e or h == -1)
        gotv = int(h == -1 or not (xv or h == 1)) - int(h == 1 and xv)
        goth = int(v == -1 or not (xh or v == 1)) - int(v == 1 and xh)
        require((gotv, goth) == (base-h, base-v), '18-input cell logic')
        counts['cell_inputs'] += 1
    for bits in range(1, 7):
        mask = (1 << bits) - 1
        for p, e in itertools.product(range(1 << bits), repeat=2):
            carry = 0
            expected = 0
            for i in range(bits):
                expected |= (((e >> i) & 1) | carry) << i
                carry = ((p >> i) & 1) & (((e >> i) & 1) | carry)
            got = ((((e & p) + p) ^ p) | e) & mask
            require(got == expected, 'carry recurrence')
            counts['carry_inputs'] += 1
    for m in range(5):
        for classes in itertools.product((b'', b'a', b'b', b'ab'), repeat=m):
            for text in strings(5):
                expected = literal_ends(classes, text)
                for w in (1, 2, 3):
                    a = ShiftAnd(classes, w)
                    got = [a.initial_end] if a.initial_end is not None else []
                    for j, c in enumerate(text, 1):
                        end = a.step(c)
                        if end is not None:
                            got.append(end)
                        active = sum(1 << (i-1) for i in range(1, m+1)
                                     if j >= i and all(text[j-i+k] in classes[k] for k in range(i)))
                        packed = sum(x << (k*w) for k, x in enumerate(a.state))
                        require(packed == active, 'every active prefix')
                        counts['class_states'] += 1
                    require(got == expected, 'all class endpoints')
                    counts['class_cases'] += 1
    for pattern, text in itertools.product(strings(5), repeat=2):
        for mode in ('global', 'infix'):
            columns = scalar_columns(pattern, text, mode)
            for w in (1, 2, 3, 5, 8, 63):
                a = Myers(pattern, mode, w)
                verify_column(a, columns[0])
                for j, c in enumerate(text, 1):
                    a.step(c)
                    verify_column(a, columns[j])
                    counts['edit_columns'] += 1
                counts['edit_cases'] += 1
    rng = random.Random(23001999)
    for trial in range(600):
        m = rng.choice((0, 1, 7, 8, 9, 62, 63, 64, 65, 126, 127, 128, 129))
        pattern = bytes(rng.choice(b'abc') for _ in range(m))
        text = bytes(rng.choice(b'abcd') for _ in range(rng.randrange(160)))
        w = rng.choice((1, 2, 3, 5, 8, 31, 63))
        classes = tuple(bytes(c for c in b'abcd' if rng.randrange(2)) for _ in pattern)
        sa = ShiftAnd(classes, w)
        ends = ([0] if not m else []) + sa.feed(text)
        require(ends == literal_ends(classes, text), 'long class matching')
        for mode in ('global', 'infix'):
            a = Myers(pattern, mode, w)
            columns = scalar_columns(pattern, text, mode)
            for j, c in enumerate(text, 1):
                a.step(c)
                verify_column(a, columns[j])
                counts['edit_columns'] += 1
            b, c = Myers(pattern, mode, w), ShiftAnd(classes, w)
            scores, split_ends = [], [0] if not m else []
            pos = 0
            while pos < len(text):
                require(b.feed(b'') == [] and c.feed(b'') == [], 'empty chunk')
                take = rng.randrange(1, 14)
                scores.extend(b.feed(text[pos:pos+take]))
                split_ends.extend(c.feed(text[pos:pos+take]))
                pos += take
            require(scores == [x[-1] for x in columns[1:]], 'split scores')
            require(split_ends == ends and c.state == sa.state, 'split endpoints')
            counts['chunk_checks'] += 1
        counts['random_cases'] += 1
    def refuse(call):
        try:
            call()
        except ValueError:
            counts['refusals'] += 1
        else:
            raise RuntimeError('missing refusal')
    for width in (0, 64, True, 2.0):
        refuse(lambda: ShiftAnd((b'a',), width))
        refuse(lambda: Myers(b'a', width=width))
    for pattern in ('ab', bytearray(b'ab'), [97], b'a'*(MAX_PATTERN+1)):
        refuse(lambda: Myers(pattern))
    for classes in ('ab', (b'a', 'b'), (bytearray(b'a'),), [b'a']*(MAX_PATTERN+1)):
        refuse(lambda: ShiftAnd(classes))
    for mode in ('local', '', 0):
        refuse(lambda: Myers(b'a', mode))
    for a in (ShiftAnd((b'a', b'b'), 1), Myers(b'ab', width=1)):
        for value in (-1, 256, True, 1.0, b'a'):
            old = repr(a.__dict__)
            refuse(lambda: a.step(value))
            require(repr(a.__dict__) == old, 'invalid byte atomicity')
        for chunk in ('ab', bytearray(b'ab'), [97, 999]):
            old = repr(a.__dict__)
            refuse(lambda: a.feed(chunk))
            require(repr(a.__dict__) == old, 'invalid chunk atomicity')
        a.position = MAX_POSITION
        old = repr(a.__dict__)
        refuse(lambda: a.step(97))
        refuse(lambda: a.feed(b'a'))
        require(a.feed(b'') == [] and repr(a.__dict__) == old, 'overflow atomicity')
    classes = (b'a', b'ab', b'b', bytes(c for c in range(256) if c != 98), b'a')
    text = b'aabxaabya'
    sa = ShiftAnd(classes, 3)
    class_trace = []
    for c in text:
        end = sa.step(c)
        class_trace.append({'end': sa.position, 'char': chr(c), 'words': sa.state[:], 'match': end})
    edit_trace = {}
    for mode in ('global', 'infix'):
        a = Myers(b'ababa', mode, 3)
        rows = [{'end': 0, 'score': a.score, 'column': a.column(), 'pv': a.pv[:], 'mv': a.mv[:]}]
        for c in b'zzababa':
            a.step(c)
            rows.append({'end': a.position, 'char': chr(c), 'score': a.score,
                         'column': a.column(), 'detail': a.last})
        edit_trace[mode] = rows
    # Minimal, directly observable failures of four tempting shortcuts.
    mutants = {
        'no_restart': {'pattern': 'a', 'text': 'a', 'wrong_endpoints': [], 'correct_endpoints': [1]},
        'no_word_shift_carry': {'pattern': 'aaaa', 'text': 'aaaa', 'width': 3,
                               'wrong_endpoints': [], 'correct_endpoints': [4]},
        'no_addition_carry': {'pattern': 'abbb', 'text': 'a', 'width': 3,
                              'wrong_score': 4, 'correct_score': 3},
        'wrong_top_boundary': {'pattern': 'a', 'text': 'ba', 'global_score': 1, 'infix_score': 0}}
    # Compute, rather than merely print, the mutant witnesses.
    require(((0 << 1) & 1) == 0, 'missing restart witness')
    bad = [0, 0]
    for _ in range(4):
        bad = [(x << 1 | (1 if i == 0 else 0)) & mask for i, (x, mask) in enumerate(zip(bad, (7, 1)))]
    require(bad == [7, 0], 'lost shift carry witness')
    # Use a match only in the first limb to force the genuine carry into limb two.
    a = Myers(b'abbb', 'infix', 3)
    good = a.step(ord('a'))
    bad_h = (((0 & 1) + 1) ^ 1) | 0
    wrong = 4 + int(bool((~(bad_h | 1) & 1))) - int(bool(1 & bad_h))
    require((good, wrong) == (3, 4), 'lost addition carry witness')
    require(Myers(b'a', 'global').feed(b'ba')[-1] == 1 and
            Myers(b'a', 'infix').feed(b'ba')[-1] == 0, 'boundary witness')
    # A compatible class border is not a valid border of the actual text.
    shortcut_classes = (b'a', b'ab')
    require(bool(set(shortcut_classes[0]) & set(shortcut_classes[1])), 'compatible border')
    wrong_q, wrong_ends = 0, []
    for j, c in enumerate(b'aba', 1):
        if c in shortcut_classes[wrong_q]:
            wrong_q += 1
        if wrong_q == 2:
            wrong_ends.append(j)
            wrong_q = 1  # Incorrect failure link built from nonempty intersection.
    require(wrong_ends == [2, 3] and literal_ends(shortcut_classes, b'aba') == [2], 'false compatible border')
    mutants['class_compatibility_border'] = {'pattern': 'a[ab]', 'text': 'aba',
        'wrong_endpoints': wrong_ends, 'correct_endpoints': [2]}
    print(json.dumps({'status': 'PASS', 'counts': counts, 'class_masks_width3': {
        chr(c): list(sa.table[c]) for c in b'abx'}, 'class_trace': class_trace,
        'edit_trace': edit_trace, 'mutant_witnesses': mutants}, ensure_ascii=False, indent=2, sort_keys=True))


if __name__ == '__main__':
    main()
