#!/usr/bin/env python3
"""sunspot-evaluate.py — recompute the headline numbers of the sunspot
group-lifetime paper from the pinned data slices in this directory.

What it recomputes, and checks against the paper's printed values
(every comparison printed as recomputed / expected / verdict; any
mismatch exits nonzero):

  1. The life table (Table 1). A vendored Turnbull EM estimator for
     interval-censored lifetimes on the integer-day lattice, run on both
     handlings of the broken records, gives the daily death hazard at
     ages 0-10; each value is checked to 5e-4 against the printed table.
  2. The five-family parametric ladder: exponential, Weibull, Gompertz,
     gamma-Gompertz and inverse-Gaussian survival models fit by exact
     interval-censored log-likelihood on each track. Checked: on track B
     the gamma-Gompertz sits 9.25 AIC above the Weibull and the
     constant-hazard exponential loses by more than 7,000; on track A the
     gamma-Gompertz wins and the exponential loses by more than 2,000.
  3. The joint daily-hazard models on the 25,883 histories whose births
     were caught on the visible disk, with their daily corrected
     whole-spot areas: area+age (J3), area+frailty (J4), and
     area+age+frailty (J5). Checked: the aging slope B = +0.1839 per day
     (to 5e-4), the area exponent beta1 = -1.153 (to 0.01), the age
     term's AIC margin 601.5 and the frailty term's 1,868.7 (to 0.1).

Data provenance (archive recipe): the slices are threaded from the
Hathaway compilation of the Royal Greenwich Observatory (1874-1976) and
USAF/NOAA (1977-2024) daily sunspot-group catalog (solarcyclescience.com,
151 yearly files g1874.txt-g2024.txt), deduplicated and split into
segments under the paper's 14-day gap rule.
  histories_A.json.gz — the conservative handling of broken records:
      never joined across the far side, every limb exit censors
      (25,883 interval-censored lifetimes).
  histories_B.json.gz — the rotation-matched handling: records joined
      across the far side when the rotation-predicted return matches
      (24,489 lifetimes).
  area_paths.json.gz  — the daily corrected whole-spot area paths of the
      25,883 birth-observed groups (95,468 group-days).
The sha256 of each .gz file is pinned below and checked before use; any
mismatch refuses to run (exit 2, both hashes printed).

Not recomputed here (stays with the paper): the far-side matching sum
(13.9 million matchings enumerated exhaustively), the 0.047-per-day
late-age plateau of the marginalized fit, the interval-arithmetic
likelihood enclosures, and the reversal-keyed cycle-window fits.

Usage:
  python3 sunspot-evaluate.py           # full recomputation
  python3 sunspot-evaluate.py --check   # adds two controls that must FAIL:
      (1) a one-byte-corrupted copy of a slice must be refused (exit 2);
      (2) the joint hazard with the age slope forced to zero must lose
          by more than 500 AIC. If either control passes, exit 1.
Exit codes: 0 every comparison passed; 1 a comparison failed (or a
--check control failed to fail); 2 data-integrity refusal.
Requires: Python 3 + numpy + scipy. A full run stays under a minute even
on a heavily loaded machine (36-55 s measured wall-clock; the three joint
fits dominate); the first output line appears within about two seconds.
"""
import argparse
import gzip
import hashlib
import json
import math
import os
import subprocess
import sys
import tempfile
import time

import numpy as np
from scipy.optimize import minimize
from scipy.stats import invgauss

HERE = os.path.dirname(os.path.abspath(__file__))

# sha256 of the shipped .gz bytes, checked before use.
PINS = {
    'histories_A.json.gz':
        'cda428949e38d248cee4610d5d5d4b4d0245b0bb2a458f0a8e19be95157a8a4d',
    'histories_B.json.gz':
        '91a2481bdfd0adcac3b77a2c1e3d6e23137bd280ca4d1744da568a7e16edc744',
    'area_paths.json.gz':
        '087840485a7db82c48781a922fa9c19a251daca6aa179ad76441d99dcd047a0c',
}

# Expected values, from the paper (labeled oracle, never fed into a fit):
# Table 1 daily hazards, ages 0-10, both tracks.
TABLE1 = {
    'A': [0.286, 0.162, 0.135, 0.137, 0.126, 0.133,
          0.135, 0.131, 0.117, 0.054, 0.025],
    'B': [0.289, 0.163, 0.137, 0.137, 0.121, 0.126,
          0.120, 0.118, 0.094, 0.073, 0.025],
}
N_EXPECTED = {'A': 25883, 'B': 24489, 'paths': 25883}
GG_MINUS_WEIBULL_B = 9.25      # track B: AIC(gamma-Gompertz) - AIC(Weibull)
B_PER_DAY = 0.1839             # J5 aging slope, per day
BETA1 = -1.153                 # J5 area exponent
DAIC_AGE = 601.5               # AIC(J4) - AIC(J5): the age term's margin
DAIC_FRAILTY = 1868.7          # AIC(J3) - AIC(J5): the frailty term's margin


def load_slice(name, data_dir):
    """Read a pinned .gz slice; refuse (exit 2) on sha mismatch."""
    path = os.path.join(data_dir, name)
    blob = open(path, 'rb').read()
    got = hashlib.sha256(blob).hexdigest()
    if got != PINS[name]:
        print(f"INTEGRITY: {name} sha256 mismatch", flush=True)
        print(f"  expected {PINS[name]}")
        print(f"  got      {got}")
        print("refusing to run on unpinned data")
        sys.exit(2)
    print(f"  {name} sha256 OK {got[:16]}...", flush=True)
    return json.loads(gzip.decompress(blob))


# ---------------- Turnbull EM on the integer-day lattice (vendored) --------
def unique_rows(rows):
    from collections import Counter
    c = Counter(rows)
    u = sorted(c)
    w = np.array([c[k] for k in u], float)
    L = np.array([k[0] for k in u], float)
    R = np.array([k[1] for k in u], float)
    return L, R, w


def turnbull(L, R, w, M=None, tol=1e-10, maxit=20000):
    """Masses p_j on (j, j+1], j=0..M-1, plus tail (M, inf)."""
    if M is None:
        M = int(max(R.max(), L.max() + 1)) + 1
    nl = len(L)
    Ind = np.zeros((nl, M + 1), dtype=float)
    for i in range(nl):
        lo = int(L[i])
        if R[i] < 0:                    # right-censored: (L, inf)
            Ind[i, lo:] = 1.0
        else:
            Ind[i, lo:int(R[i])] = 1.0  # (L, R] = units L..R-1
    p = np.full(M + 1, 1.0 / (M + 1))
    n = w.sum()
    for it in range(1, maxit + 1):
        den = Ind @ p
        p_new = p * (Ind.T @ (w / den)) / n
        if np.abs(p_new - p).max() < tol:
            p = p_new
            break
        p = p_new
    return p, it, it < maxit, M


def daily_hazard(p, ages):
    F = np.concatenate([[0.0], np.cumsum(p)])
    out = []
    for a in ages:
        S_a = 1.0 - F[a]
        out.append((F[a + 1] - F[a]) / S_a if S_a > 1e-12 else float('nan'))
    return out


# ---------------- five-family parametric ladder (vendored) -----------------
def _phi(x):
    x = np.asarray(x, float)
    return np.where(np.abs(x) < 1e-8, 1.0 + x / 2.0,
                    np.expm1(x) / np.where(x == 0, 1, x))


def logS_expon(t, th):
    lam = math.exp(th[0])
    return -lam * t


def logS_weibull(t, th):
    b, c = math.exp(th[0]), math.exp(th[1])
    return -np.power(t / b, c, where=t > 0,
                     out=np.zeros_like(np.asarray(t, float)))


def logS_gompertz(t, th):
    A, B = math.exp(th[0]), th[1]
    return -A * t * _phi(B * t)


def logS_gammagompertz(t, th):
    A, B, s2 = math.exp(th[0]), th[1], math.exp(th[2])
    z = s2 * A * t * _phi(B * t)
    return -np.log1p(z) / s2


def logS_ig(t, th):
    m, lam = math.exp(th[0]), math.exp(th[1])
    t = np.asarray(t, float)
    out = np.zeros_like(t)
    pos = t > 0
    out[pos] = invgauss.logsf(t[pos], m / lam, scale=lam)
    return out


MODELS = {
    'exponential':    (logS_expon, 1, [[-1.2], [-0.5], [-2.0], [0.0], [-3.0]]),
    'weibull':        (logS_weibull, 2, [[1.1, 0.0], [1.1, -0.35], [1.1, 0.4],
                                         [0.5, 0.0], [2.0, 0.2]]),
    'gompertz':       (logS_gompertz, 2, [[-1.2, 0.0], [-1.2, 0.1],
                                          [-1.2, -0.1], [-2.0, 0.3],
                                          [-0.7, -0.3]]),
    'gamma-gompertz': (logS_gammagompertz, 3,
                       [[-1.2, 0.1, 0.0], [-1.5, 0.3, 0.7], [-1.0, 0.05, -1.0],
                        [-2.0, 0.5, 1.5], [-1.2, -0.1, 0.0], [-0.5, 0.8, 2.0]]),
    'inv-gaussian':   (logS_ig, 2, [[1.1, 1.1], [1.6, 0.0], [0.7, 0.7],
                                    [2.0, 2.0], [1.1, -0.5]]),
}


def negll_family(th, logS, L, R, w):
    try:
        a = logS(L, th)
    except (OverflowError, FloatingPointError):
        return 1e12
    dead = R >= 0
    ll = float(np.dot(w[~dead], a[~dead]))
    if dead.any():
        b = logS(R[dead], th)
        d = b - a[dead]
        if np.any(d >= 0):
            return 1e12
        term = a[dead] + np.log1p(-np.exp(d))
        if not np.all(np.isfinite(term)):
            return 1e12
        ll += float(np.dot(w[dead], term))
    if not math.isfinite(ll):
        return 1e12
    return -ll


def fit_family(name, L, R, w):
    logS, k, starts = MODELS[name]
    best = None
    for s in starts:
        res = minimize(negll_family, np.array(s, float), args=(logS, L, R, w),
                       method='Nelder-Mead',
                       options=dict(fatol=1e-9, xatol=1e-9,
                                    maxiter=20000, maxfev=40000))
        if best is None or res.fun < best.fun:
            best = res
    ll = -best.fun
    return dict(model=name, k=k, loglik=ll, AIC=2 * k - 2 * ll,
                params=[float(x) for x in best.x])


# ---------------- joint area/age/frailty daily hazard (vendored) -----------
def flatten(H):
    """Per-day rows; log-linear interpolation of log(area+1) on
    within-segment gaps; death-interval days carry the last area forward."""
    import bisect
    day_hist = []; day_age = []; day_z = []; day_kind = []
    for i, h in enumerate(H):
        days = h['days']
        ages = [d[0] for d in days]
        zs = [math.log(d[1] + 1.0) for d in days]
        for a in range(0, h['L']):
            j = bisect.bisect_right(ages, a) - 1
            if ages[j] == a:
                z = zs[j]
            else:
                a0, a1 = ages[j], ages[j + 1]
                z = zs[j] + (zs[j + 1] - zs[j]) * (a - a0) / (a1 - a0)
            day_hist.append(i); day_age.append(a); day_z.append(z)
            day_kind.append(0)
        if h['event'] == 'death':
            zcf = zs[-1]
            for a in range(h['L'], h['R']):
                day_hist.append(i); day_age.append(a); day_z.append(zcf)
                day_kind.append(1)
    return (np.array(day_hist), np.array(day_age, float),
            np.array(day_z, float), np.array(day_kind, np.int8))


def make_negll(day_hist, day_age, day_z, day_kind, nH, is_death):
    exp_idx = day_kind == 0
    int_idx = day_kind == 1
    dh_e, da_e, dz_e = day_hist[exp_idx], day_age[exp_idx], day_z[exp_idx]
    dh_i, da_i, dz_i = day_hist[int_idx], day_age[int_idx], day_z[int_idx]

    def H_arrays(b0, b1, B):
        eta_e = np.exp(b0 + b1 * dz_e + B * da_e)
        eta_i = np.exp(b0 + b1 * dz_i + B * da_i)
        H_L = np.bincount(dh_e, weights=eta_e, minlength=nH)
        H_R = H_L + np.bincount(dh_i, weights=eta_i, minlength=nH)
        return H_L, H_R

    def negll(theta, frailty, use_area, use_age):
        b0 = theta[0]
        k = 1
        b1 = theta[k] if use_area else 0.0
        if use_area:
            k += 1
        B = theta[k] if use_age else 0.0
        if use_age:
            k += 1
        try:
            H_L, H_R = H_arrays(b0, b1, B)
        except FloatingPointError:
            return 1e12
        if not np.all(np.isfinite(H_R)):
            return 1e12
        if frailty:
            s2 = math.exp(theta[k])
            SL = np.power(1.0 + s2 * H_L, -1.0 / s2)
            SR = np.power(1.0 + s2 * H_R, -1.0 / s2)
        else:
            SL = np.exp(-H_L)
            SR = np.exp(-H_R)
        d = is_death
        diff = SL[d] - SR[d]
        if np.any(diff <= 0) or not np.all(np.isfinite(diff)):
            return 1e12
        ll = np.log(diff).sum()
        if np.any(SL[~d] <= 0):
            return 1e12
        ll += np.log(SL[~d]).sum()
        if not math.isfinite(ll):
            return 1e12
        return -ll
    return negll


JOINT = {
    #                 use_area use_age frailty  k
    'J3_area_age':     (True,  True,  False, 3),
    'J4_area_frailty': (True,  False, True,  3),
    'J5_area_age_fr':  (True,  True,  True,  4),
}

# The analysis of record ran five starts per joint model; all five reached
# the same optimum (log-likelihood spread < 1e-3 across starts in the
# record), so this evaluator runs the first two of the record's start list
# per model. The recovered optima are checked against the paper's values
# below either way.
JOINT_STARTS = {
    'J3_area_age':     [[-1.5, -0.3, 0.0], [-1.0, -0.5, -0.05]],
    'J4_area_frailty': [[-1.5, -0.3, 0.0], [-1.0, -0.5, 1.0]],
    'J5_area_age_fr':  [[-1.5, -0.3, 0.0, 0.0], [-1.0, -0.5, 0.3, 1.0]],
}


def fit_joint(name, negll, starts=None):
    use_area, use_age, frailty, k = JOINT[name]
    best = None
    for s in (starts if starts is not None else JOINT_STARTS[name]):
        res = minimize(negll, np.array(s, float),
                       args=(frailty, use_area, use_age),
                       method='Nelder-Mead',
                       options=dict(fatol=1e-9, xatol=1e-9,
                                    maxiter=40000, maxfev=80000))
        if best is None or res.fun < best.fun:
            best = res
    ll = -best.fun
    return dict(k=k, loglik=ll, AIC=2 * k - 2 * ll,
                params=[float(x) for x in best.x])


# ---------------- comparisons ----------------------------------------------
FAILURES = []


def compare(label, got, expected, tol):
    ok = abs(got - expected) <= tol
    print(f"  {label}: recomputed {got:.6f} / expected {expected} "
          f"(tol {tol}) -> {'PASS' if ok else 'FAIL'}", flush=True)
    if not ok:
        FAILURES.append(label)
    return ok


def compare_floor(label, got, floor):
    ok = got > floor
    print(f"  {label}: recomputed {got:.1f} / must exceed {floor} "
          f"-> {'PASS' if ok else 'FAIL'}", flush=True)
    if not ok:
        FAILURES.append(label)
    return ok


def rows_from_histories(H):
    return [(h['L'], h['R'] if h['event'] == 'death' else -1) for h in H]


def main():
    ap = argparse.ArgumentParser(description=__doc__.splitlines()[0])
    ap.add_argument('--check', action='store_true',
                    help='also run the two controls that must fail')
    ap.add_argument('--data-dir', default=HERE,
                    help='directory holding the pinned .gz slices')
    ap.add_argument('--integrity-only', action='store_true',
                    help='verify the sha256 pins and exit')
    args = ap.parse_args()

    t0 = time.time()
    print("sunspot-evaluate: pinned-slice recomputation of the paper's "
          "headline numbers", flush=True)
    print("[1/4] integrity: checking sha256 pins", flush=True)
    HA = load_slice('histories_A.json.gz', args.data_dir)
    HB = load_slice('histories_B.json.gz', args.data_dir)
    PATHS = load_slice('area_paths.json.gz', args.data_dir)
    if args.integrity_only:
        print("integrity-only: all pins OK", flush=True)
        return 0

    np.seterr(all='ignore')
    hist = {'A': HA, 'B': HB}
    for t in ('A', 'B'):
        n = len(hist[t])
        ok = n == N_EXPECTED[t]
        print(f"  track {t}: {n} histories / expected {N_EXPECTED[t]} "
              f"-> {'PASS' if ok else 'FAIL'}", flush=True)
        if not ok:
            FAILURES.append(f'n_{t}')

    # ---- 1. life table -----------------------------------------------------
    ladders = {}
    for t in ('A', 'B'):
        print(f"[2/4] track {t}: Turnbull life table "
              f"(t={time.time() - t0:.1f}s)", flush=True)
        L, R, w = unique_rows(rows_from_histories(hist[t]))
        p, iters, conv, M = turnbull(L, R, w)
        # The analysis of record ran the same EM with the same cap: track A
        # converged at 10,460 iterations, track B ran to the 20,000-iteration
        # cap. The estimator at that point is the estimator of record; the
        # hazard-value comparisons below are the check.
        print(f"  track {t}: EM {iters} iterations "
              f"({'converged' if conv else 'iteration cap, as in the record'})",
              flush=True)
        haz = daily_hazard(p, list(range(0, 11)))
        for a in range(11):
            compare(f"track {t} hazard age {a}", haz[a], TABLE1[t][a], 5e-4)
        ladders[t] = (L, R, w)

    # ---- 2. five-family ladder --------------------------------------------
    aic = {}
    for t in ('A', 'B'):
        print(f"[3/4] track {t}: five-family parametric ladder "
              f"(t={time.time() - t0:.1f}s)", flush=True)
        L, R, w = ladders[t]
        fits = {m: fit_family(m, L, R, w) for m in MODELS}
        aic[t] = {m: fits[m]['AIC'] for m in fits}
        for m in sorted(fits, key=lambda m: fits[m]['AIC']):
            print(f"    {m:15s} k={fits[m]['k']} AIC={fits[m]['AIC']:.2f}",
                  flush=True)
    compare("track B AIC(gamma-Gompertz) - AIC(Weibull)",
            aic['B']['gamma-gompertz'] - aic['B']['weibull'],
            GG_MINUS_WEIBULL_B, 0.05)
    compare_floor("track B AIC(exponential) - AIC(Weibull)",
                  aic['B']['exponential'] - aic['B']['weibull'], 7000.0)
    winner_A = min(aic['A'], key=aic['A'].get)
    ok = winner_A == 'gamma-gompertz'
    print(f"  track A winner: {winner_A} / expected gamma-gompertz "
          f"-> {'PASS' if ok else 'FAIL'}", flush=True)
    if not ok:
        FAILURES.append('winner_A')
    compare_floor("track A AIC(exponential) - AIC(gamma-Gompertz)",
                  aic['A']['exponential'] - aic['A']['gamma-gompertz'],
                  2000.0)

    # ---- 3. joint area/age/frailty hazard ---------------------------------
    print(f"[4/4] joint daily-hazard fits J3/J4/J5 on the "
          f"{len(PATHS)} birth-observed area paths "
          f"(t={time.time() - t0:.1f}s)", flush=True)
    n_paths_ok = len(PATHS) == N_EXPECTED['paths']
    if not n_paths_ok:
        FAILURES.append('n_paths')
        print(f"  area paths: {len(PATHS)} / expected "
              f"{N_EXPECTED['paths']} -> FAIL", flush=True)
    is_death = np.array([h['event'] == 'death' for h in PATHS])
    day_hist, day_age, day_z, day_kind = flatten(PATHS)
    print(f"  {len(day_hist)} group-day rows", flush=True)
    negll = make_negll(day_hist, day_age, day_z, day_kind,
                       len(PATHS), is_death)
    jf = {}
    for name in JOINT:
        jf[name] = fit_joint(name, negll)
        print(f"    {name:15s} k={jf[name]['k']} AIC={jf[name]['AIC']:.2f} "
              f"(t={time.time() - t0:.1f}s)", flush=True)
    th5 = jf['J5_area_age_fr']['params']   # [b0, b1, B, log s2]
    compare("J5 aging slope B (per day)", th5[2], B_PER_DAY, 5e-4)
    compare("J5 area exponent beta1", th5[1], BETA1, 0.01)
    compare("AIC(J4) - AIC(J5), the age term's margin",
            jf['J4_area_frailty']['AIC'] - jf['J5_area_age_fr']['AIC'],
            DAIC_AGE, 0.1)
    compare("AIC(J3) - AIC(J5), the frailty term's margin",
            jf['J3_area_age']['AIC'] - jf['J5_area_age_fr']['AIC'],
            DAIC_FRAILTY, 0.1)

    # ---- summary -----------------------------------------------------------
    wall = time.time() - t0
    if FAILURES:
        print(f"RESULT: {len(FAILURES)} comparison(s) FAILED: "
              f"{', '.join(FAILURES)}  (wall {wall:.1f}s)", flush=True)
        return 1
    print(f"RESULT: every comparison PASSED  (wall {wall:.1f}s)", flush=True)

    if not args.check:
        return 0

    # ---- --check: two controls that must FAIL ------------------------------
    print("--check control 1: one-byte-corrupted slice must be refused "
          "(exit 2)", flush=True)
    ok1 = False
    with tempfile.TemporaryDirectory() as tmp:
        for name in PINS:
            blob = open(os.path.join(args.data_dir, name), 'rb').read()
            if name == 'histories_A.json.gz':
                mid = len(blob) // 2
                blob = blob[:mid] + bytes([blob[mid] ^ 0x01]) + blob[mid + 1:]
            open(os.path.join(tmp, name), 'wb').write(blob)
        r = subprocess.run(
            [sys.executable, os.path.abspath(__file__),
             '--data-dir', tmp, '--integrity-only'],
            capture_output=True, text=True)
        refused = (r.returncode == 2 and 'expected' in r.stdout
                   and 'got' in r.stdout)
        for line in r.stdout.splitlines():
            if 'INTEGRITY' in line or 'expected' in line or 'got' in line \
                    or 'refusing' in line:
                print(f"    child: {line.strip()}", flush=True)
        print(f"    child exit code {r.returncode} "
              f"-> control {'FAILED as required' if refused else 'PASSED (bad)'}",
              flush=True)
        ok1 = refused

    print("--check control 2: age slope forced to zero must lose by "
          "more than 500 AIC", flush=True)
    # Refit with the age column removed (three free parameters), scored
    # against the full J5 fit.
    zero = fit_joint('J4_area_frailty', negll)
    margin = zero['AIC'] - jf['J5_area_age_fr']['AIC']
    ok2 = margin > 500.0
    print(f"    AIC(zero-slope) - AIC(J5) = {margin:.1f} "
          f"-> control {'FAILED as required' if ok2 else 'PASSED (bad)'}",
          flush=True)

    if not (ok1 and ok2):
        print("RESULT: a --check control did not fail as required", flush=True)
        return 1
    print(f"RESULT: both controls failed as required  "
          f"(wall {time.time() - t0:.1f}s)", flush=True)
    return 0


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