#!/usr/bin/env python3
"""MSC-P-039 choosing a statistical test. MIT License.

Two things happen here, in this order, and the order is the point.

First the three candidate procedures are run on one concrete two-group file, so
that the reader sees them disagree on a real case. That disagreement settles
nothing: a single dataset cannot say which procedure is right.

Then the question that can be settled is settled by simulation. Each procedure
is applied to thousands of datasets drawn under a declared null hypothesis, and
the share of times it rejects is compared with the level it claims. A procedure
that claims five percent and rejects more often is not a test at five percent,
whatever its name. The procedure examined most closely is the popular one that
runs a normality test first and picks its test from the answer.

Nothing is drawn from a library generator: the datasets come from a declared
linear congruential recursion followed by a declared Box-Muller transform, so
every implementation reproduces the same digits without any seed convention.
"""
import csv
import math
import sys
from pathlib import Path

REQUIRED = ["unit_id", "group", "minutes"]
GROUPS = ("control", "treated")
NOMINAL = 0.05
REPLICATIONS = 2000
PER_GROUP = 25
SKEW_LOG_SD = 0.85
ALTERNATIVE_LOG_SHIFT = 0.45
NORMALITY_LEVEL = 0.05
LEVEL_TOLERANCE = 0.0125
LCG_SEED = 20260915
LCG_MULTIPLIER = 1103515245
LCG_INCREMENT = 12345
LCG_MODULUS = 2147483648
# The multiplier is split so that every intermediate product stays below 2^53,
# the largest integer a double holds exactly. Written directly, the product
# reaches 2.4e18 and languages whose numbers are doubles - R and SAS among them -
# would silently compute a different recursion. The split is exact arithmetic,
# not an approximation: 16838 * 65536 + 20077 is the multiplier itself.
LCG_HIGH = LCG_MULTIPLIER >> 16
LCG_LOW = LCG_MULTIPLIER - LCG_HIGH * 65536


def load(path: Path):
    with path.open(newline="", encoding="utf-8") as handle:
        reader = csv.DictReader(handle)
        if reader.fieldnames != REQUIRED:
            raise ValueError("exact schema required")
        rows = list(reader)
    if len(rows) < 20:
        raise ValueError("too few sampling units for the declared comparison")
    groups = {name: [] for name in GROUPS}
    for index, row in enumerate(rows, start=1):
        if set(row) != set(REQUIRED) or any(row[key] is None for key in REQUIRED):
            raise ValueError("exactly three cells required per row")
        if row["unit_id"].strip() != "U%03d" % index:
            raise ValueError("units must be ordered U001, U002, ... without gaps")
        name = row["group"].strip()
        if name not in groups:
            raise ValueError("the group column must contain only control and treated")
        value = float(row["minutes"])
        if not math.isfinite(value) or value <= 0:
            raise ValueError("the outcome must be finite and strictly positive")
        groups[name].append(value)
    for name in GROUPS:
        if len(groups[name]) < 10:
            raise ValueError("each declared group needs at least ten sampling units")
    return [groups[name] for name in GROUPS]


# --- distribution functions, standard library only -------------------------

def normal_tail(value):
    """Two-sided tail of the standard normal law."""
    return math.erfc(abs(value) / math.sqrt(2.0))


def beta_fraction(a, b, x):
    """Continued fraction for the incomplete beta, by the modified Lentz method."""
    tiny = 1e-30
    qab, qap, qam = a + b, a + 1.0, a - 1.0
    c = 1.0
    d = 1.0 - qab * x / qap
    if abs(d) < tiny:
        d = tiny
    d = 1.0 / d
    result = d
    for m in range(1, 300):
        m2 = 2 * m
        term = m * (b - m) * x / ((qam + m2) * (a + m2))
        d = 1.0 + term * d
        if abs(d) < tiny:
            d = tiny
        c = 1.0 + term / c
        if abs(c) < tiny:
            c = tiny
        d = 1.0 / d
        result *= d * c
        term = -(a + m) * (qab + m) * x / ((a + m2) * (qap + m2))
        d = 1.0 + term * d
        if abs(d) < tiny:
            d = tiny
        c = 1.0 + term / c
        if abs(c) < tiny:
            c = tiny
        d = 1.0 / d
        step = d * c
        result *= step
        if abs(step - 1.0) < 1e-14:
            break
    return result


def incomplete_beta(a, b, x):
    """Regularized incomplete beta function."""
    if x <= 0.0:
        return 0.0
    if x >= 1.0:
        return 1.0
    front = math.exp(
        math.lgamma(a + b) - math.lgamma(a) - math.lgamma(b)
        + a * math.log(x) + b * math.log(1.0 - x)
    )
    if x < (a + 1.0) / (a + b + 2.0):
        return front * beta_fraction(a, b, x) / a
    return 1.0 - front * beta_fraction(b, a, 1.0 - x) / b


def student_tail(statistic, degrees):
    """Two-sided tail of Student's law with the given degrees of freedom."""
    if degrees <= 0:
        raise ValueError("the declared comparison has no degrees of freedom")
    return incomplete_beta(degrees / 2.0, 0.5, degrees / (degrees + statistic * statistic))


# --- the three candidate procedures ----------------------------------------

def moments(values):
    count = len(values)
    mean = sum(values) / count
    variance = sum((value - mean) ** 2 for value in values) / (count - 1)
    return count, mean, variance


def welch(first, second):
    """Welch two-sample test: unequal variances, Satterthwaite degrees of freedom."""
    n1, m1, v1 = moments(first)
    n2, m2, v2 = moments(second)
    if v1 <= 0 and v2 <= 0:
        raise ValueError("both declared groups are constant")
    standard_error = math.sqrt(v1 / n1 + v2 / n2)
    statistic = (m2 - m1) / standard_error
    degrees = (v1 / n1 + v2 / n2) ** 2 / ((v1 / n1) ** 2 / (n1 - 1) + (v2 / n2) ** 2 / (n2 - 1))
    return statistic, degrees, student_tail(statistic, degrees)


def mann_whitney(first, second):
    """Rank test, normal approximation with the declared tie correction."""
    combined = [(value, 0) for value in first] + [(value, 1) for value in second]
    combined.sort(key=lambda entry: entry[0])
    ranks = [0.0] * len(combined)
    ties = []
    position = 0
    while position < len(combined):
        stop = position
        while stop + 1 < len(combined) and combined[stop + 1][0] == combined[position][0]:
            stop += 1
        average = (position + stop) / 2.0 + 1.0
        for index in range(position, stop + 1):
            ranks[index] = average
        ties.append(stop - position + 1)
        position = stop + 1
    n1, n2 = len(first), len(second)
    total = n1 + n2
    rank_sum = sum(rank for rank, (_, label) in zip(ranks, combined) if label == 0)
    statistic = rank_sum - n1 * (n1 + 1) / 2.0
    mean = n1 * n2 / 2.0
    correction = sum(count ** 3 - count for count in ties)
    variance = n1 * n2 / 12.0 * ((total + 1) - correction / (total * (total - 1.0)))
    if variance <= 0:
        raise ValueError("the declared rank test has no variance")
    standardized = (statistic - mean) / math.sqrt(variance)
    return statistic, standardized, normal_tail(standardized)


def jarque_bera(values):
    """Declared normality pretest: skewness and kurtosis against a chi-square with two degrees of freedom."""
    count = len(values)
    mean = sum(values) / count
    second = sum((value - mean) ** 2 for value in values) / count
    if second <= 0:
        raise ValueError("a declared group is constant, so normality cannot be assessed")
    third = sum((value - mean) ** 3 for value in values) / count
    fourth = sum((value - mean) ** 4 for value in values) / count
    skewness = third / second ** 1.5
    kurtosis = fourth / second ** 2
    statistic = count / 6.0 * (skewness ** 2 + (kurtosis - 3.0) ** 2 / 4.0)
    return statistic, math.exp(-statistic / 2.0)


def pretest(first, second):
    """The popular rule: test normality in each group, then pick the test from the answer."""
    _, left = jarque_bera(first)
    _, right = jarque_bera(second)
    if left > NORMALITY_LEVEL and right > NORMALITY_LEVEL:
        return "welch", welch(first, second)[2]
    return "mann-whitney", mann_whitney(first, second)[2]


# --- the declared simulation -----------------------------------------------

class Declared:
    """Declared recursion and Box-Muller transform; the only source of randomness."""

    def __init__(self, seed):
        self.state = seed
        self.spare = None

    def uniform(self):
        high = (LCG_HIGH * self.state) % LCG_MODULUS
        self.state = (high * 65536 + LCG_LOW * self.state + LCG_INCREMENT) % LCG_MODULUS
        return (self.state + 0.5) / LCG_MODULUS

    def normal(self):
        if self.spare is not None:
            value, self.spare = self.spare, None
            return value
        first, second = self.uniform(), self.uniform()
        radius = math.sqrt(-2.0 * math.log(first))
        angle = 2.0 * math.pi * second
        self.spare = radius * math.sin(angle)
        return radius * math.cos(angle)


def simulate(distribution, shift):
    """Rejection share of each procedure over the declared replications."""
    stream = Declared(LCG_SEED)
    counts = {"welch": 0, "mann-whitney": 0, "pretest": 0}
    chosen_welch = 0
    for _ in range(REPLICATIONS):
        first, second = [], []
        for _ in range(PER_GROUP):
            if distribution == "normal":
                first.append(stream.normal())
                second.append(stream.normal() + shift)
            else:
                first.append(math.exp(SKEW_LOG_SD * stream.normal()))
                second.append(math.exp(SKEW_LOG_SD * stream.normal() + shift))
        if welch(first, second)[2] < NOMINAL:
            counts["welch"] += 1
        if mann_whitney(first, second)[2] < NOMINAL:
            counts["mann-whitney"] += 1
        name, probability = pretest(first, second)
        if name == "welch":
            chosen_welch += 1
        if probability < NOMINAL:
            counts["pretest"] += 1
    return {name: count / REPLICATIONS for name, count in counts.items()}, chosen_welch / REPLICATIONS


def main():
    path = Path(sys.argv[1]) if len(sys.argv) > 1 else Path(__file__).resolve().parent.parent / "datasets" / "msc-p039-two-group-outcome.csv"
    control, treated = load(path)

    statistic, degrees, welch_probability = welch(control, treated)
    rank_statistic, standardized, rank_probability = mann_whitney(control, treated)
    control_jb, control_probability = jarque_bera(control)
    treated_jb, treated_probability = jarque_bera(treated)
    chosen, chosen_probability = pretest(control, treated)

    print(f"control_units={len(control)}")
    print(f"treated_units={len(treated)}")
    print(f"welch_statistic={statistic:.6f}")
    print(f"welch_degrees_of_freedom={degrees:.6f}")
    print(f"welch_probability={welch_probability:.6f}")
    print(f"rank_statistic={rank_statistic:.6f}")
    print(f"rank_standardized={standardized:.6f}")
    print(f"rank_probability={rank_probability:.6f}")
    print(f"control_normality_statistic={control_jb:.6f}")
    print(f"control_normality_probability={control_probability:.6f}")
    print(f"treated_normality_statistic={treated_jb:.6f}")
    print(f"treated_normality_probability={treated_probability:.6f}")
    print(f"pretest_choice={chosen}")
    print(f"pretest_probability={chosen_probability:.6f}")

    print(f"replications={REPLICATIONS}")
    print(f"units_per_group={PER_GROUP}")
    print(f"nominal_level={NOMINAL:.2f}")
    scenarios = [
        ("normal_null", "normal", 0.0),
        ("skewed_null", "skewed", 0.0),
        ("skewed_alternative", "skewed", ALTERNATIVE_LOG_SHIFT),
    ]
    levels = {}
    for label, distribution, shift in scenarios:
        rates, welch_share = simulate(distribution, shift)
        for name in ("welch", "mann-whitney", "pretest"):
            print(f"{label}_{name.replace('-', '_')}={rates[name]:.6f}")
        print(f"{label}_pretest_chose_welch={welch_share:.6f}")
        if shift == 0.0:
            levels[label] = rates

    flags = {}
    for label in ("normal_null", "skewed_null"):
        for name in ("welch", "mann-whitney", "pretest"):
            honest = abs(levels[label][name] - NOMINAL) <= LEVEL_TOLERANCE
            flags[(label, name)] = "PASS" if honest else "FAIL"
            print(f"{label}_{name.replace('-', '_')}_flag={flags[(label, name)]}")
    pretest_honest = all(flags[(label, "pretest")] == "PASS" for label in ("normal_null", "skewed_null"))
    print(f"level_tolerance={LEVEL_TOLERANCE:.4f}")
    print(f"verdict={'PRETEST_RULE_HOLDS_ITS_LEVEL' if pretest_honest else 'PRETEST_RULE_DOES_NOT_HOLD_ITS_LEVEL'}")


if __name__ == "__main__":
    main()
