#!/usr/bin/env python3
"""Finite checks accompanying five technical notes by Kylian de Groot, 2026-09-05.

Run with Python 3; only the standard library is needed.
These finite checks support examples and bounded domains, not universal proofs.
The universal claims and their assumptions are proved in the notes themselves.
"""
from fractions import Fraction as F
from itertools import combinations, product
from math import prod
import json


def reconciliation():
    amounts = [100, 400, 200, 300]
    signature = lambda x: (sum(x), sum(a * b for a, b in zip(amounts, x)))
    vectors = list(product([0, 1], repeat=4))
    assert signature((1, 1, 0, 0)) == signature((0, 0, 1, 1)) == (2, 500)
    collisions = [(x, y) for x, y in combinations(vectors, 2) if signature(x) == signature(y)]
    witnesses = [z for z in product([-1, 0, 1], repeat=4) if any(z) and signature(z) == (0, 0)]
    assert {tuple(a - b for a, b in zip(x, y)) for x, y in collisions}.issubset(set(witnesses))
    for z in witnesses:
        x = tuple(int(v == 1) for v in z)
        y = tuple(int(v == -1) for v in z)
        assert x != y and signature(x) == signature(y)
    inclusion_checks = 0
    for e in vectors:
        for local in vectors:
            if all(l <= s for l, s in zip(local, e)):
                if sum(local) == sum(e):
                    assert local == e
                inclusion_checks += 1
    encoded = {sum(bit * 2 ** i for i, bit in enumerate(x)) for x in vectors}
    assert len(encoded) == len(vectors)
    return {'binary_candidates': len(vectors), 'collision_pairs': len(collisions),
            'kernel_witnesses': len(witnesses), 'inclusion_cases': inclusion_checks}


def idempotency():
    def transition(d, j, key, amount):
        d, j = dict(d), dict(j)
        if key not in d:
            d[key], j[key] = amount, amount
        elif d[key] != amount:
            return d, j, 'conflict'
        return d, j, 'ok'
    requests = [('a', 1250), ('a', 1400), ('b', -300)]
    sequences = 0
    for length in range(8):
        for sequence in product(requests, repeat=length):
            d, j = {}, {}
            for key, amount in sequence:
                d, j, _ = transition(d, j, key, amount)
                assert d.keys() == j.keys() and all(j[k] == d[k] for k in d)
                repeated = transition(d, j, key, amount)
                assert repeated[:2] == (d, j)
            sequences += 1
    def run(sequence):
        d, j = {}, {}
        for key, amount in sequence:
            d, j, _ = transition(d, j, key, amount)
        return d, j
    assert run([('a', 1250), ('b', -300)]) == run([('b', -300), ('a', 1250)])
    d, j = run([('a', 1250), ('a', 1250), ('b', -300), ('a', 1400)])
    assert len(d) == 2 and sum(j.values()) == 950
    # Deliberately non-atomic effect-first schedule: effect, crash, retry.
    effects, seen = [1250], set()
    if 'a' not in seen:
        effects.append(1250)
        seen.add('a')
    assert sum(effects) == 2500
    # Deliberately non-atomic key-first schedule: key, crash, suppress retry.
    effects, seen = [], {'a'}
    if 'a' not in seen:
        effects.append(1250)
    assert sum(effects) == 0
    return {'delivery_sequences': sequences, 'example_total_cents': 950,
            'non_atomic_counterexamples': 2}


def allocation():
    def allocate(total, weights):
        denominator = sum(weights)
        qr = [divmod(total * w, denominator) for w in weights]
        floors = [q for q, _ in qr]
        remaining = total - sum(floors)
        selected = set(sorted(range(len(weights)), key=lambda i: (-qr[i][1], i))[:remaining])
        return [q + int(i in selected) for i, q in enumerate(floors)]
    cases = quota_candidates = 0
    for n in range(1, 5):
        for weights in product(range(4), repeat=n):
            if not sum(weights):
                continue
            for total in range(16):
                actual = allocate(total, weights)
                ideals = [F(total * w, sum(weights)) for w in weights]
                floors = [q.numerator // q.denominator for q in ideals]
                fractional = [i for i, q in enumerate(ideals) if q.denominator != 1]
                remaining = total - sum(floors)
                assert sum(actual) == total and all(a >= 0 for a in actual)
                assert all(abs(a - q) < 1 for a, q in zip(actual, ideals))
                best_abs = sum(abs(a - q) for a, q in zip(actual, ideals))
                best_square = sum((a - q) ** 2 for a, q in zip(actual, ideals))
                for selected in combinations(fractional, remaining):
                    candidate = [f + int(i in selected) for i, f in enumerate(floors)]
                    assert best_abs <= sum(abs(a - q) for a, q in zip(candidate, ideals))
                    assert best_square <= sum((a - q) ** 2 for a, q in zip(candidate, ideals))
                    quota_candidates += 1
                cases += 1
    assert allocate(1001, [5, 3, 2]) == [501, 300, 200]
    assert allocate(4, [5, 3, 1]) == [2, 1, 1]
    assert allocate(5, [5, 3, 1]) == [3, 2, 0]
    return {'weight_and_total_cases': cases, 'quota_candidates': quota_candidates,
            'example': [501, 300, 200], 'monotonicity_counterexample': [[2, 1, 1], [3, 2, 0]]}


def retries():
    probability_cases = collision_cases = 0
    for k in range(1, 9):
        for p in [F(0), F(1, 5), F(1, 2), F(1)]:
            expectation, success = F(0), F(0)
            for failures in product([False, True], repeat=k):
                probability = prod(p if failed else 1 - p for failed in failures)
                attempts = next((i + 1 for i, failed in enumerate(failures) if not failed), k)
                expectation += probability * attempts
                if not all(failures):
                    success += probability
            assert success == 1 - p ** k
            assert expectation == sum(p ** j for j in range(k))
            if p != 1:
                assert expectation == (1 - p ** k) / (1 - p)
            probability_cases += 1
    for n in range(1, 6):
        for m in range(1, 6):
            collision_sum = collided = 0
            for slots in product(range(m), repeat=n):
                pairs = sum(slots[i] == slots[j] for i in range(n) for j in range(i + 1, n))
                collision_sum += pairs
                collided += pairs > 0
            expectation = F(collision_sum, m ** n)
            probability = F(collided, m ** n)
            assert expectation == F(n * (n - 1), 2 * m)
            assert probability <= min(1, expectation)
            exact_no_collision = prod(F(m - j, m) for j in range(n)) if n <= m else 0
            assert 1 - probability == exact_no_collision
            collision_cases += 1
    assert 4 * 500 + 100 * (2 ** (4 - 1) - 1) == 2700
    assert sum(F(1, 5) ** j for j in range(4)) == F(156, 125)
    assert 1 - F(1, 5) ** 4 == F(624, 625)
    assert 3 ** 5 == 243
    return {'independent_attempt_cases': probability_cases, 'slot_models': collision_cases,
            'max_example_milliseconds': 2700, 'expected_attempts': '156/125',
            'example_success_probability': '624/625'}


def synthetic_data():
    tv = lambda p, q: sum(abs(a - b) for a, b in zip(p, q)) / 2
    p = [F(1, 2), F(0), F(0), F(1, 2)]
    q = [F(0), F(1, 2), F(1, 2), F(0)]
    assert tv(p, q) == 1
    assert p[0] + p[1] == q[0] + q[1]
    assert p[0] + p[2] == q[0] + q[2]
    tables = 0
    for n in range(1, 13):
        for c00 in range(n + 1):
            for c01 in range(n - c00 + 1):
                for c10 in range(n - c00 - c01 + 1):
                    c11 = n - c00 - c01 - c10
                    a, b, t = F(c10 + c11, n), F(c01 + c11, n), F(c11, n)
                    assert max(0, a + b - 1) <= t <= min(a, b)
                    assert [1 - a - b + t, b - t, a - t, t] == [F(c, n) for c in [c00, c01, c10, c11]]
                    tables += 1
    # All distributions on three outcomes with denominator four and every indicator loss.
    distributions = [[F(a, 4), F(b, 4), F(4 - a - b, 4)]
                     for a in range(5) for b in range(5 - a)]
    loss_checks = conditional_checks = 0
    for p in distributions:
        for q in distributions:
            for indicator in product([0, 1], repeat=3):
                assert abs(sum(f * (a - b) for f, a, b in zip(indicator, p, q))) <= tv(p, q)
                loss_checks += 1
            for mask in product([0, 1], repeat=3):
                alpha = sum(a for a, include in zip(p, mask) if include)
                beta = sum(b for b, include in zip(q, mask) if include)
                if alpha and alpha == beta:
                    pc = [a / alpha for a, include in zip(p, mask) if include]
                    qc = [b / alpha for b, include in zip(q, mask) if include]
                    assert tv(pc, qc) <= min(1, tv(p, q) / alpha)
                    conditional_checks += 1
    p, q = [F(1, 100), F(0), F(99, 100)], [F(0), F(1, 100), F(99, 100)]
    assert tv(p, q) == F(1, 100)
    assert tv([F(1), F(0)], [F(0), F(1)]) == 1
    return {'binary_tables': tables, 'bounded_indicator_checks': loss_checks,
            'equal_mass_conditional_checks': conditional_checks,
            'rare_group_global_tv': '1/100', 'rare_group_conditional_tv': 1}


if __name__ == '__main__':
    checks = {name: fn() for name, fn in [
        ('reconciliation', reconciliation), ('idempotency', idempotency),
        ('allocation', allocation), ('retries', retries), ('synthetic_data', synthetic_data)]}
    print(json.dumps({'status': 'PASS', 'scope': 'Finite exact-arithmetic checks; not formal verification of the universal proofs.', 'checks': checks}, indent=2))
