"""Recompute published paired deltas, exact sign flips, and four-endpoint Holm."""
import itertools
import json
import math
from pathlib import Path

ROOT = Path(__file__).resolve().parent if "__file__" in globals() else Path.cwd()


def exact_sign_flip(deltas):
    """Two-sided test: enumerate all 2**n signs; include ties in the tail."""
    observed = abs(sum(deltas))
    extreme = sum(abs(sum(sign * value for sign, value in zip(signs, deltas))) >= observed - 1e-12
                  for signs in itertools.product((-1, 1), repeat=len(deltas)))
    return extreme / (2 ** len(deltas))


def reproduce():
    records = json.loads((ROOT / "source.json").read_text(encoding="utf-8"))
    results = {}
    for protocol in ("S2", "S4"):
        assert len(records[protocol]) == 10
        assert sorted(row["fold"] for row in records[protocol]) == list(range(10))
        for metric, column in (("CI", "deltaCi"), ("RMSE", "deltaRmse")):
            # Use the published delta, not subtraction of independently rounded model scores.
            deltas = [row[column] for row in records[protocol]]
            results[f"{protocol}_{metric}"] = {
                "n": len(deltas), "mean_delta": sum(deltas) / len(deltas),
                "exact_p": exact_sign_flip(deltas)}
    ordered = sorted(results, key=lambda key: results[key]["exact_p"])
    previous = 0
    for rank, key in enumerate(ordered):
        adjusted = min(1, max(previous, (len(ordered) - rank) * results[key]["exact_p"]))
        results[key]["holm_p"] = adjusted
        previous = adjusted
    return results


if __name__ == "__main__":
    result = reproduce()
    expected = json.loads((ROOT / "expected.json").read_text(encoding="utf-8"))
    for contrast, values in expected.items():
        for key, value in values.items():
            assert math.isclose(result[contrast][key], value, rel_tol=0, abs_tol=1e-10), (contrast, key)
    assert [key for key, value in result.items() if value["holm_p"] < 0.05] == ["S2_RMSE"]
    print(json.dumps(result, indent=2))
    print("PASS: 20 reported folds and four contrasts checked; only S2 RMSE passes Holm < 0.05.")
    print("No training, bootstrap intervals, missing folds, or quantum hardware reproduced.")
