"""Exact coverage of four nominal 95% confidence intervals for a binomial proportion.

For a sample size n and a true proportion p, an interval method's coverage is the probability
that the interval it computes from X ~ Binomial(n, p) contains p:

    C(n, p) = sum over k = 0..n of [L(k) <= p <= U(k)] * P(X = k)

This is computed exactly, up to floating-point rounding, for every n from N_MIN to N_MAX and
every p on a grid of step 1/GRID, for the Wald, Wilson score, Agresti-Coull, and
Clopper-Pearson intervals. Nothing is simulated.

Only addition, subtraction, multiplication, division, and square roots are used, which IEEE 754
rounds the same way on every platform, and sums of floats go through math.fsum, which is
correctly rounded, so the results depend neither on the machine, nor on the C library's pow,
exp, or log, nor on the Python version (whose built-in sum compensates since 3.12). Writes
results/R1.json.
"""

import json
import math
import pathlib

N_MIN, N_MAX = 5, 100
GRID = 1000  # p = 1/GRID, 2/GRID, ..., (GRID - 1)/GRID
Z = 1.959963984540054  # the 0.975 quantile of the standard normal distribution
ALPHA = 0.05
DETAIL = (10, 20, 50, 100)  # sample sizes reported one by one
WALD_P = (0.01, 0.05, 0.2, 0.5)  # proportions at which Wald's coverage is followed across n
LOW = 0.93  # coverage below this counts as poor
METHODS = ("wald", "wilson", "agresti_coull", "clopper_pearson")


def pmf(n, p):
    """P(X = k) for k = 0..n, by the ratio of consecutive terms, starting from the larger tail."""
    q = 1.0 - p
    probabilities = [0.0] * (n + 1)
    if p <= 0.5:
        first = 1.0
        for _ in range(n):
            first *= q
        probabilities[0] = first
        for k in range(n):
            probabilities[k + 1] = probabilities[k] * (n - k) / (k + 1) * p / q
    else:
        last = 1.0
        for _ in range(n):
            last *= p
        probabilities[n] = last
        for k in range(n, 0, -1):
            probabilities[k - 1] = probabilities[k] * k / (n - k + 1) * q / p
    return probabilities


def wald(k, n):
    phat = k / n
    half = Z * math.sqrt(phat * (1 - phat) / n)
    return phat - half, phat + half


def wilson(k, n):
    phat = k / n
    z2 = Z * Z
    center = (phat + z2 / (2 * n)) / (1 + z2 / n)
    half = Z * math.sqrt(phat * (1 - phat) / n + z2 / (4 * n * n)) / (1 + z2 / n)
    return center - half, center + half


def agresti_coull(k, n):
    z2 = Z * Z
    ntilde = n + z2
    ptilde = (k + z2 / 2) / ntilde
    half = Z * math.sqrt(ptilde * (1 - ptilde) / ntilde)
    return ptilde - half, ptilde + half


def upper_tail(k, n, p):
    """P(X >= k)."""
    return math.fsum(pmf(n, p)[k:])


def lower_tail(k, n, p):
    """P(X <= k)."""
    return math.fsum(pmf(n, p)[: k + 1])


def bisect(f, target, increasing):
    """The p in (0, 1) where f(p) crosses target, by bisection to the limit of double precision."""
    lo, hi = 0.0, 1.0
    for _ in range(80):
        mid = (lo + hi) / 2
        if (f(mid) < target) == increasing:
            lo = mid
        else:
            hi = mid
    return (lo + hi) / 2


def clopper_pearson(k, n):
    """The interval from inverting two one-sided binomial tests, each at level ALPHA / 2."""
    lower = 0.0 if k == 0 else bisect(lambda p: upper_tail(k, n, p), ALPHA / 2, increasing=True)
    upper = 1.0 if k == n else bisect(lambda p: lower_tail(k, n, p), ALPHA / 2, increasing=False)
    return lower, upper


INTERVALS = {"wald": wald, "wilson": wilson, "agresti_coull": agresti_coull, "clopper_pearson": clopper_pearson}


def coverage_table():
    """coverage[method][n] is the list of C(n, p) over the grid."""
    table = {method: {} for method in METHODS}
    grid = [i / GRID for i in range(1, GRID)]
    for n in range(N_MIN, N_MAX + 1):
        bounds = {method: [INTERVALS[method](k, n) for k in range(n + 1)] for method in METHODS}
        for method in METHODS:
            table[method][n] = []
        for p in grid:
            probabilities = pmf(n, p)
            for method in METHODS:
                table[method][n].append(math.fsum(probabilities[k] for k, (lo, hi) in enumerate(bounds[method]) if lo <= p <= hi))
    return grid, table


def rounded(value):
    return round(value, 6)


def main():
    grid, table = coverage_table()
    sizes = range(N_MIN, N_MAX + 1)
    pairs = len(sizes) * len(grid)
    results = {
        "n_min": N_MIN,
        "n_max": N_MAX,
        "grid_points": len(grid),
        "pairs": pairs,
        "low_threshold": LOW,
        # Over every (n, p) pair: the share with coverage below LOW, and the lowest coverage.
        "share_below_low": {m: rounded(sum(c < LOW for n in sizes for c in table[m][n]) / pairs) for m in METHODS},
        "share_below_nominal": {m: rounded(sum(c < 1 - ALPHA for n in sizes for c in table[m][n]) / pairs) for m in METHODS},
        "min_coverage": {m: rounded(min(min(table[m][n]) for n in sizes)) for m in METHODS},
        "mean_coverage": {m: rounded(math.fsum(c for n in sizes for c in table[m][n]) / pairs) for m in METHODS},
        # For the sizes reported one by one: mean and lowest coverage over the grid, and the share below LOW.
        "by_n": {
            str(n): {
                m: {
                    "mean": rounded(math.fsum(table[m][n]) / len(grid)),
                    "min": rounded(min(table[m][n])),
                    "share_below_low": rounded(sum(c < LOW for c in table[m][n]) / len(grid)),
                }
                for m in METHODS
            }
            for n in DETAIL
        },
    }
    # Wald's coverage at fixed p as n grows: how often one more observation lowers it, the
    # largest such fall, and the largest n in the range at which it is still below LOW.
    wald_by_p = {}
    for p in WALD_P:
        index = round(p * GRID) - 1
        series = [table["wald"][n][index] for n in sizes]
        falls = [(series[i] - series[i + 1], N_MIN + i) for i in range(len(series) - 1)]
        size, at = max(falls)
        below = [N_MIN + i for i, c in enumerate(series) if c < LOW]
        # A result is named by its path with dots between keys, so the key spells p without one.
        wald_by_p["p_" + f"{p:g}".replace(".", "_")] = {
            "drops": sum(fall > 0 for fall, _ in falls),
            "steps": len(falls),
            "largest_drop": (
                {"from_n": at, "from": rounded(series[at - N_MIN]), "to": rounded(series[at + 1 - N_MIN]), "size": rounded(size)}
                if size > 0
                else None
            ),
            "last_n_below_low": below[-1] if below else None,
            "at_n_max": rounded(series[-1]),
        }
    results["wald_at_fixed_p"] = wald_by_p
    out = pathlib.Path(__file__).resolve().parent.parent / "results" / "R1.json"
    out.write_text(json.dumps(results, indent=2) + "\n")
    print(json.dumps(results, indent=2))


if __name__ == "__main__":
    main()
