#!/usr/bin/env python3
"""Reproduce every numerical claim in the Bayesian walkthrough catalog."""
from __future__ import annotations

import argparse
import json
import math
from pathlib import Path

ROOT = Path(__file__).resolve().parents[1]


def normal_cdf(value: float) -> float:
    return 0.5 * (1 + math.erf(value / math.sqrt(2)))


def beta_tail_integer(alpha: int, beta: int, threshold: float) -> float:
    size = alpha + beta - 1
    return sum(
        math.comb(size, count)
        * threshold**count
        * (1 - threshold) ** (size - count)
        for count in range(alpha)
    )


def beta_binomial_mass(count: int, trials: int, alpha: int, beta: int) -> float:
    log_value = (
        math.log(math.comb(trials, count))
        + math.lgamma(alpha + count)
        + math.lgamma(beta + trials - count)
        - math.lgamma(alpha + beta + trials)
        - math.lgamma(alpha)
        - math.lgamma(beta)
        + math.lgamma(alpha + beta)
    )
    return math.exp(log_value)


def beta_binomial(inputs):
    alpha = inputs["alpha"] + inputs["successes"]
    beta = inputs["beta"] + inputs["trials"] - inputs["successes"]
    future = inputs["futureTrials"]
    return {
        "posteriorAlpha": alpha,
        "posteriorBeta": beta,
        "posteriorMean": alpha / (alpha + beta),
        "posteriorVariance": alpha * beta / ((alpha + beta) ** 2 * (alpha + beta + 1)),
        "probabilityAbove04": beta_tail_integer(alpha, beta, 0.4),
        "probabilityAbove05": beta_tail_integer(alpha, beta, 0.5),
        "futureExpected": future * alpha / (alpha + beta),
        "futureAtLeast3": sum(
            beta_binomial_mass(count, future, alpha, beta)
            for count in range(3, future + 1)
        ),
    }


def gamma_survival_integer(shape: int, rate: float, threshold: float) -> float:
    scaled = rate * threshold
    return math.exp(-scaled) * sum(
        scaled**count / math.factorial(count) for count in range(shape)
    )


def gamma_poisson_predictive_mass(count: int, shape: int, rate: float) -> float:
    return math.exp(
        math.lgamma(shape + count)
        - math.lgamma(shape)
        - math.lgamma(count + 1)
        + shape * math.log(rate / (rate + 1))
        + count * math.log(1 / (rate + 1))
    )


def gamma_poisson(inputs):
    shape = inputs["shape"] + sum(inputs["counts"])
    rate = inputs["rate"] + len(inputs["counts"])
    return {
        "posteriorShape": shape,
        "posteriorRate": rate,
        "posteriorMean": shape / rate,
        "posteriorVariance": shape / rate**2,
        "probabilityRateAbove3": gamma_survival_integer(shape, rate, 3),
        "predictiveZero": gamma_poisson_predictive_mass(0, shape, rate),
        "predictiveAtLeast5": 1 - sum(
            gamma_poisson_predictive_mass(count, shape, rate) for count in range(5)
        ),
    }


def normal_normal(inputs):
    prior_variance = inputs["priorSd"] ** 2
    known_variance = inputs["knownSd"] ** 2
    posterior_variance = 1 / (
        1 / prior_variance + inputs["sampleSize"] / known_variance
    )
    posterior_mean = posterior_variance * (
        inputs["priorMean"] / prior_variance
        + inputs["sampleSize"] * inputs["sampleMean"] / known_variance
    )
    posterior_sd = math.sqrt(posterior_variance)
    predictive_sd = math.sqrt(known_variance + posterior_variance)
    return {
        "posteriorMean": posterior_mean,
        "posteriorVariance": posterior_variance,
        "posteriorSd": posterior_sd,
        "credibleLow95": posterior_mean - 1.96 * posterior_sd,
        "credibleHigh95": posterior_mean + 1.96 * posterior_sd,
        "probabilityMeanAbove52": 1 - normal_cdf((52 - posterior_mean) / posterior_sd),
        "predictiveSd": predictive_sd,
        "predictiveAbove60": 1 - normal_cdf((60 - posterior_mean) / predictive_sd),
    }


def pooled_group(hyper_mean, hyper_sd, known_sd, size, sample_mean):
    variance = 1 / (1 / hyper_sd**2 + size / known_sd**2)
    mean = variance * (
        hyper_mean / hyper_sd**2 + size * sample_mean / known_sd**2
    )
    predictive_sd = math.sqrt(known_sd**2 + variance)
    return mean, variance, predictive_sd


def partial_pooling(inputs):
    group_a, group_b = inputs["groups"]
    common = (inputs["hyperMean"], inputs["hyperSd"], inputs["knownSd"])
    a_mean, a_variance, a_predictive_sd = pooled_group(
        *common, group_a["n"], group_a["mean"]
    )
    b_mean, b_variance, b_predictive_sd = pooled_group(
        *common, group_b["n"], group_b["mean"]
    )
    a_tau20 = pooled_group(
        inputs["hyperMean"], 20, inputs["knownSd"], group_a["n"], group_a["mean"]
    )[0]
    a_tau2 = pooled_group(
        inputs["hyperMean"], 2, inputs["knownSd"], group_a["n"], group_a["mean"]
    )[0]
    return {
        "groupAMean": a_mean,
        "groupAVariance": a_variance,
        "groupAPredictiveSd": a_predictive_sd,
        "groupAAbove60": 1 - normal_cdf((60 - a_mean) / a_predictive_sd),
        "groupBMean": b_mean,
        "groupBVariance": b_variance,
        "groupBPredictiveSd": b_predictive_sd,
        "groupBAbove60": 1 - normal_cdf((60 - b_mean) / b_predictive_sd),
        "groupAMeanTau20": a_tau20,
        "groupAMeanTau2": a_tau2,
        "mcseGroupA": math.sqrt(a_variance) / math.sqrt(inputs["ess"]),
    }


RUNNERS = {
    "beta_binomial": beta_binomial,
    "gamma_poisson": gamma_poisson,
    "normal_normal": normal_normal,
    "partial_pooling": partial_pooling,
}


def close(actual, expected, tolerance=1e-9):
    return math.isclose(actual, expected, rel_tol=tolerance, abs_tol=tolerance)


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("--json", action="store_true")
    args = parser.parse_args()
    catalog = json.loads(
        (ROOT / "content" / "bayesian_walkthroughs.json").read_text(encoding="utf-8")
    )
    report = {}
    for item in catalog:
        check = item["check"]
        actual = RUNNERS[check["kind"]](check["inputs"])
        mismatches = {
            key: {"actual": actual.get(key), "expected": expected}
            for key, expected in check["expected"].items()
            if key not in actual or not close(actual[key], expected)
        }
        report[item["id"]] = {
            "passed": not mismatches,
            "actual": actual,
            "mismatches": mismatches,
        }
    if args.json:
        print(json.dumps(report, ensure_ascii=False, indent=2))
    else:
        for item_id, result in report.items():
            print(f"{item_id}: {'PASS' if result['passed'] else 'FAIL'}")
    return 0 if all(result["passed"] for result in report.values()) else 1


if __name__ == "__main__":
    raise SystemExit(main())
