#!/usr/bin/env python3
"""mixtures-evaluate.py — recompute the mixtures page's exact numbers live.

Everything printed on diagrams/mixtures.md that this script covers is
recomputed here from scratch and COMPARED, and the exit code is the verdict:
rc 0 only when every comparison passes, nonzero on any mismatch.

Default run (a few seconds on a laptop):
  1. Z(n,n), n = 1..8 — the two-state closed form
         Z(n,n) = (n!)^2/(2n+1)! [ 1/(n+1) + 2(H_{2n+1} - H_{n+1}) ]
     against the allocation-sum engine (two independent code paths inside
     the vendored evaluator: the phi-lattice DP and the brute expansion),
     exact Fraction equality.
  2. Z(2,2,2) = 37/63000 on the three-state model (both engines).
  3. The Lin–Sturmfels–Xu coin-toss benchmark: the printed rational for
     U = (2,2,2,2,2) and the 25-digit full marginal likelihood for
     U = (51,18,73,25,75) (their Example 2.5; largest printed value).
  4. The page's four displayed constants from their closed forms, each
     evaluated at two mpmath precisions (dps 30 and 60) and compared to the
     page's displayed digits:
         C   = sqrt(2) pi^{3/2} Gamma(5/4) / (32 * 3^{1/4})   = 0.16948...
         c   = 2 * 3^{1/4} / (pi Gamma(1/4))                  = 0.231089...
         b   = 1/(2 log 2)                                    = 0.721347...
         t33 = (2 pi)^{5/2} * 2 ln^2(2+sqrt(3)) / 81          = 4.23778...

Modes:
  --full     adds the fourth-root-law crossings by direct numerical
             quadrature of the two-component evidence integral on the
             honest ray U = N(1,4,6,4,1)/16 (log-space Gauss–Legendre
             tensor rule, two resolutions each):
             BF(22) ~ 1.2 and BF(1,281) = 3.00006 (the page's crossing).
             Adds under a minute of wall time.
  --series   verifies the published 800-term series bank in series/
             (the paper's Section 3.5 verification data, the standing
             test for any claimed recurrence on this family):
               a. sha256 of the four bank files against the pins below
                  (fail-closed: a tampered or truncated bank fails);
               b. the first terms of each ray recomputed EXACTLY by the
                  vendored allocation-sum engine (Z_phi on Model.coin,
                  the same convention as the LSX coin-toss benchmark
                  above) and reduced modulo every listed prime —
                  singular ray U = n(1,4,6,4,1), n = 0..10 (N <= 160),
                  at all fourteen 25-bit primes; all-ones ray
                  U = n(1,1,1,1,1), n = 0..20 (N <= 100), at both
                  24-bit primes;
               c. the recurrence-scan records: sixteen orders at two
                  primes per ray, no operator found at any order, the
                  per-order degree caps equal to the values recorded in
                  the records themselves (pinned below) and the search
                  region (r+1)(d+1) <= 670 stated in the paper.
             Adds seconds of wall time.
  --mutate   perturbs one pinned reference (the 37/63000 target) by 1e-30
             — and, with --series, one loaded series residue by +1 —
             and runs the SAME comparison paths: the run must then exit
             nonzero. A control that cannot fail is vacuous.

Requirements: python3 + mpmath + numpy (numpy only for --full and the
vendored engine's lattice DP). The vendored evaluator lsx_direct.py sits
beside this script (byte-identical to the copy inside the Mixalot bundle,
pinned in MANIFEST.sha256). No network, no imports outside this directory.
"""
import argparse
import hashlib
import json
import os
import sys
import time
from fractions import Fraction
from math import comb, factorial, lgamma

sys.stdout.reconfigure(line_buffering=True)
print("== mixtures-evaluate: exact evidence checks for the mixtures page ==",
      flush=True)

HERE = os.path.dirname(os.path.abspath(__file__))
sys.path.insert(0, HERE)

T0 = time.time()
FAILS = []


def check(name, ok, detail=""):
    s = "PASS" if ok else "FAIL"
    print(f"{s}  {name}  {detail}")
    if not ok:
        FAILS.append(name)


def H(j):
    """Harmonic number H_j as an exact Fraction."""
    return sum(Fraction(1, i) for i in range(1, j + 1))


def closed_nn(n):
    """The page's closed form for the two-state balanced evidence Z(n,n)."""
    return Fraction(factorial(n) ** 2, factorial(2 * n + 1)) * (
        Fraction(1, n + 1) + 2 * (H(2 * n + 1) - H(n + 1)))


# ------------------------------------------------------------- series bank
SERIES_DIR = os.path.join(HERE, "series")
# sha256 of the published bank files (also listed in series/SHA256SUMS.txt
# and MANIFEST.sha256). The two series files are byte-identical to the
# archive of record; the scan records differ from it only in metadata
# strings — two in the singular-ray record (a relative path and a
# provenance note), one in the all-ones record (the provenance note); see
# series/README.md.
SERIES_PINS = {
    "SERIES_800.json":
        "33bb006eaaa7bc4aa4d17799a378c11278304125feeac806cb5021b956e4c421",
    "ALLONES_SERIES_800.json":
        "914fb1a3834d49d8fbf16cc9596466154dcbef23d1d0267fb4deb63ba7a7366a",
    "RESCAN16_SINGULAR_800.json":
        "7c27a9ad77dc06ede1f44b970a53d9db463348a113a9c22cb426e379b767e9b7",
    "RESCAN16_ALLONES_800.json":
        "1ecff7a185a284fa874982c2241662f42c62050b7c95f46ff994d48ffd1c234f",
}
# Per-order coefficient-degree caps of the recurrence scans, as recorded in
# the scan records themselves (the "dcap" field of every per_order entry
# in series/RESCAN16_SINGULAR_800.json and RESCAN16_ALLONES_800.json):
# 140 for r <= 3 on the singular ray, 160 on the all-ones ray, then 132,
# 110, 94, 82, 72 for r = 4..8, down to 37 at r = 16. The paper (Section
# 3.5) prints only the endpoints of the singular-ray list (140 at r <= 3,
# 37 at r = 16) and the bound (r+1)(d+1) <= 670, the reachable region of
# an 801-term series with overdetermination margin 8; the intermediate
# caps and the all-ones ceiling are the records' own values, pinned here.
RECORD_CAPS = {4: 132, 5: 110, 6: 94, 7: 82, 8: 72, 16: 37}
REGION_BOUND = 670
SERIES_NCHECK = {"singular": 10, "allones": 20}


def series_checks(L, mutate=False):
    """The --series mode: pins, exact recomputation mod p, scan records."""
    def load(name):
        b = open(os.path.join(SERIES_DIR, name), "rb").read()
        return json.loads(b), hashlib.sha256(b).hexdigest()

    # a. sha256 pins (fail-closed)
    files, pins_ok = {}, True
    for name, pin in SERIES_PINS.items():
        try:
            obj, h = load(name)
        except OSError as e:
            check(f"series/{name} present", False, str(e))
            pins_ok = False
            continue
        files[name] = obj
        ok = (h == pin)
        pins_ok &= ok
        check(f"series/{name} sha256 == pin {pin[:16]}...", ok, h[:16] + "...")
    if len(files) < len(SERIES_PINS):
        return

    rays = {
        "singular": (files["SERIES_800.json"], (1, 4, 6, 4, 1),
                     files["RESCAN16_SINGULAR_800.json"], 140),
        "allones": (files["ALLONES_SERIES_800.json"], (1, 1, 1, 1, 1),
                    files["RESCAN16_ALLONES_800.json"], 160),
    }
    if mutate:
        s = rays["singular"][0]
        p0 = sorted(s)[0]
        s[p0][5] = (s[p0][5] + 1) % int(p0)
        print(f"[mutate] series/SERIES_800.json residue at prime {p0}, "
              f"n = 5 perturbed by +1 in memory (this run MUST exit nonzero)")

    # b. exact recomputation of the first terms, reduced mod every prime
    for ray, (S, U0, R, cap_low) in rays.items():
        primes = sorted(int(p) for p in S)
        t0 = time.time()
        shape_ok = all(len(S[str(p)]) == 801 for p in primes) and \
            all(0 <= v < p for p in primes for v in S[str(p)])
        check(f"{ray} ray: {len(primes)} primes x 801 residues, each in [0, p)",
              shape_ok, f"primes {primes[0]}..{primes[-1]}")
        nchk = SERIES_NCHECK[ray]
        n_ok, n_bad = 0, []
        for n in range(0, nchk + 1):
            if n == 0:
                z = Fraction(1)          # Z of the empty sample
            else:
                z = L.Z_phi(L.Model.coin([n * u for u in U0]))
            for p in primes:
                r = z.numerator * pow(z.denominator, p - 2, p) % p
                if r == S[str(p)][n]:
                    n_ok += 1
                else:
                    n_bad.append((p, n))
        Nmax = nchk * sum(U0)
        check(f"{ray} ray: terms n = 0..{nchk} (N <= {Nmax}) recomputed "
              f"exactly == bank residues at every prime",
              not n_bad,
              f"{n_ok}/{(nchk + 1) * len(primes)} residues match"
              + (f"; first mismatch (p, n) = {n_bad[0]}" if n_bad else "")
              + f" ({time.time() - t0:.1f}s)")

        # c. recurrence-scan record
        per = R["per_order"]
        scan_primes = sorted(int(p) for p in per)
        orders_ok = all(sorted(int(r) for r in per[p]) == list(range(1, 17))
                        for p in per)
        none_found = all(not per[p][r]["found"] for p in per for r in per[p])
        caps = {int(r): per[str(scan_primes[0])][r]["dcap"]
                for r in per[str(scan_primes[0])]}
        caps_same = all({int(r): per[p][r]["dcap"] for r in per[p]} == caps
                        for p in per)
        caps_record = all(caps[r] == cap_low for r in (1, 2, 3)) and \
            all(caps[r] == d for r, d in RECORD_CAPS.items())
        mono = all(caps[r] >= caps[r + 1] for r in range(1, 16))
        region = max((r + 1) * (d + 1) for r, d in caps.items() if r >= 4)
        check(f"{ray} ray scan record: 16 orders x {len(scan_primes)} primes, "
              f"no operator found, caps {cap_low}@r<=3 then "
              f"{', '.join(str(caps[r]) for r in range(4, 9))} ... {caps[16]}@r=16, "
              f"region (r+1)(d+1) <= {REGION_BOUND}",
              len(scan_primes) == 2 and set(scan_primes) <= set(primes)
              and R.get("nmax") == 800 and R.get("overdet") == 8
              and orders_ok and none_found and caps_same and caps_record
              and mono and region <= REGION_BOUND,
              f"max (r+1)(d+1) = {region}")


# ---------------------------------------------------------------- quadrature
def logZ2_ray(N, nodes):
    """log Z2 for the binomial mixture on the honest ray U = N(1,4,6,4,1)/16
    by a log-stabilized Gauss–Legendre tensor rule over (sigma,theta,rho)."""
    import numpy as np
    x, w = np.polynomial.legendre.leggauss(nodes)
    x = 0.5 * (x + 1.0)
    w = 0.5 * w
    lw = np.log(w)
    U = N * np.array([1.0, 4.0, 6.0, 4.0, 1.0]) / 16.0
    s = x[:, None, None]
    th = x[None, :, None]
    rh = x[None, None, :]
    logI = np.zeros((nodes, nodes, nodes))
    for v in range(5):
        fv_t = th ** (4 - v) * (1 - th) ** v
        fv_r = rh ** (4 - v) * (1 - rh) ** v
        logI = logI + U[v] * np.log(s * fv_t + (1 - s) * fv_r)
    logI = logI + lw[:, None, None] + lw[None, :, None] + lw[None, None, :]
    M = logI.max()
    return float(M + np.log(np.exp(logI - M).sum()))


def bf_ray(N, nodes):
    """Bayes factor Z1/Z2 on the honest ray (one component over two)."""
    logZ1 = 2.0 * lgamma(2 * N + 1) - lgamma(4 * N + 2)
    return __import__("math").exp(logZ1 - logZ2_ray(N, nodes))


def main():
    ap = argparse.ArgumentParser(
        description="recompute the mixtures page's exact numbers")
    ap.add_argument("--full", action="store_true",
                    help="add the N=22 and N=1,281 quadrature crossings")
    ap.add_argument("--series", action="store_true",
                    help="verify the published 800-term series bank in "
                         "series/ (pins, exact recomputation mod p, scan "
                         "records)")
    ap.add_argument("--mutate", action="store_true",
                    help="perturb one pinned reference (and, with --series, "
                         "one series residue); run must exit nonzero")
    args = ap.parse_args()

    import lsx_direct as L

    # 1) Z(n,n) n=1..8: closed form vs both allocation-sum engines
    ok_all = True
    for n in range(1, 9):
        m = L.model_1var((n, n))
        cf = closed_nn(n)
        ok = (L.Z_phi(m) == cf) and (L.Z_xsum(m) == cf)
        ok_all &= ok
    check("Z(n,n) n=1..8: closed form == phi-DP == brute expansion", ok_all,
          f"e.g. Z(8,8) = {closed_nn(8)}")

    # 2) Z(2,2,2) on the three-state model
    target = Fraction(37, 63000)
    if args.mutate:
        target += Fraction(1, 10 ** 30)
        print("[mutate] Z(2,2,2) reference perturbed by 1e-30 "
              "(this run MUST exit nonzero)")
    cells = [(((1, 0, 0),), 2), (((0, 1, 0),), 2), (((0, 0, 1),), 2)]
    m3 = L.Model([1], cells)
    z3 = L.Z_phi(m3)
    check("Z(2,2,2) == 37/63000 (page value; both engines)",
          z3 == target and L.Z_xsum(m3) == target, f"Z = {z3}")

    # 3) LSX coin-toss benchmark
    mc = L.Model.coin(L.COIN10_U)
    z10 = L.Z_phi(mc)
    check("coin toss U=(2,2,2,2,2) == printed LSX rational",
          z10 == L.COIN10_REF and L.Z_xsum(mc) == L.COIN10_REF,
          f"Z = {z10}")
    t0 = time.time()
    z242 = L.Z_phi(L.Model.coin(L.COIN242_U))
    fml = L.full_ml_reduced(z242, L.COIN242_U, L.COIN_MULT)
    import mpmath as mp
    mp.mp.dps = 30
    v = mp.mpf(fml.numerator) / mp.mpf(fml.denominator)
    ref = mp.mpf(L.COIN242_FULLML_25DIG)
    rel = abs(v - ref) / ref
    check("coin toss U=(51,18,73,25,75) full ML == 25-digit LSX value (>=24d)",
          rel < mp.mpf("2e-24"),
          f"{mp.nstr(v, 25)} vs {L.COIN242_FULLML_25DIG} ({time.time()-t0:.1f}s)")

    # 4) four displayed constants, two precisions each
    #    (displayed-digit references: the page and paper print C = 0.16948...,
    #     c = 0.231089..., 1/(2 log 2) = 0.7213..., t33 = 4.2378...)
    REFS = {
        "C":   "0.16948473887776722409",
        "c":   "0.231089047016534",
        "b":   "0.721347520444482",
        "t33": "4.23778021756722",
    }

    def constants(dps):
        mp.mp.dps = dps
        one4 = mp.mpf(1) / 4
        return {
            "C": mp.sqrt(2) * mp.pi ** mp.mpf("1.5") * mp.gamma(5 * one4)
                 / (32 * mp.power(3, one4)),
            "c": 2 * mp.power(3, one4) / (mp.pi * mp.gamma(one4)),
            "b": 1 / (2 * mp.log(2)),
            "t33": mp.power(2 * mp.pi, mp.mpf("2.5")) * 2
                   * mp.log(2 + mp.sqrt(3)) ** 2 / 81,
        }

    v30, v60 = constants(30), constants(60)
    mp.mp.dps = 60
    for k in REFS:
        stable = abs(v60[k] - mp.mpf(v30[k])) < abs(v60[k]) * mp.mpf("1e-25")
        shown = mp.nstr(v60[k], len(REFS[k].replace(".", "").lstrip("0")))
        check(f"constant {k}: two precisions stable (dps 30/60) and "
              f"== displayed digits {REFS[k]}",
              stable and shown == REFS[k], f"value {shown}")

    # 5) --full: the fourth-root-law crossings by direct quadrature
    if args.full:
        print("[full] quadrature crossings on the honest ray "
              "U = N(1,4,6,4,1)/16 (log-space Gauss-Legendre, "
              "two resolutions each)...")
        t0 = time.time()
        b22a, b22b = bf_ray(22, 96), bf_ray(22, 160)
        check("BF(22) ~ 1.2 (page: 'about 1.2'; two resolutions agree)",
              abs(b22a - b22b) < 1e-9 * b22b and round(b22b, 1) == 1.2,
              f"BF(22) = {b22b:.10f} ({time.time()-t0:.1f}s)")
        t0 = time.time()
        b81a, b81b = bf_ray(1281, 224), bf_ray(1281, 320)
        check("BF(1,281) = 3.00006 (the page's 3-to-1 crossing; "
              "two resolutions agree)",
              abs(b81a - b81b) < 1e-8 * b81b and round(b81b, 5) == 3.00006,
              f"BF(1281) = {b81b:.10f} ({time.time()-t0:.1f}s)")

    # 6) --series: the published 800-term series bank
    if args.series:
        print("[series] verifying the published series bank in series/ "
              "(pins, exact recomputation mod p, scan records)...")
        series_checks(L, mutate=args.mutate)

    wall = time.time() - T0
    if FAILS:
        print(f"OVERALL FAIL ({len(FAILS)} failed): {', '.join(FAILS)} "
              f"[{wall:.1f}s]")
        return 1
    print(f"OVERALL PASS ({wall:.1f}s wall)")
    return 0


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