#!/usr/bin/env python3
"""Finite teaching contracts, not a PostgreSQL/Trino implementation.
Only standard library; exact rational predictions; stdout only; checks survive -O.
"""
from bisect import bisect_right
from collections import Counter
from dataclasses import dataclass
from fractions import Fraction as F
from itertools import product, permutations
from random import Random
import json


def require(ok, message):
    if not ok:
        raise AssertionError(message)


def integer(x):
    return isinstance(x, int) and not isinstance(x, bool)


class Histogram:
    def __init__(self, values, boundaries, top_m, snapshot):
        boundaries = tuple(boundaries)
        if (len(boundaries) < 2 or not all(integer(x) for x in boundaries)
                or any(a >= b for a, b in zip(boundaries, boundaries[1:]))
                or not integer(top_m) or top_m < 0):
            raise ValueError('strict integer boundaries and nonnegative MCV count required')
        self.boundaries, self.snapshot = boundaries, snapshot
        freq = Counter()
        self.n = self.nulls = 0
        for x in values:
            if x is None:
                self.nulls += 1
            elif not integer(x) or not boundaries[0] <= x < boundaries[-1]:
                raise ValueError('value outside the declared finite integer universe')
            else:
                freq[x] += 1
            self.n += 1
        ranked = sorted(freq, key=lambda x: (-freq[x], x))
        self.mcv = {x: freq[x] for x in ranked[:top_m]}
        counts = [0] * (len(boundaries)-1)
        slots = [b-a for a, b in zip(boundaries, boundaries[1:])]
        mcv_entries = []
        for x, count in freq.items():
            j = bisect_right(boundaries, x)-1
            if x in self.mcv:
                slots[j] -= 1
                mcv_entries.append((x, count, j))
            else:
                counts[j] += count
        self.counts, self.slots = tuple(counts), tuple(slots)
        self.mcv_entries = tuple(mcv_entries)
        require(self.nulls + sum(self.mcv.values()) + sum(counts) == self.n, 'mass partition')
        require(all(l > 0 or c == 0 for l, c in zip(slots, counts)), 'no mass in absent slots')

    def signature(self):
        return (self.snapshot, self.boundaries, self.n, self.nulls,
                tuple(sorted(self.mcv.items())), self.counts, self.slots)

    def _context(self, snapshot):
        if snapshot != self.snapshot:
            raise ValueError('summary belongs to a different snapshot')

    def estimate_lt(self, threshold, snapshot):
        self._context(snapshot)
        if not integer(threshold):
            raise ValueError('integer threshold required')
        taken = [max(0, min(threshold, b)-a)
                 for a, b in zip(self.boundaries, self.boundaries[1:])]
        exact_mcv = 0
        for x, count, j in self.mcv_entries:
            if x < threshold:
                taken[j] -= 1
                exact_mcv += count
        estimated = F(exact_mcv)
        lower = upper = exact_mcv
        for c, size, selected in zip(self.counts, self.slots, taken):
            require(0 <= selected <= size, 'selected residual positions')
            if size:
                estimated += F(c*selected, size)
            if selected == size:
                lower += c
                upper += c
            elif selected > 0:
                upper += c
        return dict(estimate=estimated, lower=lower, upper=upper,
                    selectivity=None if self.n == 0 else estimated/self.n,
                    residual_positions=taken)

    def estimate_eq(self, x, snapshot):
        self._context(snapshot)
        if x is None:
            return F(0)  # Ordinary SQL '=' with NULL is never TRUE.
        if not integer(x):
            raise ValueError('integer equality key required')
        if x in self.mcv:
            return F(self.mcv[x])
        if not self.boundaries[0] <= x < self.boundaries[-1]:
            return F(0)
        j = bisect_right(self.boundaries, x)-1
        require(self.slots[j] > 0, 'a non-MCV key has one residual slot')
        return F(self.counts[j], self.slots[j])


@dataclass(frozen=True)
class HashConfig:
    m: int = 16
    seed: int = 0
    version: str = 'integer-affine-pair-v1'

    def __post_init__(self):
        if not integer(self.m) or self.m <= 0 or not integer(self.seed) or self.version != 'integer-affine-pair-v1':
            raise ValueError('unsupported hash configuration')

    def locations(self, key):
        if not integer(key):
            raise ValueError('non-NULL integer key required')
        return ((key+self.seed) % self.m, (5*key+1+self.seed) % self.m)


def bloom_bits(keys, config):
    bits = bytearray(config.m)
    for key in keys:
        if key is not None:
            for i in config.locations(key):
                bits[i] = 1
    return bytes(bits)


class ProtocolError(RuntimeError):
    pass


class RuntimeFilter:
    def __init__(self, token, shards, config=HashConfig()):
        self.token = tuple(token)
        self.shards = frozenset(shards)
        if len(self.token) != 4:
            raise ValueError('query, execution epoch, snapshot, join ID required')
        self.config = config
        self.received = {}
        self.merged = bytearray(config.m)
        self.state = 'WAITING' if self.shards else 'READY'

    def receive(self, token, shard, config, final_bits):
        if self.state == 'FAILED':
            raise ProtocolError('failed execution cannot continue')
        if self.state == 'TOP':
            return 'ignored_after_timeout'
        if tuple(token) != self.token:
            return 'ignored_execution'
        if shard not in self.shards:
            return 'ignored_shard'
        if config != self.config:
            return 'ignored_configuration'
        if (not isinstance(final_bits, (bytes, bytearray)) or len(final_bits) != config.m
                or any(x not in (0, 1) for x in final_bits)):
            self.state = 'FAILED'
            raise ProtocolError('malformed final bitmap')
        bits = bytes(final_bits)
        if shard in self.received:
            if bits != self.received[shard]:
                self.state = 'FAILED'
                raise ProtocolError('conflicting immutable final summary')
            return 'duplicate'
        self.received[shard] = bits
        for i in range(config.m):
            self.merged[i] |= bits[i]
        if len(self.received) == len(self.shards):
            self.state = 'READY'
        return self.state

    def timeout(self):
        if self.state == 'FAILED':
            raise ProtocolError('failure is not a timeout fallback')
        if self.state == 'WAITING':
            self.state = 'TOP'
        return self.state

    def may_match(self, key):
        if self.state == 'FAILED':
            raise ProtocolError('no output from failed query')
        if key is None:
            return False  # Safe here only for the declared ordinary inner equality join.
        locations = self.config.locations(key)
        if self.state in ('WAITING', 'TOP'):
            return True
        return all(self.merged[i] for i in locations)


class Checkpoint:
    def __init__(self, token, hash_slots=8, startup=48, switch=6):
        if not all(integer(x) and x >= 0 for x in (hash_slots, startup, switch)):
            raise ValueError('nonnegative capacity and declared work charges required')
        self.token = tuple(token)
        if len(self.token) != 4:
            raise ValueError('four-field execution token required')
        self.capacity, self.startup, self.switch = hash_slots, startup, switch
        self.rows = {'B': [], 'P': []}
        self.ids = {'B': set(), 'P': set()}
        self.closed = set()
        self.state = 'COLLECTING'
        self.choice = None
        self.delivered = []
        self.stats = dict(comparisons=0, key_inputs=0, outputs=0, hash_rows=0)

    def append(self, side, row, token):
        if self.state != 'COLLECTING' or side not in self.rows or side in self.closed:
            raise ValueError('input side is not open')
        if tuple(token) != self.token:
            raise ValueError('input from another execution/snapshot')
        if len(row) != 2 or not isinstance(row[0], str) or not row[0]:
            raise ValueError('nonempty stable string ID and key required')
        identity, key = row
        if key is not None and not integer(key):
            raise ValueError('integer/NULL key required')
        if side == 'P' and key is None:
            raise ValueError('probe NULLs must already be removed for this checkpoint')
        if identity in self.ids[side]:
            raise ValueError('duplicate occurrence identity, not duplicate value')
        self.rows[side].append((identity, key))
        self.ids[side].add(identity)

    def close(self, side):
        if self.state != 'COLLECTING' or side not in self.rows or side in self.closed:
            raise ValueError('invalid END transition')
        self.rows[side] = tuple(self.rows[side])
        self.closed.add(side)
        if len(self.closed) == 2:
            self.state = 'FROZEN'

    def fail_input(self):
        if self.state != 'COLLECTING':
            raise ValueError('input failure outside collection')
        self.state = 'FAILED'

    def choose(self, token=None, columns=('bid', 'pid'), order='unordered-bag'):
        if self.state != 'FROZEN':
            raise ValueError('only an unconsumed frozen checkpoint can choose once')
        if token is not None and tuple(token) != self.token:
            raise ValueError('candidate snapshot/execution mismatch')
        if columns != ('bid', 'pid') or order != 'unordered-bag':
            raise ValueError('candidate output contract mismatch')
        n, q = len(self.rows['B']), len(self.rows['P'])
        costs = {'NL': n*q, 'HP': self.startup+n+q, 'HB': self.startup+n+q}
        feasible = {'NL': True, 'HP': q <= self.capacity, 'HB': n <= self.capacity}
        future = {name: cost+(0 if name == 'NL' else self.switch)
                  for name, cost in costs.items() if feasible[name]}
        rank = {'NL': 0, 'HP': 1, 'HB': 2}
        selected = min(future, key=lambda name: (future[name], rank[name]))
        self.choice = dict(plan=selected, switched=selected != 'NL',
                           input_rows=(n, q), costs=costs, feasible=feasible, future=future,
                           chosen_future=future[selected])
        self.state = 'READY'
        return self.choice.copy()

    def _produce(self):
        B, P = self.rows['B'], self.rows['P']
        plan = self.choice['plan']
        if plan == 'NL':
            for bid, bk in B:
                for pid, pk in P:
                    self.stats['comparisons'] += 1
                    if bk == pk:
                        yield (bid, pid)
        else:
            build, scan = (P, B) if plan == 'HP' else (B, P)
            table = {}
            for identity, key in build:
                self.stats['key_inputs'] += 1
                table.setdefault(key, []).append(identity)
                self.stats['hash_rows'] += 1
            require(self.stats['hash_rows'] <= self.capacity, 'actual hash row slots')
            for identity, key in scan:
                self.stats['key_inputs'] += 1
                for match in table.get(key, ()):
                    yield (identity, match) if plan == 'HP' else (match, identity)

    def next_pair(self):
        if self.state == 'READY':
            self.iterator = self._produce()
            self.state = 'RUNNING'
        if self.state == 'DONE':
            return None
        if self.state != 'RUNNING':
            raise ValueError('suffix is not runnable')
        try:
            item = next(self.iterator)
        except StopIteration:
            self.state = 'DONE'
            return None
        except Exception:
            self.state = 'FAILED'
            raise
        self.delivered.append(item)
        self.stats['outputs'] += 1
        return item

    def finish(self):
        while self.next_pair() is not None:
            pass
        return tuple(self.delivered)


def direct_join(B, P):
    return tuple((bid, pid) for bid, bk in B for pid, pk in P
                 if bk is not None and pk is not None and bk == pk)


def frozen(B, P, token, **options):
    checkpoint = Checkpoint(token, **options)
    for side, rows in [('B', B), ('P', P)]:
        for row in rows:
            checkpoint.append(side, row, token)
        checkpoint.close(side)
    return checkpoint


def rejects(fn, kind=ValueError):
    try:
        fn()
    except kind:
        return
    raise AssertionError('required rejection did not occur')


def demo():
    token = ('q', 2, 'sigma7', 'J')
    values = [None, None]+[8]*10+[0, 1, 4, 5]+[6]*4+[7]*4
    summary = Histogram(values, (0, 6, 12), 1, 'sigma7')
    prediction = summary.estimate_lt(8, 'sigma7')
    alternative = [None, None]+[8]*10+[0, 1, 4, 5]+[9]*4+[10]*4
    require(summary.signature() == Histogram(alternative, (0, 6, 12), 1, 'sigma7').signature(), 'same summary')
    B = [(f'r{i}', x) for i, x in enumerate(values) if x is not None and x < 8]
    P = [(f'p{i}', x) for i, x in enumerate([0, 2, 6, 6, 7, 7, 10, 16, None])]
    config = HashConfig()
    shard_bits = [bloom_bits([x for _, x in B[:8]], config), bloom_bits([x for _, x in B[8:]], config)]
    receiver = RuntimeFilter(token, ('s0', 's1'), config)
    events = []
    for t, shard, cfg, bits in [(token, 's0', config, shard_bits[0]),
                                (token, 's0', config, shard_bits[0]),
                                (('q', 1, 'sigma7', 'J'), 's1', config, shard_bits[1]),
                                (token, 's1', HashConfig(32), bloom_bits([7], HashConfig(32))),
                                (token, 's1', config, shard_bits[1])]:
        answer = receiver.receive(t, shard, cfg, bits)
        events.append(dict(result=answer, state=receiver.state, completed=len(receiver.received),
                           key7_passes=receiver.may_match(7)))
    filtered = [row for row in P if receiver.may_match(row[1])]
    partial = [row for row in P if row[1] is not None and all(shard_bits[0][i] for i in config.locations(row[1]))]
    wrong_and = [row for row in P if row[1] is not None and all(shard_bits[0][i] & shard_bits[1][i] for i in config.locations(row[1]))]
    timeout = RuntimeFilter(token, ('s0', 's1'), config)
    timeout.receive(token, 's0', config, shard_bits[0]);timeout.timeout()
    late = timeout.receive(token, 's1', config, shard_bits[1])
    unfiltered = [row for row in P if timeout.may_match(row[1])]
    paths = []
    for name, probe in [('complete_filter', filtered), ('timeout_top', unfiltered)]:
        checkpoint = frozen(B, probe, token)
        decision = checkpoint.choose()
        pairs = checkpoint.finish()
        require(Counter(pairs) == Counter(direct_join(B, P)), 'full output occurrence bag')
        paths.append(dict(path=name, probe=probe, decision=decision, stats=checkpoint.stats,
                          final_state=checkpoint.state, pairs=pairs))
    require([p['decision']['plan'] for p in paths] == ['NL', 'HP'], 'two genuine decisions')
    require(len(direct_join(B, P)) == 17 and len(direct_join(B, partial)) == 9, 'partial-filter loss')
    migration_fee = frozen(B, filtered, token, switch=5)
    migration_small = frozen(B, filtered, token, hash_slots=5)
    fee_choice, small_choice = migration_fee.choose(), migration_small.choose()
    require(fee_choice['plan'] == 'HP' and small_choice['plan'] == 'NL', 'fee/capacity migrations')
    require(Counter(migration_fee.finish()) == Counter(direct_join(B,P)), 'migrated fee still exact')
    require(Counter(migration_small.finish()) == Counter(direct_join(B,P)), 'migrated capacity still exact')
    prefix_run = frozen(B, filtered, token)
    prefix_run.choose();first_pair = prefix_run.next_pair();rejects(prefix_run.choose)
    prefix_run.finish()
    unsafe_restart = (first_pair,)+direct_join(B,P)
    foreign_probe = P+[('new-snapshot-p9',7)]
    require(len(unsafe_restart) == 18 and len(direct_join(B,foreign_probe)) == 21, 'unsafe restart/snapshot certificates')
    no_mcv = Histogram(values,(0,6,12),0,'sigma7')
    return dict(histogram=dict(signature=summary.signature(),prediction=prediction, actual=12,
                               same_summary_alternative=4, equality8=summary.estimate_eq(8,'sigma7'),
                               equality9=summary.estimate_eq(9,'sigma7')),
                filter=dict(shard_set_bits=[[i for i,v in enumerate(bits) if v] for bits in shard_bits],
                            final_set_bits=[i for i,v in enumerate(receiver.merged) if v], events=events,
                            partial_pairs=len(direct_join(B, partial)), and_pairs=len(direct_join(B, wrong_and)),
                            timeout_late=late), paths=paths, migrations=dict(fee5=fee_choice, capacity5=small_choice, no_mcv_counts=no_mcv.counts, no_mcv_prediction=no_mcv.estimate_lt(8,'sigma7'), denied_running_switch=True, safe_prefix_total=len(prefix_run.delivered), unsafe_restarted_total=len(unsafe_restart), mixed_snapshot_total=len(direct_join(B,foreign_probe))))


def main():
    counts = Counter()
    # Enumerate all multiplicities 0,1,2 for NULL and six integer values.
    for multiplicities in product(range(3), repeat=7):
        values = [x for x, count in zip([None]+list(range(6)), multiplicities) for _ in range(count)]
        for top in (0, 1, 6):
            summary = Histogram(values, (0, 3, 6), top, 's')
            for threshold in range(-1, 8):
                answer = summary.estimate_lt(threshold, 's')
                truth = sum(x is not None and x < threshold for x in values)
                require(answer['lower'] <= truth <= answer['upper'], 'true cardinality interval')
                require(answer['lower'] <= answer['estimate'] <= answer['upper'], 'estimate inside interval')
                if threshold in (0, 3, 6):
                    require(answer['estimate'] == truth, 'bucket boundary exactness')
                counts['histogram_thresholds'] += 1
            for x in summary.mcv:
                require(summary.estimate_eq(x, 's') == values.count(x), 'exact MCV equality')
            counts['histogram_summaries'] += 1
    token = ('q', 1, 's', 'j');config = HashConfig(7)
    for assignment in product(range(4), repeat=5):
        shards = [[key for key, side in enumerate(assignment) if side == j+1] for j in range(3)]
        true_keys = {key for part in shards for key in part}
        bitmaps = [bloom_bits(part, config) for part in shards]
        for order in permutations(range(3)):
            rf = RuntimeFilter(token, range(3), config)
            for i in order:
                rf.receive(token, i, config, bitmaps[i]);rf.receive(token, i, config, bitmaps[i])
                require(all(rf.may_match(key) for key in true_keys), 'never lose a build key')
                if rf.state != 'READY':
                    require(all(rf.may_match(key) for key in range(11)), 'partial state is TOP for non-NULL')
                counts['filter_message_steps'] += 2
            require(rf.state == 'READY', 'all expected empty/nonempty shards completed')
            counts['filter_delivery_orders'] += 1
    rng = Random(120091)
    for _ in range(1400):
        B = [(f'b{i}', rng.choice([None]+list(range(6)))) for i in range(rng.randrange(14))]
        P = [(f'p{i}', rng.randrange(6)) for i in range(rng.randrange(12))]
        cp = frozen(B, P, token, hash_slots=rng.randrange(15), startup=rng.randrange(60), switch=rng.randrange(11))
        decision = cp.choose();pairs = cp.finish()
        require(Counter(pairs) == Counter(direct_join(B, P)), 'all three suffixes match occurrence product')
        require(decision['chosen_future'] == min(decision['future'].values()), 'minimum declared future cost')
        if decision['plan'] == 'NL':
            require(cp.stats['comparisons'] == len(B)*len(P), 'literal pair comparison count')
        else:
            require(cp.stats['key_inputs'] == len(B)+len(P), 'hash input-work count')
        counts['checkpoint_'+decision['plan']] += 1
    empty = RuntimeFilter(token, (), config)
    require(not any(empty.may_match(k) for k in range(10)), 'empty build is not unknown build')
    invalid = RuntimeFilter(token, ('s',), config);bits = bloom_bits([1],config)
    invalid.receive(token, 's', config, bits)
    rejects(lambda: invalid.receive(token, 's', config, bytes(config.m)), ProtocolError)
    rejects(lambda: invalid.may_match(1), ProtocolError)
    cp = Checkpoint(token)
    rejects(cp.choose);cp.append('B', ('b', 7), token);cp.close('B')
    rejects(lambda: cp.append('B', ('later',7),token))
    rejects(lambda: cp.append('P', ('p',7),('q',1,'wrong','j')))
    cp.append('P', ('p',7),token);rejects(lambda: cp.append('P',('p',7),token));cp.close('P')
    rejects(lambda: cp.choose(order='sorted'));rejects(lambda: cp.choose(token=('q',1,'wrong','j')))
    cp.choose();require(cp.next_pair()==('b','p'),'first output');rejects(cp.choose)
    cp.finish();require(len(cp.delivered)==1,'no duplicate after denied switch')
    failed = Checkpoint(token);failed.fail_input();rejects(lambda: failed.close('B'));rejects(failed.choose)
    rejects(lambda: Histogram([0],(0,0,1),1,'s'))
    rejects(lambda: Histogram([0],(0,1),1,'s').estimate_lt(1,'other'))
    result=dict(status='PASS', counts=dict(counts), demo=demo(),
                limits='Finite teaching tests. No probabilistic independence claim for demo hashes; no live database or distributed failure implementation.')
    print(json.dumps(jsonable(result),ensure_ascii=False,indent=2))


def jsonable(x):
    if isinstance(x,F):
        return str(x)
    if isinstance(x,dict):
        return {str(k):jsonable(v) for k,v in x.items()}
    if isinstance(x,(list,tuple)):
        return [jsonable(v) for v in x]
    return x


if __name__ == '__main__':
    main()
