"""Offline, registered endpoint and noise-model audit of a frozen NASA series."""
import csv
import hashlib
import json
import math
from pathlib import Path

import numpy as np
from scipy.stats import norm, t
from statsmodels.regression.linear_model import OLS
from statsmodels.stats.sandwich_covariance import cov_hac
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt

ROOT = Path(__file__).resolve().parents[1]
OUT = ROOT / "results"
SEED = 20261006
B = 10000
MONTHS = ["Jan", "Feb", "Mar", "Apr", "May", "Jun", "Jul", "Aug", "Sep", "Oct", "Nov", "Dec"]


def portable(value):
    if isinstance(value, float):
        if not math.isfinite(value):
            raise ValueError("Nonfinite result")
        return float(format(value, ".10g"))
    if isinstance(value, dict):
        return {k: portable(v) for k, v in value.items()}
    if isinstance(value, list):
        return [portable(v) for v in value]
    return value


def read_series():
    path = ROOT / "data/gistemp.csv"
    provenance = json.loads((ROOT / "data/source.json").read_text())
    assert hashlib.sha256(path.read_bytes()).hexdigest() == provenance["sha256"]
    with path.open(newline="") as f:
        assert next(f).strip() == "Land-Ocean: Global Means"
        rows = list(csv.DictReader(f))
    years = [int(row["Year"]) for row in rows]
    assert len(years) == len(set(years))
    chosen = []
    for row in rows:
        year = int(row["Year"])
        if 1970 <= year <= 2025:
            monthly = [float(row[m]) for m in MONTHS]
            annual = float(row["J-D"])
            assert np.isfinite(monthly).all() and math.isfinite(annual)
            # Monthly and annual entries are separately rounded to hundredths.
            assert abs(np.mean(monthly) - annual) < 0.011
            chosen.append((year, annual))
    assert [r[0] for r in chosen] == list(range(1970, 2026))
    return np.array([r[0] for r in chosen]), np.array([r[1] for r in chosen])


def design(years, knot=None):
    columns = [np.ones(len(years)), (years - 1970) / 10]
    if knot is not None:
        columns.append(np.maximum(0, years - knot) / 10)
    return np.column_stack(columns)


def fixed_fit(years, y, lag):
    x = design(years, 2015)
    beta = np.linalg.lstsq(x, y, rcond=None)[0]
    residual = y - x @ beta
    inverse = np.linalg.inv(x.T @ x)
    scores = x * residual[:, None]
    meat = scores.T @ scores
    for j in range(1, lag + 1):
        cross = scores[j:].T @ scores[:-j]
        meat += (1 - j / (lag + 1)) * (cross + cross.T)
    covariance = inverse @ meat @ inverse * len(y) / (len(y) - 3)
    reference = OLS(y, x).fit()
    assert np.max(np.abs(covariance - cov_hac(reference, nlags=lag, use_correction=True))) < 1e-10
    assert np.max(np.abs(beta - reference.params)) < 1e-10
    se = math.sqrt(covariance[2, 2])
    quantile = norm.ppf(.975)
    iid_se = math.sqrt(np.dot(residual, residual) / (len(y) - 3) * inverse[2, 2])
    return {
        "n": len(y), "start_year": int(years[0]), "end_year": int(years[-1]),
        "knot_year": 2015, "HAC_lag": lag,
        "before_c_per_decade": float(beta[1]), "after_c_per_decade": float(beta[1] + beta[2]),
        "delta_c_per_decade": float(beta[2]), "se_delta_c_per_decade": se,
        "ci95_low_c_per_decade": float(beta[2] - quantile * se),
        "ci95_high_c_per_decade": float(beta[2] + quantile * se),
        "p_two_sided_normal": float(2 * norm.sf(abs(beta[2] / se))),
        "iid_se_delta_c_per_decade": iid_se,
        "iid_ci95_low_c_per_decade": float(beta[2] - t.ppf(.975, len(y)-3)*iid_se),
        "iid_ci95_high_c_per_decade": float(beta[2] + t.ppf(.975, len(y)-3)*iid_se),
        "iid_p_two_sided_t": float(2*t.sf(abs(beta[2]/iid_se), len(y)-3)),
        "residual_lag_one_correlation": float(np.corrcoef(residual[1:], residual[:-1])[0, 1]),
    }, beta


def scan_setup(years):
    x = design(years)
    q, _ = np.linalg.qr(x)
    knots = np.arange(1985, min(2015, int(years[-1])-10)+1)
    hinges = np.maximum(0, years[:, None] - knots[None, :])/10
    residual_hinges = hinges - q @ (q.T @ hinges)
    directions = residual_hinges / np.sqrt(np.sum(residual_hinges**2, axis=0))
    return q, directions, knots


def scan(values, q, directions):
    residual = values - (values @ q) @ q.T
    sse = np.sum(residual**2, axis=1)
    improvement = (residual @ directions)**2
    statistics = improvement / (sse[:, None] - improvement) * (values.shape[1]-3)
    indices = np.argmax(statistics, axis=1)
    return statistics[np.arange(len(values)), indices], indices


def scalar_scan(years, y, knots):
    x = design(years)
    residual = y - x @ np.linalg.lstsq(x, y, rcond=None)[0]
    sse0 = np.dot(residual, residual)
    values = []
    for knot in knots:
        full = design(years, knot)
        residual = y - full @ np.linalg.lstsq(full, y, rcond=None)[0]
        sse1 = np.dot(residual, residual)
        values.append((sse0-sse1)/sse1*(len(y)-3))
    return np.max(values), int(np.argmax(values))


def ar1_noise(rng, count, n, rho, sigma):
    normal = rng.standard_normal((count, n))
    noise = np.empty_like(normal)
    noise[:, 0] = normal[:, 0] * sigma / math.sqrt(1-rho*rho)
    for i in range(1, n):
        noise[:, i] = rho * noise[:, i-1] + sigma * normal[:, i]
    return noise


def wilson(successes, count):
    z = norm.ppf(.975)
    p = successes/count
    denominator = 1+z*z/count
    center = (p+z*z/(2*count))/denominator
    half = z*math.sqrt(p*(1-p)/count+z*z/(4*count*count))/denominator
    return [max(0., float(center-half)), min(1., float(center+half))]


def calibrated_test(years, y, rho_override=None, iid=False):
    x = design(years)
    beta = np.linalg.lstsq(x, y, rcond=None)[0]
    residual = y - x @ beta
    raw_rho = float(np.dot(residual[1:], residual[:-1])/np.dot(residual[:-1], residual[:-1]))
    rho = 0.0 if iid else float(np.clip(raw_rho, -.95, .95)) if rho_override is None else rho_override
    sigma = math.sqrt(np.dot(residual, residual)/(len(y)-2)) if iid else math.sqrt(np.mean((residual[1:]-rho*residual[:-1])**2))
    q, directions, knots = scan_setup(years)
    observed, best = scan(y[None, :], q, directions)
    observed = float(observed[0]); best = int(best[0])
    scalar, scalar_best = scalar_scan(years, y, knots)
    assert abs(scalar-observed) < 1e-9 and scalar_best == best
    # Reset the seed per calibration, providing common random numbers in comparisons.
    rng = np.random.Generator(np.random.PCG64(SEED))
    noise = ar1_noise(rng, B, len(y), rho, sigma)
    simulated, _ = scan(noise, q, directions)
    max_scalar_error = 0.0
    for j in range(10):
        scalar, scalar_best = scalar_scan(years, noise[j], knots)
        max_scalar_error = max(max_scalar_error, abs(scalar-float(simulated[j])))
        assert max_scalar_error < 1e-9
    exceed = int(np.count_nonzero(simulated >= observed))
    p = (exceed+1)/(B+1)
    knot = int(knots[best])
    selected_beta = np.linalg.lstsq(design(years, knot), y, rcond=None)[0]
    result = {
        "end_year": int(years[-1]), "n":len(y), "candidate_knots": [int(k) for k in knots],
        "selected_knot":knot, "max_F": observed,
        "selected_delta_c_per_decade":float(selected_beta[2]),
        "raw_null_rho":raw_rho, "null_rho":rho,
        "rho_was_clipped":bool(not iid and rho_override is None and rho != raw_rho),
        "innovation_sigma_c":sigma, "B":B, "seed":SEED, "exceedances":exceed,
        "p_selection_adjusted":p, "monte_carlo_se":math.sqrt(p*(1-p)/B),
        "exceedance_probability_wilson95":wilson(exceed,B),
        "null_maxF_95percentile":float(np.quantile(simulated, .95)),
        "scalar_statistic_max_abs_error":max_scalar_error,
    }
    return result, simulated, (q, directions, rho, sigma)


def main():
    OUT.mkdir(exist_ok=True)
    years, y = read_series()
    fixed, bootstrap = {}, {}
    for endpoint in [2025, 2024, 2022]:
        mask = years <= endpoint
        fixed[str(endpoint)] = {}
        for lag in [0,1,3,5]:
            summary, _ = fixed_fit(years[mask], y[mask], lag)
            fixed[str(endpoint)][str(lag)] = summary
        ar, sim, setup = calibrated_test(years[mask], y[mask])
        iid, _, _ = calibrated_test(years[mask], y[mask], iid=True)
        bootstrap[str(endpoint)] = {"AR1":ar, "iid":iid}
        if endpoint == 2025:
            primary_setup, primary_sim = setup, sim
    rho_sensitivity = {}
    for rho in [.2,.4,.6]:
        r, _, _ = calibrated_test(years, y, rho_override=rho)
        rho_sensitivity['rho_' + str(rho).replace('.', 'p')] = r
    q, directions, rho, sigma = primary_setup
    critical = float(np.quantile(primary_sim, .95))
    rng = np.random.Generator(np.random.PCG64(SEED+1))
    power = {}
    for delta in [.1,.2,.3]:
        noise = ar1_noise(rng, 5000, len(y), rho, sigma)
        values = noise + delta*np.maximum(0, years-2015)[None, :]/10
        statistic, _ = scan(values, q, directions)
        successes = int(np.count_nonzero(statistic > critical))
        probability = successes/5000
        power['delta_' + str(delta).replace('.', 'p')] = {
            "true_delta_c_per_decade":delta, "simulations":5000, "rejections":successes,
            "conditional_power":probability,
            "monte_carlo_se":math.sqrt(probability*(1-probability)/5000),
            "wilson95":wilson(successes,5000),
        }
    primary = fixed["2025"]["3"]
    display = {
        "primary_delta": round(primary['delta_c_per_decade'], 3),
        "primary_before": round(primary['before_c_per_decade'], 3),
        "primary_after": round(primary['after_c_per_decade'], 3),
        "primary_se": round(primary['se_delta_c_per_decade'], 3),
        "primary_low": round(primary['ci95_low_c_per_decade'], 3),
        "primary_high": round(primary['ci95_high_c_per_decade'], 3),
        "primary_p": format(primary['p_two_sided_normal'], '.3g'),
        "endpoints": {
            endpoint: {
                "delta":round(fixed[endpoint]['3']['delta_c_per_decade'],3),
                "low":round(fixed[endpoint]['3']['ci95_low_c_per_decade'],3),
                "high":round(fixed[endpoint]['3']['ci95_high_c_per_decade'],3),
                "fixed_p":format(fixed[endpoint]['3']['p_two_sided_normal'],'.3g'),
                "search_p":round(bootstrap[endpoint]['AR1']['p_selection_adjusted'],4),
                "iid_search_p":round(bootstrap[endpoint]['iid']['p_selection_adjusted'],4),
                "rho":round(bootstrap[endpoint]['AR1']['null_rho'],3),
                "mc_se":round(bootstrap[endpoint]['AR1']['monte_carlo_se'],4),
            } for endpoint in ['2025','2024','2022']
        },
        "rho_sensitivity": {
            key:round(value['p_selection_adjusted'],4) for key,value in rho_sensitivity.items()
        },
        "power_percent": {
            key:round(100*value['conditional_power'],1) for key,value in power.items()
        },
        "confidence_percent":95,
    }
    result = {
        "primary":primary, "fixed_knot_sensitivity":fixed, "bootstrap":bootstrap,
        "rho_sensitivity":rho_sensitivity, "conditional_power":power, "display":display,
        "checks":{"unique_complete_years":True, "HAC_matches_statsmodels":True,
                  "projection_matches_scalar_lstsq":True, "partial_2026_excluded":True},
    }
    (OUT/"R1.json").write_text(json.dumps(portable(result), indent=2)+'\n')
    with (OUT/"annual.csv").open('w',newline='') as f:
        w=csv.writer(f);w.writerow(['year','anomaly_c']);w.writerows(zip(years,y))
    fig, axes = plt.subplots(1,2,figsize=(11,4.2),layout='constrained')
    axes[0].plot(years,y,'o',markersize=3,label='NASA annual anomaly')
    for endpoint,color in [(2025,'#226a9b'),(2024,'#db863a'),(2022,'#43835f')]:
        mask=years<=endpoint;_,beta=fixed_fit(years[mask],y[mask],3)
        axes[0].plot(years[mask],design(years[mask],2015)@beta,color=color,label=f'Fit through {endpoint}')
    axes[0].axvline(2015,color='gray',linestyle=':',linewidth=1)
    axes[0].set(xlabel='Year',ylabel='Anomaly relative to 1951-1980 (°C)')
    axes[0].legend(fontsize=8)
    for j,endpoint in enumerate([2025,2024,2022]):
        r=fixed[str(endpoint)]['3'];center=r['delta_c_per_decade']
        axes[1].errorbar(center,j,xerr=[[center-r['ci95_low_c_per_decade']],[r['ci95_high_c_per_decade']-center]],fmt='o',capsize=4,color='#226a9b')
    axes[1].set(yticks=range(3),yticklabels=['Through 2025','Through 2024','Through 2022'],xlabel='Slope increase at fixed 2015 knot (°C/decade)')
    axes[1].axvline(0,color='gray',linestyle=':')
    axes[1].set_title('95% normal HAC intervals, lag 3',fontsize=10)
    fig.savefig(OUT/'endpoint-sensitivity.png',dpi=180)
    plt.close(fig)
    print(json.dumps(portable({"primary":primary,"bootstrap":bootstrap,"rho_sensitivity":rho_sensitivity,"conditional_power":power}),indent=2))


if __name__ == '__main__':
    main()
