# Adversarial checks of the bundle's breakpoint-search test, written independently of its code.
# Reads results/annual.csv (year, anomaly) from the re-run. Prints what each check finds.
import csv, sys, math
import numpy as np
from scipy.optimize import minimize

rows = list(csv.DictReader(open(sys.argv[1])))
years = np.array([int(r["year"]) for r in rows]); y = np.array([float(r["anomaly_c"]) for r in rows])
rng = np.random.Generator(np.random.PCG64(7))
B = 20000

def basis(yrs, knot=None):
    cols = [np.ones(len(yrs)), (yrs - 1970) / 10.0]
    if knot is not None: cols.append(np.maximum(0, yrs - knot) / 10.0)
    return np.column_stack(cols)

def maxF_scalar(yrs, v, knots):
    x0 = basis(yrs); r0 = v - x0 @ np.linalg.lstsq(x0, v, rcond=None)[0]; s0 = r0 @ r0
    best = (-1, None)
    for k in knots:
        x1 = basis(yrs, k); r1 = v - x1 @ np.linalg.lstsq(x1, v, rcond=None)[0]; s1 = r1 @ r1
        f = (s0 - s1) / s1 * (len(v) - 3)
        if f > best[0]: best = (f, k)
    return best

def scanner(yrs, knots):
    # projection form, checked below against maxF_scalar
    x0 = basis(yrs); P = x0 @ np.linalg.pinv(x0)
    H = np.maximum(0, yrs[:, None] - knots[None, :]) / 10.0
    Hr = H - P @ H; U = Hr / np.linalg.norm(Hr, axis=0)
    def f(V):
        R = V - V @ P.T; s0 = np.sum(R * R, axis=1); imp = (R @ U) ** 2
        return np.max(imp / (s0[:, None] - imp) * (V.shape[1] - 3), axis=1)
    return f

def ar1(n, rho, count):
    e = rng.standard_normal((count, n)); out = np.empty_like(e)
    out[:, 0] = e[:, 0] / math.sqrt(1 - rho * rho)
    for i in range(1, n): out[:, i] = rho * out[:, i - 1] + e[:, i]
    return out

def arma11(n, phi, theta, count, burn=200):
    e = rng.standard_normal((count, n + burn)); out = np.zeros_like(e)
    for i in range(1, n + burn): out[:, i] = phi * out[:, i - 1] + e[:, i] + theta * e[:, i - 1]
    return out[:, burn:]

def null_rho(v, yrs):
    x0 = basis(yrs); r = v - x0 @ np.linalg.lstsq(x0, v, rcond=None)[0]
    return float(r[1:] @ r[:-1] / (r[:-1] @ r[:-1]))

def rho_batch(V, yrs):
    x0 = basis(yrs); P = x0 @ np.linalg.pinv(x0); R = V - V @ P.T
    return np.sum(R[:, 1:] * R[:, :-1], axis=1) / np.sum(R[:, :-1] ** 2, axis=1)

def knots_for(end): return np.arange(1985, min(2015, end - 10) + 1)

print("== 1. Observed statistics (independent scalar refits)")
obs = {}
for end in (2025, 2024, 2022):
    m = years <= end; f, k = maxF_scalar(years[m], y[m], knots_for(end)); obs[end] = f
    sc = scanner(years[m], knots_for(end)); assert abs(sc(y[m][None, :])[0] - f) < 1e-9
    print(f"  through {end}: maxF {f:.4f} at knot {k}; null rho from straight-line residuals {null_rho(y[m], years[m]):.4f}")

n = len(y); sc = scanner(years, knots_for(2025))
def p_at(noise): s = sc(noise); return (1 + np.count_nonzero(s >= obs[2025])) / (len(s) + 1)

print("== 2. Bias of the plug-in rho, and the p-value at a bias-corrected rho (endpoint 2025)")
target = null_rho(y, years)
grid = np.round(np.arange(0.25, 0.56, 0.025), 3)
means = {}
for r in grid:
    means[r] = float(np.mean(rho_batch(ar1(n, r, 4000), years)))
for r in grid: print(f"  true rho {r:.3f}: mean estimated rho {means[r]:.4f}")
corrected = float(np.interp(target, [means[r] for r in grid], grid))
print(f"  observed plug-in rho {target:.4f}; mean-unbiased rho {corrected:.4f}")
for r in (round(target, 4), corrected):
    print(f"  p at rho {r:.4f}: {p_at(ar1(n, r, B)):.4f}")

print("== 3. Accounting for rho's estimation: double-bootstrap-style calibration")
# For data drawn at the corrected rho, re-estimate rho in each replicate and use the
# replicate's own plug-in p-value; the rejection rate at nominal 0.05 is the test's size.
def plug_in_critical(r, cache={}):
    key = round(r, 2)
    if key not in cache: cache[key] = np.quantile(sc(ar1(n, key, 4000)), 0.95)
    return cache[key]
for true_r in (round(target, 3), round(corrected, 3)):
    V = ar1(n, true_r, 2000); S = sc(V); R = np.clip(rho_batch(V, years), -0.95, 0.95)
    rej = np.mean([S[i] > plug_in_critical(R[i]) for i in range(len(S))])
    print(f"  true rho {true_r}: plug-in test rejects {rej:.3f} at nominal 0.05")

print("== 4. ARMA(1,1) and AR(2) nulls fitted to the straight-line residuals (endpoint 2025)")
x0 = basis(years); resid = y - x0 @ np.linalg.lstsq(x0, y, rcond=None)[0]
def css(params, kind):
    # Conditional sum of squares for zero-mean ARMA residuals; Gaussian likelihood via its innovations.
    e = np.zeros(n)
    for i in range(n):
        if kind == "arma11":
            phi, theta = params
            e[i] = resid[i] - (phi * resid[i-1] + theta * e[i-1] if i else 0.0)
        elif kind == "ar2":
            a1, a2 = params
            e[i] = resid[i] - ((a1 * resid[i-1] if i else 0.0) + (a2 * resid[i-2] if i > 1 else 0.0))
        else:
            e[i] = resid[i] - (params[0] * resid[i-1] if i else 0.0)
    return float(e @ e)
fits = {}
for kind, start in (("arma11", [0.3, 0.0]), ("ar2", [0.3, 0.0]), ("ar1", [0.3])):
    best = minimize(css, start, args=(kind,), method="Nelder-Mead", options={"xatol": 1e-8, "fatol": 1e-12, "maxiter": 5000})
    k = len(start); aic = n * math.log(best.fun / n) + 2 * (k + 1)
    fits[kind] = best.x
    print(f"  {kind}: params {np.round(best.x, 4).tolist()}, CSS AIC {aic:.2f}")
phi, theta = fits["arma11"]
print(f"    p under fitted ARMA(1,1): {p_at(arma11(n, phi, theta, B)):.4f}")
a1, a2 = fits["ar2"]
e = rng.standard_normal((B, n + 200)); out = np.zeros_like(e)
for i in range(2, n + 200): out[:, i] = a1 * out[:, i - 1] + a2 * out[:, i - 2] + e[:, i]
print(f"    p under fitted AR(2): {p_at(out[:, 200:]):.4f}")

print("== 5. Residual dependence without the curvature: residuals of the 2012-hinge model")
x1 = basis(years, 2012); r1 = y - x1 @ np.linalg.lstsq(x1, y, rcond=None)[0]
print(f"  lag-1 autocorrelation of hinge residuals {np.corrcoef(r1[1:], r1[:-1])[0,1]:.4f}; "
      f"of 1970-2012 straight-line residuals {null_rho(y[years<=2012], years[years<=2012]):.4f}")

print("== 6. Is the 2022-to-2025 change what an acceleration would produce anyway?")
# Truth: the bundle's own fixed-knot estimate (0.232 C/decade more from 2015) plus AR(1) noise
# with the hinge model's residual autocorrelation and scale. Compare each endpoint's search p,
# calibrated as the bundle does (plug-in rho from straight-line residuals).
x2 = basis(years, 2015); b2 = np.linalg.lstsq(x2, y, rcond=None)[0]; r2 = y - x2 @ b2
rho_alt = float(r2[1:] @ r2[:-1] / (r2[:-1] @ r2[:-1])); sig = math.sqrt(np.mean((r2[1:] - rho_alt * r2[:-1]) ** 2))
signal = x2 @ b2
V = signal[None, :] + sig * ar1(n, rho_alt, 1000)
crit = {}
res = {2022: [], 2025: []}
for end in (2022, 2025):
    m = years <= end; s_end = scanner(years[m], knots_for(end))
    S = s_end(V[:, m]); R = np.clip(rho_batch(V[:, m], years[m]), -0.95, 0.95)
    for i in range(len(S)):
        key = (end, round(R[i], 2))
        if key not in crit: crit[key] = s_end(ar1(m.sum(), round(R[i], 2), 2000))
        res[end].append((1 + np.count_nonzero(crit[key] >= S[i])) / (len(crit[key]) + 1))
p22, p25 = np.array(res[2022]), np.array(res[2025])
print(f"  alternative: delta {b2[2]:.3f} C/decade from 2015, AR(1) rho {rho_alt:.3f}")
print(f"  P(p through 2022 > 0.4) = {np.mean(p22 > 0.4):.3f}; median p through 2022 = {np.median(p22):.3f}")
print(f"  P(p through 2025 < 0.05) = {np.mean(p25 < 0.05):.3f}; median p through 2025 = {np.median(p25):.3f}")
print(f"  P(p2022 > 0.4 and p2025 < 0.05) = {np.mean((p22 > 0.4) & (p25 < 0.05)):.3f}")
