# Copyright (c) 2026 INNOVATIO SAS
# SPDX-License-Identifier: MIT
"""Reproduce the synthetic MSC-P-004 significance/effect-size examples."""

from __future__ import annotations

import argparse
import csv
import math
from pathlib import Path


Z_95 = 1.959963984540054
Z_90 = 1.6448536269514722


def normal_survival(z: float) -> float:
    return 0.5 * math.erfc(z / math.sqrt(2.0))


def analyze(row: dict[str, str]) -> dict[str, float | str]:
    n_a, y_a = int(row["n_a"]), int(row["conversions_a"])
    n_b, y_b = int(row["n_b"]), int(row["conversions_b"])
    delta = float(row["practical_threshold_pp"]) / 100.0
    if n_a <= 0 or n_b <= 0:
        raise ValueError("n_a and n_b must be strictly positive")
    if not (0 <= y_a <= n_a and 0 <= y_b <= n_b):
        raise ValueError("conversions must lie between zero and their group size")
    if delta <= 0:
        raise ValueError("practical_threshold_pp must be strictly positive")
    p_a, p_b = y_a / n_a, y_b / n_b
    difference = p_b - p_a
    pooled = (y_a + y_b) / (n_a + n_b)
    se_null = math.sqrt(pooled * (1.0 - pooled) * (1.0 / n_a + 1.0 / n_b))
    if se_null == 0.0:
        raise ValueError("the pooled null variance is zero")
    z = difference / se_null
    p_value = 2.0 * normal_survival(abs(z))
    se = math.sqrt(p_a * (1.0 - p_a) / n_a + p_b * (1.0 - p_b) / n_b)
    if se == 0.0:
        raise ValueError("the estimated sampling variance is zero")
    ci95 = (difference - Z_95 * se, difference + Z_95 * se)
    ci90 = (difference - Z_90 * se, difference + Z_90 * se)
    tost_p = max(normal_survival((difference + delta) / se), normal_survival((delta - difference) / se))
    equivalent = ci90[0] > -delta and ci90[1] < delta and tost_p < 0.05
    if ci95[0] > delta:
        decision = "practically-convincing-benefit"
    elif ci95[1] < -delta:
        decision = "practically-convincing-harm"
    elif p_value < 0.05 and equivalent:
        decision = "detectable-but-negligible"
    elif equivalent:
        decision = "practically-equivalent"
    else:
        decision = "inconclusive"
    return {
        "scenario": row["scenario"],
        "difference_pp": 100.0 * difference,
        "relative_lift_pct": math.nan if p_a == 0.0 else 100.0 * difference / p_a,
        "p_value": p_value,
        "ci95_low_pp": 100.0 * ci95[0],
        "ci95_high_pp": 100.0 * ci95[1],
        "ci90_low_pp": 100.0 * ci90[0],
        "ci90_high_pp": 100.0 * ci90[1],
        "tost_p": tost_p,
        "decision": decision,
    }


def main() -> None:
    parser = argparse.ArgumentParser()
    parser.add_argument("--csv", type=Path, default=Path(__file__).with_name("msc-p004-significance-effect-size.csv"))
    args = parser.parse_args()
    with args.csv.open(newline="", encoding="utf-8") as handle:
        rows = list(csv.DictReader(handle))
    for row in rows:
        result = analyze(row)
        print(
            f"{result['scenario']}: difference={result['difference_pp']:.6f}pp; "
            f"lift={result['relative_lift_pct']:.6f}%; p={result['p_value']:.8g}; "
            f"ci95=[{result['ci95_low_pp']:.6f}, {result['ci95_high_pp']:.6f}]pp; "
            f"ci90=[{result['ci90_low_pp']:.6f}, {result['ci90_high_pp']:.6f}]pp; "
            f"tost_p={result['tost_p']:.8g}; decision={result['decision']}"
        )


if __name__ == "__main__":
    main()
