"""Exact finite example of selection on reused evaluation data.

This module deliberately models a narrow case: K independent candidate models
each receive a noise-only accuracy count on the same *kind* of evaluation, and
the candidate with the largest count is selected.  The counts themselves are
assumed independent Binomial(n, p0) variables.  No model is trained here.
"""

from __future__ import annotations

import json
import math
from pathlib import Path
from typing import Any, Iterable


DEFAULT_N = 64
DEFAULT_P0 = 0.5
DEFAULT_ALPHA = 0.05
DEFAULT_K_VALUES = (1, 2, 5, 10, 20, 50, 100, 200)
DEFAULT_OUTPUT = Path(__file__).with_name("data") / "adaptive.json"


def _validate_binomial_parameters(n: int, p: float) -> None:
    if not isinstance(n, int) or isinstance(n, bool) or n < 0:
        raise ValueError("n must be a nonnegative integer")
    if not 0.0 <= p <= 1.0:
        raise ValueError("p must lie in [0, 1]")


def _validate_k(k: int) -> None:
    if not isinstance(k, int) or isinstance(k, bool) or k < 1:
        raise ValueError("k must be a positive integer")


def binomial_pmf(x: int, n: int, p: float) -> float:
    """Return P(X=x) for X ~ Binomial(n, p), using only the stdlib."""
    _validate_binomial_parameters(n, p)
    if not isinstance(x, int) or isinstance(x, bool) or x < 0 or x > n:
        return 0.0
    if p == 0.0:
        return 1.0 if x == 0 else 0.0
    if p == 1.0:
        return 1.0 if x == n else 0.0
    return math.comb(n, x) * (p**x) * ((1.0 - p) ** (n - x))


def binomial_cdf(x: int, n: int, p: float) -> float:
    """Return P(X<=x) for X ~ Binomial(n, p)."""
    _validate_binomial_parameters(n, p)
    if x < 0:
        return 0.0
    if x >= n:
        return 1.0
    return math.fsum(binomial_pmf(j, n, p) for j in range(x + 1))


def binomial_tail(x: int, n: int, p: float) -> float:
    """Return the inclusive one-sided tail P(X>=x)."""
    _validate_binomial_parameters(n, p)
    if x <= 0:
        return 1.0
    if x > n:
        return 0.0
    return math.fsum(binomial_pmf(j, n, p) for j in range(x, n + 1))


def critical_count(n: int, p0: float, cutoff: float) -> int:
    """Smallest count whose exact one-sided null tail is at most cutoff.

    Returns n + 1 if the cutoff is smaller than every attainable positive tail.
    """
    _validate_binomial_parameters(n, p0)
    if not 0.0 <= cutoff <= 1.0:
        raise ValueError("cutoff must lie in [0, 1]")
    for x in range(n + 1):
        if binomial_tail(x, n, p0) <= cutoff:
            return x
    return n + 1


def selected_false_accept_rate(
    *, n: int, p0: float, k: int, critical: int
) -> float:
    """Return P(max(X_1,...,X_K) >= critical) under independent nulls."""
    _validate_binomial_parameters(n, p0)
    _validate_k(k)
    if critical <= 0:
        return 1.0
    if critical > n:
        return 0.0
    single_tail = binomial_tail(critical, n, p0)
    if single_tail == 1.0:
        return 1.0
    return -math.expm1(k * math.log1p(-single_tail))


def expected_max_count(*, n: int, p0: float, k: int) -> float:
    """Return E[max(X_1,...,X_K)] using the finite survival identity."""
    _validate_binomial_parameters(n, p0)
    _validate_k(k)
    survival_terms = []
    for count in range(1, n + 1):
        cdf_below = binomial_cdf(count - 1, n, p0)
        survival_terms.append(1.0 - cdf_below**k)
    return math.fsum(survival_terms)


def experiment_row(
    k: int,
    *,
    n: int = DEFAULT_N,
    p0: float = DEFAULT_P0,
    alpha: float = DEFAULT_ALPHA,
) -> dict[str, int | float]:
    """Compute all exact finite quantities shown for one candidate count K."""
    _validate_k(k)
    if n < 1:
        raise ValueError("n must be positive for accuracy calculations")
    nominal_critical = critical_count(n, p0, alpha)
    nominal_attainable = binomial_tail(nominal_critical, n, p0)
    corrected_cutoff = alpha / k
    corrected_critical = critical_count(n, p0, corrected_cutoff)
    corrected_attainable = binomial_tail(corrected_critical, n, p0)
    expected_reused_count = expected_max_count(n=n, p0=p0, k=k)

    return {
        "k": k,
        "nominal_critical_count": nominal_critical,
        "nominal_critical_accuracy": nominal_critical / n,
        "nominal_single_candidate_attainable_rate": nominal_attainable,
        "reused_selected_false_accept_rate": selected_false_accept_rate(
            n=n, p0=p0, k=k, critical=nominal_critical
        ),
        "bonferroni_cutoff": corrected_cutoff,
        "bonferroni_critical_count": corrected_critical,
        "bonferroni_critical_accuracy": corrected_critical / n,
        "bonferroni_single_candidate_attainable_rate": corrected_attainable,
        "bonferroni_familywise_false_accept_rate": selected_false_accept_rate(
            n=n, p0=p0, k=k, critical=corrected_critical
        ),
        "expected_selected_reused_count": expected_reused_count,
        "expected_selected_reused_accuracy": expected_reused_count / n,
        "fresh_critical_count": nominal_critical,
        "fresh_critical_accuracy": nominal_critical / n,
        "fresh_selected_false_accept_rate": nominal_attainable,
        "expected_fresh_selected_accuracy": p0,
    }


def primary_source() -> dict:
    """Identify the primary publication motivating the broader question."""
    return {
        "title": "The reusable holdout: Preserving validity in adaptive data analysis",
        "authors": [
            "Cynthia Dwork",
            "Vitaly Feldman",
            "Moritz Hardt",
            "Toniann Pitassi",
            "Omer Reingold",
            "Aaron Roth",
        ],
        "publication": "Science",
        "year": 2015,
        "doi": "10.1126/science.aaa9375",
        "url": "https://research.ibm.com/publications/the-reusable-holdout-preserving-validity-in-adaptive-data-analysis",
    }


def build_payload(
    *,
    n: int = DEFAULT_N,
    p0: float = DEFAULT_P0,
    alpha: float = DEFAULT_ALPHA,
    k_values: Iterable[int] = DEFAULT_K_VALUES,
) -> dict[str, Any]:
    """Build the deterministic plotting payload."""
    k_values = tuple(k_values)
    if not k_values:
        raise ValueError("k_values must not be empty")
    rows = [experiment_row(k, n=n, p0=p0, alpha=alpha) for k in k_values]
    return {
        "schema_version": 1,
        "model": "independent_binomial_null_winner_selection",
        "source": primary_source(),
        "parameters": {
            "n_evaluation_items": n,
            "null_accuracy": p0,
            "alpha": alpha,
            "k_values": list(k_values),
            "selection_rule": "Select one candidate with the largest reused-evaluation count; ties may be broken arbitrarily.",
            "reused_test_rule": "Apply an exact inclusive upper-tail Binomial(n, p0) p-value to the selected reused count.",
            "bonferroni_rule": "Apply the same exact test at alpha/K before selecting any passing winner.",
            "fresh_test_rule": "Evaluate the already selected candidate once on a new independent Binomial(n, p0) count, with no re-selection, at alpha.",
        },
        "definitions": {
            "reused_selected_false_accept_rate": "P(max_i X_i reaches the nominal exact-binomial critical count) under K independent null candidates.",
            "bonferroni_familywise_false_accept_rate": "P(max_i X_i reaches the exact-binomial critical count for alpha/K) under K independent null candidates.",
            "fresh_selected_false_accept_rate": "P(Y_I reaches the nominal critical count), where I is selected from reused counts and all fresh Y_i are independent of that selection.",
            "expected_selected_reused_accuracy": "E[max_i X_i / n] on the reused evaluation.",
            "expected_fresh_selected_accuracy": "E[Y_I / n] on the untouched evaluation; under these assumptions it equals p0.",
        },
        "assumptions": [
            "Every candidate is null: its true accuracy is p0.",
            "Candidate reused-evaluation counts X_i are mutually independent Binomial(n, p0) variables.",
            "Fresh counts Y_i are mutually independent Binomial(n, p0) variables and independent of every reused count X_i.",
            "The candidate index is selected only from the reused counts, never from fresh results.",
            "This is a finite model-selection special case, not a model of general sequential adaptivity.",
            "This calculation does not train an AI system, run a human experiment, or reimplement the private reusable-holdout mechanism.",
            "The guarantees shown rely on valid exact p-values and the stated independence; they are not asserted for dependent candidates or misspecified nulls.",
        ],
        "rows": rows,
    }


def write_payload(path: Path = DEFAULT_OUTPUT) -> Path:
    """Write the deterministic JSON payload and return its path."""
    payload = build_payload()
    path.parent.mkdir(parents=True, exist_ok=True)
    path.write_text(json.dumps(payload, indent=2) + "\n", encoding="utf-8")
    return path


if __name__ == "__main__":
    print(write_payload())
