"""Independent coverage grid using SciPy probabilities and beta quantiles."""
import json
import math
import pathlib
from fractions import Fraction

import numpy as np
from scipy.stats import beta, binom
import scipy

Z = 1.959963984540054
METHODS = ("wald", "wilson", "agresti_coull", "clopper_pearson")


def intervals(n):
    k = np.arange(n + 1, dtype=float)
    observed = k / n
    wald_radius = Z * np.sqrt(observed * (1 - observed) / n)
    adjusted_n = n + Z*Z
    center = (k + Z*Z / 2) / adjusted_n
    score_radius = Z * np.sqrt(k * (1 - observed) + Z*Z / 4) / adjusted_n
    adjusted_radius = Z * np.sqrt(center * (1 - center) / adjusted_n)
    exact_low = np.zeros(n + 1)
    exact_high = np.ones(n + 1)
    exact_low[1:] = beta.ppf(0.025, k[1:], n - k[1:] + 1)
    exact_high[:-1] = beta.ppf(0.975, k[:-1] + 1, n - k[:-1])
    return {
        "wald": (observed - wald_radius, observed + wald_radius),
        "wilson": (center - score_radius, center + score_radius),
        "agresti_coull": (center - adjusted_radius, center + adjusted_radius),
        "clopper_pearson": (exact_low, exact_high),
    }


def exact_wald_at(n, numerator):
    low, high = intervals(n)["wald"]
    p = numerator / 1000
    # Sum exact integer binomial numerators at the rational grid probability.
    total = sum(math.comb(n, k) * numerator**k * (1000 - numerator)**(n - k)
                for k in range(n + 1) if low[k] <= p <= high[k])
    return Fraction(total, 1000**n)


def main():
    grid = np.arange(1, 1000, dtype=float) / 1000
    tables = {method: [] for method in METHODS}
    normalization_error = 0.0
    for n in range(5, 101):
        probabilities = binom.pmf(np.arange(n + 1)[:, None], n, grid[None, :])
        normalization_error = max(normalization_error, float(np.max(np.abs(probabilities.sum(axis=0) - 1))))
        for method, (lower, upper) in intervals(n).items():
            includes = (lower[:, None] <= grid[None, :]) & (grid[None, :] <= upper[:, None])
            tables[method].append((probabilities * includes).sum(axis=0))
    tables = {method: np.array(rows) for method, rows in tables.items()}
    results = {
        "share_below_low": {m: round(float(np.mean(tables[m] < 0.93)), 6) for m in METHODS},
        "share_below_nominal": {m: round(float(np.mean(tables[m] < 0.95)), 6) for m in METHODS},
        "mean_coverage": {m: round(float(np.mean(tables[m])), 6) for m in METHODS},
        "min_coverage": {m: round(float(np.min(tables[m])), 6) for m in METHODS},
        "wald_at_fixed_p": {},
    }
    exact_audit = {}
    for numerator, key in [(10, "p_0_01"), (50, "p_0_05"), (200, "p_0_2"), (500, "p_0_5")]:
        exact = [exact_wald_at(n, numerator) for n in range(5, 101)]
        drops = [exact[i] - exact[i + 1] for i in range(95)]
        index = max(range(95), key=lambda i: drops[i])
        exact_drop_count = sum(drop > 0 for drop in drops)
        results["wald_at_fixed_p"][key] = {
            "at_n_max": round(float(exact[-1]), 6),
            "largest_drop": {"from_n": 5 + index, "from": round(float(exact[index]), 6), "to": round(float(exact[index + 1]), 6)} if drops[index] > 0 else None,
        }
        exact_audit[key] = {"exact_rational_drop_count": exact_drop_count, "equal_adjacent_coverages": sum(drop == 0 for drop in drops)}
        assert abs(float(exact[-1]) - tables["wald"][-1, numerator - 1]) < 1e-12
    audit = {
        "numpy": np.__version__, "scipy": scipy.__version__,
        "grid_pairs": 96 * 999,
        "maximum_probability_normalization_error": normalization_error,
        "clopper_pearson_all_grid_coverages_above_0_95": bool(np.all(tables["clopper_pearson"] >= 0.95)),
        "exact_rational_fixed_p_audit": exact_audit,
    }
    assert normalization_error < 1e-12
    pathlib.Path("results/independent.json").write_text(json.dumps(results, indent=2) + "\n")
    pathlib.Path("results/independent-audit.json").write_text(json.dumps(audit, indent=2) + "\n")


if __name__ == "__main__":
    main()
