#!/usr/bin/env python3
"""Exact finite A22 reference. Supplied integer planes are NOT Gaussian samples.
Run normally or with -O; checks deliberately do not use Python assert.
"""
from dataclasses import dataclass
from fractions import Fraction
from itertools import product, permutations
from math import comb
from random import Random
from types import MappingProxyType
import json


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


def integer(x, low=None):
    need(type(x) is int and (low is None or x >= low), 'integer domain')
    return x


def vector(x, d):
    need(type(x) in (tuple, list) and len(x) == d, 'vector dimension')
    y = tuple(integer(t) for t in x)
    need(any(y), 'zero direction')
    return y


def threshold(a, b):
    integer(a, 0); integer(b, 1)
    need(a <= b, 'threshold outside [0,1]')


def cosine_at_least(x, y, a, b):
    """Exact threshold in [0,1]; sign must be checked before squaring."""
    threshold(a, b)
    x = vector(x, len(y)); y = vector(y, len(x))
    dot = sum(s*t for s, t in zip(x, y))
    return (dot >= 0 and b*b*dot*dot >=
            a*a*sum(s*s for s in x)*sum(t*t for t in y))


@dataclass(frozen=True, eq=False)
class Hyperplanes:
    planes: tuple

    def __post_init__(self):
        need(type(self.planes) in (tuple, list) and len(self.planes) > 0,
             'empty plane family')
        first = self.planes[0]
        need(type(first) in (tuple, list) and len(first) > 0, 'plane dimension')
        ps = tuple(vector(g, len(first)) for g in self.planes)
        object.__setattr__(self, 'planes', ps)

    @property
    def d(self):
        return len(self.planes[0])

    @property
    def q(self):
        return len(self.planes)

    def prepare(self, x):
        return vector(x, self.d)

    def values(self, x):
        x = self.prepare(x)
        return tuple(int(sum(a*b for a, b in zip(g, x)) >= 0)
                     for g in self.planes)

    def encode(self, x):
        bits = self.values(x)
        data = bytearray((self.q+7)//8)
        for i, bit in enumerate(bits):
            data[i//8] |= bit << (7-i % 8)
        return Fingerprint(self, bytes(data))

    def accepts(self, x, y, a, b):
        return cosine_at_least(x, y, a, b)


@dataclass(frozen=True)
class Fingerprint:
    family: Hyperplanes
    data: bytes

    def __post_init__(self):
        need(type(self.family) is Hyperplanes, 'fingerprint family')
        need(type(self.data) is bytes and len(self.data) == (self.family.q+7)//8,
             'fingerprint byte length')
        unused = (-self.family.q) % 8
        need(not unused or not (self.data[-1] & ((1 << unused)-1)),
             'noncanonical padding')

    def bits(self):
        return tuple((self.data[i//8] >> (7-i % 8)) & 1
                     for i in range(self.family.q))


def mismatch(a, b):
    need(type(a) is Fingerprint and type(b) is Fingerprint, 'fingerprint type')
    need(a.family is b.family, 'different shared family')
    return sum((s ^ t).bit_count() for s, t in zip(a.data, b.data))


def angle_fraction(a, b):
    return Fraction(mismatch(a, b), a.family.q)


@dataclass(frozen=True, eq=False)
class Minwise:
    size: int
    orders: tuple

    def __post_init__(self):
        integer(self.size, 1)
        need(type(self.orders) in (tuple, list) and len(self.orders) > 0,
             'empty permutation family')
        rows = []
        for order in self.orders:
            need(type(order) in (tuple, list), 'permutation shape')
            o = tuple(integer(t, 0) for t in order)
            need(len(o) == self.size and set(o) == set(range(self.size)),
                 'not a universe permutation')
            rows.append(o)
        object.__setattr__(self, 'orders', tuple(rows))
        ranks = []
        for o in rows:
            rank = [0]*self.size
            for i, item in enumerate(o):
                rank[item] = i
            ranks.append(tuple(rank))
        object.__setattr__(self, '_ranks', tuple(ranks))

    @property
    def q(self):
        return len(self.orders)

    def prepare(self, x):
        need(type(x) in (tuple, list, set, frozenset), 'set input shape')
        y = frozenset(integer(t, 0) for t in x)
        need(bool(y) and all(t < self.size for t in y), 'nonempty universe subset')
        return y

    def values(self, x):
        x = self.prepare(x)
        return tuple(min(x, key=rank.__getitem__) for rank in self._ranks)

    def accepts(self, x, y, a, b):
        threshold(a, b)
        return b*len(x & y) >= a*len(x | y)


class StaticLSH:
    """Read-only tables. Distinct IDs are records; equal values do not coalesce."""
    def __init__(self, family, k, records):
        need(type(family) in (Hyperplanes, Minwise), 'unsupported family')
        integer(k, 1)
        need(family.q % k == 0, 'nonrectangular table family')
        need(type(records) in (tuple, list), 'record list')
        values, rows = {}, []
        tables = [{} for _ in range(family.q//k)]
        for row in records:
            need(type(row) in (tuple, list) and len(row) == 2, 'record shape')
            rid, raw = row
            need(type(rid) is str and bool(rid) and rid not in values,
                 'duplicate or invalid record ID')
            value = family.prepare(raw)
            values[rid] = value; rows.append((rid, value))
            sig = family.values(value)
            for j, table in enumerate(tables):
                key = tuple(sig[j*k:(j+1)*k])
                table.setdefault(key, []).append(rid)
        self.family, self.k, self.L = family, k, family.q//k
        self.records = tuple(rows)
        self.values = MappingProxyType(values)
        self.tables = tuple(MappingProxyType({k: tuple(v) for k, v in t.items()})
                            for t in tables)

    def query(self, raw, a, b, budget=None):
        threshold(a, b)
        if budget is not None:
            integer(budget, 0)
        value = self.family.prepare(raw)
        sig = self.family.values(value)
        keys = tuple(tuple(sig[j*self.k:(j+1)*self.k]) for j in range(self.L))
        buckets = tuple(t.get(key, ()) for t, key in zip(self.tables, keys))
        total = sum(map(len, buckets))
        limit = total if budget is None else min(total, budget)
        seen, checked, matches = set(), [], []
        scanned = 0
        for bucket in buckets:
            for rid in bucket:
                if scanned == limit:
                    break
                scanned += 1
                if rid not in seen:
                    seen.add(rid); checked.append(rid)
                    if self.family.accepts(self.values[rid], value, a, b):
                        matches.append(rid)
            if scanned == limit:
                break
        return dict(keys=keys, checked_ids=checked, matched_ids=matches,
                    occurrences_available=total, occurrences_scanned=scanned,
                    duplicate_occurrences=scanned-len(checked),
                    status='buckets_exhausted' if scanned == total else 'budget_exhausted')


def amplified(p, k, L):
    need(type(p) is Fraction and 0 <= p <= 1, 'rational probability')
    integer(k, 1); integer(L, 1)
    return 1-(1-p**k)**L


def rational(x):
    return f'{x.numerator}/{x.denominator}'


def check(ok, label):
    if not ok:
        raise RuntimeError(label)


def run():
    counts = dict(sign_cases=0, packed_roundtrips=0, exact_probabilities=0,
                  index_queries=0, rejected_inputs=0)
    plane = Hyperplanes(((1, 0), (0, 1), (1, 1), (1, -1), (-2, 3)))
    vectors = [('A', (2, 1)), ('B', (1, 2)), ('minus_A', (-2, -1)), ('tie', (1, 1))]
    trace = {name: dict(vector=x, bits=plane.encode(x).bits(),
                       hex=plane.encode(x).data.hex()) for name, x in vectors}
    check(trace['A']['hex'] == 'f0' and trace['B']['hex'] == 'e8', 'five-bit trace')
    check(angle_fraction(plane.encode((2, 1)), plane.encode((1, 2))) == Fraction(2, 5),
          'angle trace')
    # Exact supplied integer arithmetic and canonical byte encodings.
    rng = Random(2201)
    for q in range(1, 34):
        for d in range(1, 5):
            ps = []
            while len(ps) < q:
                g = tuple(rng.randrange(-3, 4) for _ in range(d))
                if any(g): ps.append(g)
            f = Hyperplanes(ps)
            for _ in range(12):
                x = tuple(rng.randrange(-4, 5) for _ in range(d))
                if not any(x): x = (1,)+(0,)*(d-1)
                bits = tuple(int(sum(g[j]*x[j] for j in range(d)) >= 0) for g in ps)
                fp = f.encode(x)
                check(fp.bits() == bits and f.encode(tuple(7*t for t in x)) == fp,
                      'sign / positive scale')
                check(Fingerprint(f, fp.data) == fp, 'canonical roundtrip')
                counts['sign_cases'] += q; counts['packed_roundtrips'] += 1
    # Exact finite product experiment: this verifies combinatorics, not Gaussian sampling.
    quadrant = ((1, 1), (-1, 1), (-1, -1), (1, -1))
    histogram = [0]*6
    for ps in product(quadrant, repeat=5):
        f = Hyperplanes(ps)
        histogram[mismatch(f.encode((1, 0)), f.encode((0, 1)))] += 1
    check(histogram == [32*comb(5, j) for j in range(6)], 'binomial five bits')
    # Enumerate all 6^4 independent MinHash assignments and all 7^2 nonempty-set pairs.
    orders = tuple(permutations(range(3)))
    sets = tuple(frozenset(i for i in range(3) if mask >> i & 1) for mask in range(1, 8))
    signatures = [tuple(Minwise(3, (p,)).values(s)[0] for s in sets) for p in orders]
    totals = [[0]*7 for _ in sets]
    shared_components, repeated_tables = 0, 0
    ai, bi = sets.index(frozenset((0, 1))), sets.index(frozenset((1, 2)))
    for a, b, c, d in product(range(6), repeat=4):
        for i in range(7):
            for j in range(7):
                first = signatures[a][i] == signatures[a][j] and signatures[b][i] == signatures[b][j]
                second = signatures[c][i] == signatures[c][j] and signatures[d][i] == signatures[d][j]
                totals[i][j] += first or second
        shared_components += signatures[a][ai] == signatures[a][bi]
        repeated_tables += (signatures[a][ai] == signatures[a][bi] and
                            signatures[b][ai] == signatures[b][bi])
    for i, x in enumerate(sets):
        for j, y in enumerate(sets):
            p = Fraction(len(x & y), len(x | y))
            check(Fraction(totals[i][j], 1296) == amplified(p, 2, 2), 'exact AND/OR')
            counts['exact_probabilities'] += 1
    check(totals[ai][bi] == 272 and shared_components == 432 and repeated_tables == 144,
          'independence failures')
    f = Hyperplanes(((1, 0), (0, 1), (1, 1), (1, -1)))
    records = [('F', (1, 8)), ('A', (2, 1)), ('A2', (4, 2)), ('B', (1, -1)), ('C', (-3, 1))]
    idx = StaticLSH(f, 2, records)
    queries = {str(b): idx.query((3, 1), 4, 5, b) for b in (0, 1, 2, 3, 4, 6)}
    queries['full'] = idx.query((3, 1), 4, 5)
    check(queries['full']['matched_ids'] == ['A', 'A2'], 'exact duplicate-valued records')
    check(queries['full']['occurrences_available'] == 6 and
          queries['full']['duplicate_occurrences'] == 2, 'posting count')
    check(queries['1']['matched_ids'] == [] and queries['1']['status'] == 'budget_exhausted',
          'budget false negative')
    # Literal full-key oracle, independent traversal order and exact rational verification.
    for _ in range(350):
        family = (Hyperplanes(tuple(rng.choice(quadrant) for _ in range(6))) if _ % 2 else
                  Minwise(3, tuple(rng.choice(orders) for _ in range(6))))
        options = list(sets) if type(family) is Minwise else [(1, 0), (0, 1), (-1, 0), (1, 2), (2, 1), (-1, -1)]
        rows = [(str(i), rng.choice(options)) for i in range(rng.randrange(9))]
        index = StaticLSH(family, 2, rows)
        for query in options:
            sigq = family.values(query)
            occurrences = []
            for table in range(3):
                for rid, x in rows:
                    sx = family.values(x)
                    if all(sx[2*table+j] == sigq[2*table+j] for j in range(2)):
                        occurrences.append(rid)
            for budget in (None, 0, 1, 2, 5, len(occurrences), len(occurrences)+1):
                result = index.query(query, 2, 3, budget)
                visited = occurrences if budget is None else occurrences[:budget]
                unique = list(dict.fromkeys(visited)); expected = []
                for rid in unique:
                    x = dict(rows)[rid]
                    if type(family) is Minwise:
                        good = Fraction(len(x & query), len(x | query)) >= Fraction(2, 3)
                    else:
                        dot = sum(x[i]*query[i] for i in range(2))
                        good = dot >= 0 and Fraction(dot*dot, sum(t*t for t in x)*sum(t*t for t in query)) >= Fraction(4, 9)
                    if good: expected.append(rid)
                check(result['checked_ids'] == unique and result['matched_ids'] == expected,
                      'literal candidate oracle')
                check(result['occurrences_available'] == len(occurrences) and
                      result['occurrences_scanned'] == len(visited) and
                      result['duplicate_occurrences'] == len(visited)-len(unique), 'scan accounting')
                check((result['status'] == 'buckets_exhausted') == (len(visited) == len(occurrences)),
                      'exhaustion status')
                counts['index_queries'] += 1
    # Nonzero opposite directions chosen in a known nullspace have identical fingerprints.
    adaptive = Hyperplanes(((1, 0, 0), (0, 1, 0)))
    check(mismatch(adaptive.encode((0, 0, 1)), adaptive.encode((0, 0, -1))) == 0,
          'adaptive nullspace')
    refusal = [
        ('zero direction', lambda: plane.encode((0, 0))),
        ('wrong dimension', lambda: plane.encode((1, 2, 3))),
        ('float coordinate', lambda: plane.encode((1.0, 2))),
        ('bool coordinate', lambda: plane.encode((True, 2))),
        ('empty planes', lambda: Hyperplanes(())),
        ('zero plane', lambda: Hyperplanes(((0, 0),))),
        ('ragged planes', lambda: Hyperplanes(((1, 0), (1,)))),
        ('short payload', lambda: Fingerprint(plane, b'')),
        ('long payload', lambda: Fingerprint(plane, b'\xf0\x00')),
        ('nonzero padding', lambda: Fingerprint(plane, b'\xf1')),
        ('mutable payload', lambda: Fingerprint(plane, bytearray(b'\xf0'))),
        ('different families', lambda: mismatch(plane.encode((2, 1)), Hyperplanes(plane.planes).encode((2, 1)))),
        ('zero k', lambda: StaticLSH(f, 0, [])),
        ('unequal table lengths', lambda: StaticLSH(plane, 2, [])),
        ('duplicate IDs', lambda: StaticLSH(f, 2, [('x', (1, 0)), ('x', (0, 1))])),
        ('empty ID', lambda: StaticLSH(f, 2, [('', (1, 0))])),
        ('negative budget', lambda: idx.query((3, 1), 4, 5, -1)),
        ('bool budget', lambda: idx.query((3, 1), 4, 5, True)),
        ('large threshold', lambda: idx.query((3, 1), 6, 5)),
        ('zero denominator', lambda: idx.query((3, 1), 0, 0)),
        ('duplicate permutation entry', lambda: Minwise(3, ((0, 0, 2),))),
        ('empty set', lambda: Minwise(3, (orders[0],)).values(set())),
        ('out-of-universe', lambda: Minwise(3, (orders[0],)).values({3})),
    ]
    rejects = []
    for label, call in refusal:
        try: call()
        except ValueError as error:
            rejects.append(dict(case=label, error=str(error)))
        else: raise RuntimeError('accepted invalid input: '+label)
    counts['rejected_inputs'] = len(rejects)
    return dict(status='PASS', counts=counts, plane_trace=trace,
                trace_normalized_angle_estimate='2/5', binomial_histogram=histogram,
                probability=dict(base='1/3', k=2, L=2, configurations=1296,
                                 colliding=272, rate='17/81', shared_components='1/3', repeated_tables='1/9'),
                candidate_trace=queries, refusal_witnesses=rejects,
                caveat='Exact integer fixtures and finite permutation product; no ideal Gaussian sampler is implemented.')


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