#!/usr/bin/env python3
"""trials-evaluate.py — recompute the drug-trials census tallies.

Reads the nine per-design cell files straight out of trials-census.zip
(beside this script, no extraction needed), verifies each against the
sha256 recorded in the archive's MANIFEST, and recomputes, in exact
rational arithmetic:

  * the straddle flag of every cell — the paper's own rule,
    (p-hat - b)^2 <= 4 p-hat (1 - p-hat) / nsim  (2 Monte-Carlo standard
    errors; p-hat in {0,1} straddles iff p-hat = b) — cross-checked
    against every stored per-cell flag and the archive's per-row tallies;
  * the census counts: 1,076 cells, 977 claims under the full
    three-column instrument, 351 of those within their own simulation
    noise, 77 cells printing exactly their boundary;
  * the required simulation count for every straddling cell with a
    nonzero margin — smallest n with (p-hat - b)^2 > 4 p-hat (1-p-hat)/n
    — and its median, 45,697;
  * two of the paper's exact values from scratch, with no stored answer:
    the power of the Simon two-stage design behind the five reversed
    power entries (r1=1/n1=10, r=5/N=16, p1=0.45), a finite binomial sum,
    and the repetitions a printed 0.0498 against a 0.05 cap requires.

The decided-pool counts the page's results section quotes (140 adjudicated
reproducible, 123 decided, 85/30/8) live in the paper's adjudication annex,
not in these census files, so this script does not recompute them.

--check runs a mutation control that must FAIL: one cell's printed value
is perturbed by one unit in its last printed digit, chosen so its straddle
flag flips; the census tally has to break. If the tally still matches,
the script exits 1.

Exit code: 0 only if every recomputed value agrees (or, under --check,
the control fired). Python standard library only.
"""

import argparse
import hashlib
import json
import sys
import time
import zipfile
from fractions import Fraction as F
from os.path import abspath, dirname, join

sys.stdout.reconfigure(line_buffering=True)

ZIP = join(dirname(abspath(__file__)), "trials-census.zip")

SIMON_POWER = F(65613605784585863859, 81920000000000000000)

FAILURES = []


def fail(msg):
    FAILURES.append(msg)
    print("  DISAGREES:", msg)


def frac(s):
    """'99/1000' or '0.159' or 12 -> exact Fraction."""
    if isinstance(s, int):
        return F(s)
    if isinstance(s, float):
        return F(str(s))
    s = str(s).strip()
    if "/" in s:
        num, den = s.split("/")
        return F(int(num), int(den))
    return F(s)


def straddle(p, b, n, k=2):
    """The paper's rule, exact: (p-b)^2 <= k^2 p(1-p)/n; p in {0,1} has
    Monte-Carlo standard error 0 and straddles iff p == b."""
    if p in (F(0), F(1)):
        return p == b
    return (p - b) ** 2 * n <= k * k * p * (1 - p)


def ulp_of(printed, unit):
    """One unit in the last printed decimal, in proportion units."""
    t = str(printed).strip()
    pct = unit == "percent" or t.endswith("%")
    t = t.rstrip("%").strip()
    if "/" in t or "." not in t:
        return None
    u = F(1, 10 ** len(t.split(".")[1]))
    return u / 100 if pct else u


def load_cells(z, manifest):
    """The nine per-design cell files -> one list of dicts, replicating the
    field mapping of the census's own loader (corpus/referee_checks.py in
    the archive). sha256 of each file is checked against the MANIFEST."""
    shas = {e["path"]: e["sha256"] for e in manifest["files"]}

    def read(path):
        raw = z.read(path)
        got = hashlib.sha256(raw).hexdigest()
        if got != shas[path]:
            sys.exit("FAIL: %s sha256 %s != MANIFEST %s" % (path, got, shas[path]))
        return json.loads(raw)

    cells = []

    def add(row, cid, phat, b, nsim, printed, unit, cal, dup, s2, s3, pfe):
        cells.append(dict(row=row, id=cid, phat=frac(phat), b=frac(b),
                          nsim=nsim, printed=printed, unit=unit,
                          calibration=bool(cal), dup=bool(dup),
                          s2_stored=s2, s3_stored=s3, has_pfe=pfe))

    for c in read("leg1_bop2/results/adjudication.json"):
        add("row1", c["cell"], c["printed"], c["boundary"], c["nsim"],
            None, "prop", c["calibration"], False,
            c["straddle_2se"], c["straddle_3se"], True)
    for c in read("leg3_basket/results/leg3_cells.json")["cells"]:
        add("leg3", c["cell_id"], c["phat"], c["boundary"], c["nsim"],
            c["printed_percent"], "percent",
            "calibration threshold" in c["boundary_class"],
            c["duplicate_reprint_of"] is not None,
            c["straddle_2se"], c["straddle_3se"], False)
    for c in read("corpus/row4_bacis/results/adjudication.json"):
        add("row4", c["id"], c["p_printed"], c["b"], c["nsim"],
            c["p_printed"], "prop", c["calibration"], False,
            c["straddle_2mcse"], c["straddle_3mcse"], True)
    for c in read("corpus/row5_cbhm/cells/cells.json")["cells"]:
        add("row5", c["cell_id"], c["phat"], c["boundary"], c["nsim"],
            c["printed_value"], "prop", c["calibration"],
            c["duplicate_reprint_of"] is not None,
            c["straddle_2se"], c["straddle_3se"], True)
    for c in read("corpus/row6_adaptr/cells/cells.json")["cells"]:
        add("row6", c["id"], c["phat"], c["boundary"], c["nsim"],
            c["printed"], "prop", c["calibration"],
            c["duplicate_reprint_of"] is not None,
            c["straddle2"], c["straddle3"], True)
    for c in read("corpus/row7_assistant/cells/adjudication.json")["cells"]:
        add("row7", c["id"], c["phat"], c["boundary"], c["nsim"],
            c["printed_value"], "percent", c["calibration"],
            c["duplicate_reprint_of"] is not None,
            c["straddle_2mcse"], c["straddle_3mcse"], True)
    for c in read("corpus/row8_bayesctdesign/adjudication_final.json"):
        add("row8", c["id"], c["p_hat_printed"], c["boundary"], c["nsim"],
            c["p_hat_printed"], "prop", False,
            c["duplicate_reprint_of"] is not None,
            c["straddle_2mcse"], c["straddle_3mcse"], True)
    for c in read("corpus/row9_localmem/results/cells_final.json"):
        add("row9", c["id"], c["p_printed"], c["boundary"], c["nsim"],
            c["p_printed"], "prop", c["calibration"],
            c["duplicate_reprint_of"] is not None,
            c["straddle2"], c["straddle3"], True)
    for c in read("corpus/row10_scurmst/results/cells.json")["cells"]:
        add("row10", c["cell"], c["p_hat"], c["boundary"], c["nsim"],
            c["printed"], "prop", c["calibration"],
            c["duplicate_reprint_of"] is not None,
            c["straddle_2se"], c["straddle_3se"], True)
    return cells


def required_count(p, b):
    """Smallest integer n with (p-b)^2 * n > 4 p (1-p), exact."""
    q = 4 * p * (1 - p) / (p - b) ** 2
    return q.numerator // q.denominator + 1


def simon_power(p, r1=1, n1=10, r=5, N=16):
    """Exact rejection probability of Simon's two-stage design: continue
    iff first-stage responses exceed r1, declare promising iff total
    responses exceed r. A finite sum of binomial terms."""
    from math import comb
    n2 = N - n1
    w1 = [F(comb(n1, x)) * p ** x * (1 - p) ** (n1 - x) for x in range(n1 + 1)]
    w2 = [F(comb(n2, x)) * p ** x * (1 - p) ** (n2 - x) for x in range(n2 + 1)]
    tot = F(0)
    for x1 in range(r1 + 1, n1 + 1):
        need = r + 1 - x1
        tot += w1[x1] * sum(w2[x2] for x2 in range(max(need, 0), n2 + 1))
    return tot


def census(cells):
    """Recompute every tally; returns the straddle-2SE count over the
    full-instrument cells (the page's 'within their own noise' number)."""
    n = len(cells)
    print("cells loaded from the archive: %s" % format(n, ","))
    if n != 1076:
        fail("cell count %d != 1,076" % n)
    pfe = [c for c in cells if c["has_pfe"]]
    print("claims under the full three-column instrument: %d" % len(pfe))
    if len(pfe) != 977:
        fail("full-instrument count %d != 977" % len(pfe))

    mism = 0
    for c in cells:
        c["s2"] = straddle(c["phat"], c["b"], c["nsim"], 2)
        c["s3"] = straddle(c["phat"], c["b"], c["nsim"], 3)
        if c["s2"] != c["s2_stored"] or c["s3"] != c["s3_stored"]:
            mism += 1
            fail("straddle flag mismatch at %s/%s" % (c["row"], c["id"]))
    print("straddle flags recomputed exactly: %d of %d agree with the "
          "stored flags" % (n - mism, n))

    s2_pfe = sum(1 for c in pfe if c["s2"])
    print("within their own simulation noise (2 MC-SE): %d of %d claims"
          % (s2_pfe, len(pfe)))
    if s2_pfe != 351:
        fail("within-noise count %d != 351" % s2_pfe)

    onb = sum(1 for c in cells if c["phat"] == c["b"])
    print("cells printing exactly their boundary: %d" % onb)
    if onb != 77:
        fail("exact-boundary count %d != 77" % onb)

    reqs = sorted(required_count(c["phat"], c["b"]) for c in cells
                  if c["s2"] and c["phat"] != c["b"]
                  and c["phat"] not in (F(0), F(1)))
    med = reqs[len(reqs) // 2]
    print("required simulation counts (straddling cells with a nonzero "
          "margin, %d cells): median %s" % (len(reqs), format(med, ",")))
    if med != 45697:
        fail("median required count %s != 45,697" % med)
    return s2_pfe


def crosscheck_row_tallies(cells, z):
    """The archive's own per-row tally record must match the recomputation."""
    tal = json.loads(z.read("corpus/ROW_TALLIES.json"))
    bad = 0
    for rw, v in tal["rows"].items():
        mine = [c for c in cells if c["row"] == rw]
        got = (len(mine), sum(c["s2"] for c in mine),
               sum(c["s3"] for c in mine),
               sum(c["calibration"] for c in mine))
        want = (v["cells"], v["straddle2"], v["straddle3"], v["calibration"])
        if got != want:
            bad += 1
            fail("row tally mismatch %s: recomputed %s, recorded %s"
                 % (rw, got, want))
    print("per-row tallies vs the archive's ROW_TALLIES record: %d rows "
          "checked, %d disagree" % (len(tal["rows"]), bad))


def exact_values():
    pw = simon_power(F(9, 20))
    print("Simon two-stage power at p1 = 0.45 (recomputed from the design, "
          "no stored answer):")
    print("  %s = %.5f" % (pw, float(pw)))
    if pw != SIMON_POWER:
        fail("Simon power %s != %s" % (pw, SIMON_POWER))
    nreq = required_count(F(498, 10000), F(5, 100))
    print("required count for a printed 0.0498 against a 0.05 cap: n = %s"
          % format(nreq, ","))
    if nreq != 4731997:
        fail("required count %d != 4,731,997" % nreq)


def main():
    ap = argparse.ArgumentParser(description=__doc__.splitlines()[0])
    ap.add_argument("--check", action="store_true",
                    help="mutation control: perturb one cell's printed value "
                         "by one last-digit unit so its straddle flag flips; "
                         "the census tally MUST break")
    args = ap.parse_args()

    t0 = time.time()
    print("trials-evaluate: exact recomputation of the drug-trials census"
          + (" (--check control)" if args.check else ""))
    z = zipfile.ZipFile(ZIP)
    manifest = json.loads(z.read("MANIFEST.json"))
    cells = load_cells(z, manifest)

    if args.check:
        target = None
        for c in cells:
            if not c["has_pfe"]:
                continue  # the within-noise tally counts full-instrument cells
            u = ulp_of(c["printed"], c["unit"]) if c["printed"] else None
            if u is None:
                continue
            for p2 in (c["phat"] + u, c["phat"] - u):
                if 0 <= p2 <= 1 and straddle(p2, c["b"], c["nsim"], 2) != \
                        straddle(c["phat"], c["b"], c["nsim"], 2):
                    target = (c, p2)
                    break
            if target:
                break
        if target is None:
            print("CONTROL BROKEN: no cell found whose flag flips — FAIL")
            return 1
        c, p2 = target
        print("control: cell %s/%s printed value %s -> %s (one last-digit "
              "unit); its straddle flag flips, so the census must break"
              % (c["row"], c["id"], c["phat"], p2))
        c["phat"] = p2
        c["s2_stored"] = straddle(p2, c["b"], c["nsim"], 2)  # isolate the tally
        c["s3_stored"] = straddle(p2, c["b"], c["nsim"], 3)
        census(cells)
        if FAILURES:
            print("CONTROL FIRED as required (%d disagreements) — PASS"
                  % len(FAILURES))
            return 0
        print("CONTROL DID NOT FIRE: tallies still match — FAIL")
        return 1

    census(cells)
    crosscheck_row_tallies(cells, z)
    exact_values()
    print("total wall time: %.1f s" % (time.time() - t0))
    if FAILURES:
        print("OVERALL: FAIL — %d disagreement(s)" % len(FAILURES))
        return 1
    print("OVERALL: PASS — every recomputed tally and exact value agrees")
    return 0


if __name__ == "__main__":
    sys.exit(main())
