#!/usr/bin/env python3
"""Exact finite-horizon rejection probabilities; no sampled or human data."""
import itertools
import json
import math
from fractions import Fraction
from pathlib import Path

HORIZONS = (20, 50, 100, 200)
PROBABILITIES = ((1, 2), (3, 5), (3, 4))
RULES = ('fixed_binomial', 'peek_binomial', 'likelihood_ratio')


def rejects(rule, n, heads, horizon):
    if rule == 'likelihood_ratio':
        # LR for alternative p=3/4 against null p=1/2.
        return 3 ** heads >= 20 * 2 ** n
    if rule == 'fixed_binomial' and n != horizon:
        return False
    # Exact one-sided binomial upper-tail p-value <= 1/20.
    return 20 * sum(math.comb(n, j) for j in range(heads, n + 1)) <= 2 ** n


def boundaries(rule, horizon):
    out = [horizon + 1]
    for n in range(1, horizon + 1):
        if rule == 'likelihood_ratio':
            out.append(next((k for k in range(n + 1) if 3 ** k >= 20 * 2 ** n), n + 1))
        elif rule == 'fixed_binomial' and n != horizon:
            out.append(n + 1)
        else:
            tail = 0
            boundary = n + 1
            for k in range(n, -1, -1):
                tail += math.comb(n, k)
                if 20 * tail <= 2 ** n:
                    boundary = k
                else:
                    break
            out.append(boundary)
    return out


def exact(rule, horizon, a, d):
    """Surviving integer path weights; head weight a, tail weight d-a."""
    boundary = boundaries(rule, horizon)
    alive = [1]
    hits = 0
    stopping_sum = 0
    for n in range(1, horizon + 1):
        nxt = [0] * (n + 1)
        hits *= d
        stopping_sum *= d
        for k, weight in enumerate(alive):
            nxt[k] += weight * (d - a)
            nxt[k + 1] += weight * a
        for k in range(boundary[n], n + 1):
            hits += nxt[k]
            stopping_sum += n * nxt[k]
            nxt[k] = 0
        alive = nxt
        assert sum(alive) + hits == d ** n
    prob = Fraction(hits, d ** horizon)
    expected_n = Fraction(stopping_sum + horizon * sum(alive), d ** horizon)
    return prob, expected_n


def brute(rule, horizon, a, d):
    """Separate oracle: enumerate full paths and directly evaluate definitions."""
    hit_weight = 0
    stop_weight = 0
    for path in itertools.product((0, 1), repeat=horizon):
        heads = 0
        stop = horizon
        hit = False
        for n, bit in enumerate(path, 1):
            heads += bit
            if rejects(rule, n, heads, horizon):
                hit = True
                stop = n
                break
        weight = a ** sum(path) * (d-a) ** (horizon-sum(path))
        hit_weight += weight * hit
        stop_weight += weight * stop
    return Fraction(hit_weight, d ** horizon), Fraction(stop_weight, d ** horizon)


def main():
    rows = []
    for horizon in HORIZONS:
        for a, d in PROBABILITIES:
            for rule in RULES:
                prob, expected = exact(rule, horizon, a, d)
                rows.append({'horizon': horizon, 'p': f'{a}/{d}', 'rule': rule,
                             'rejection_probability': float(prob),
                             'probability_exact': f'{prob.numerator}/{prob.denominator}',
                             'expected_observations': float(expected)})
    checks = 0
    for horizon in (5, 10, 16):
        for a, d in ((1, 2), (3, 5)):
            for rule in RULES:
                assert exact(rule, horizon, a, d) == brute(rule, horizon, a, d)
                checks += 1
    selected = {r['rule']: r['rejection_probability'] for r in rows if r['horizon'] == 100 and r['p'] == '1/2'}
    assert selected['fixed_binomial'] <= .05
    assert selected['likelihood_ratio'] <= .05
    assert selected['peek_binomial'] > .05
    output = {'parameters': {'alpha': '1/20', 'alternative_for_lr': '3/4',
                             'monitoring': 'after every observation from the first'},
              'rows': rows, 'null_horizon_100': selected,
              'peeking_to_fixed_ratio_horizon_100': selected['peek_binomial']/selected['fixed_binomial'],
              'oracle_checks_passed': checks, 'row_count': len(rows),
              'synthetic_only': True}
    root = Path(__file__).resolve().parents[1]
    (root/'results').mkdir(exist_ok=True)
    (root/'results'/'R1.json').write_text(json.dumps(output, indent=2)+'\n')
    print(json.dumps({'null_horizon_100': selected, 'oracle_checks_passed': checks,
                      'row_count': len(rows)}, indent=2))


if __name__ == '__main__':
    main()
