#!/usr/bin/env python3
import calendar, hashlib, json, math, shutil, zipfile
from pathlib import Path
import numpy as np
import pandas as pd
from scipy import stats
import matplotlib.pyplot as plt

ROOT=Path(__file__).resolve().parent
RAW=ROOT/"raw"
station=RAW/"USW00094728.dly"

def sha256(p):
    h=hashlib.sha256()
    with open(p,"rb") as f:
        for b in iter(lambda:f.read(1<<20),b""): h.update(b)
    return h.hexdigest()

rows=[]
for line in station.read_text().splitlines():
    if len(line)<269 or line[:11]!="USW00094728": continue
    y,m,el=int(line[11:15]),int(line[15:17]),line[17:21]
    if el not in ("TMIN","TMAX"): continue
    for d in range(1,calendar.monthrange(y,m)[1]+1):
        i=21+(d-1)*8
        v=int(line[i:i+5]); q=line[i+6:i+7].strip()
        rows.append((pd.Timestamp(y,m,d),el,np.nan if v==-9999 else v/10,q))
long=pd.DataFrame(rows,columns=["date","element","value_c","qflag"])
wide=long.pivot(index="date",columns="element",values="value_c")
q=long.pivot(index="date",columns="element",values="qflag").fillna("")
annual=[]; audits=[]
for y in range(1991,wide.index.max().year+1):
    expected=366 if calendar.isleap(y) else 365
    s=wide[wide.index.year==y]; sq=q[q.index.year==y]
    n_missing=int(s[["TMIN","TMAX"]].isna().sum().sum()) if len(s) else 2*expected
    n_qflag=int((sq[["TMIN","TMAX"]]!="").sum().sum()) if len(s) else 0
    complete=len(s)==expected and n_missing==0 and n_qflag==0
    audits.append(dict(year=y,expected_days=expected,rows=len(s),missing_values=n_missing,quality_flags=n_qflag,complete=complete))
    if complete:
        daily=(s.TMAX+s.TMIN)/2
        annual.append(dict(year=y,days=expected,annual_mean_temp_c=daily.mean(),
            hot10_mean_tmax_c=s.TMAX.nlargest(10).mean(),annual_max_tmax_c=s.TMAX.max()))
annual=pd.DataFrame(annual); audit=pd.DataFrame(audits)

def fit(col):
    x=annual.year.to_numpy(float); y=annual[col].to_numpy(float); n=len(y)
    X=np.column_stack([np.ones(n),x]); inv=np.linalg.inv(X.T@X); beta=inv@X.T@y; resid=y-X@beta
    lag=max(1,math.floor(4*(n/100)**(2/9)))
    meat=np.zeros((2,2))
    for t in range(n): meat += resid[t]**2*np.outer(X[t],X[t])
    for L in range(1,lag+1):
        w=1-L/(lag+1); G=np.zeros((2,2))
        for t in range(L,n): G += resid[t]*resid[t-L]*np.outer(X[t],X[t-L])
        meat += w*(G+G.T)
    cov=inv@meat@inv; se=math.sqrt(cov[1,1]); z=stats.norm.ppf(.975)
    ts=stats.theilslopes(y,x,alpha=.95)
    rho=np.corrcoef(resid[1:],resid[:-1])[0,1]
    dw=np.sum(np.diff(resid)**2)/np.sum(resid**2)
    sst=np.sum((y-y.mean())**2); r2=1-np.sum(resid**2)/sst
    return dict(metric=col,n=n,nw_lag=lag,ols_slope_c_decade=beta[1]*10,
      nw_se_c_decade=se*10,nw_ci_low=(beta[1]-z*se)*10,nw_ci_high=(beta[1]+z*se)*10,
      nw_p_two_sided=2*stats.norm.sf(abs(beta[1]/se)),r2=r2,residual_lag1=rho,durbin_watson=dw,
      theil_sen_slope_c_decade=ts.slope*10,theil_sen_ci_low=ts.low_slope*10,theil_sen_ci_high=ts.high_slope*10)
fits=pd.DataFrame([fit("annual_mean_temp_c"),fit("hot10_mean_tmax_c")])

annual.to_csv(ROOT/"annual_metrics.csv",index=False,float_format="%.4f")
audit.to_csv(ROOT/"quality_checks.csv",index=False)
fits.to_csv(ROOT/"trend_results.csv",index=False,float_format="%.5f")
checks={"station":"USW00094728","latest_date_in_file":str(wide.index.max().date()),
 "complete_years":annual.year.astype(int).tolist(),"excluded_years":audit.loc[~audit.complete,"year"].astype(int).tolist(),
 "all_complete_years_have_no_missing_or_qflags":bool((audit.loc[audit.complete,["missing_values","quality_flags"]]==0).all().all()),
 "raw_sha256":sha256(station)}
(ROOT/"quality_checks.json").write_text(json.dumps(checks,indent=2)+"\n")

fig,axs=plt.subplots(2,1,figsize=(9,7),sharex=True)
for ax,col,title,color in [(axs[0],"annual_mean_temp_c","Annual mean temperature","#2774ae"),
 (axs[1],"hot10_mean_tmax_c","Mean of ten hottest daily maxima","#c43c35")]:
    x=annual.year.to_numpy(); y=annual[col].to_numpy(); f=fits[fits.metric==col].iloc[0]
    ax.scatter(x,y,color=color,s=28,label="Annual value")
    ax.plot(x,np.polyval(np.polyfit(x,y,1),x),color="black",lw=1.6,label="OLS")
    ts=stats.theilslopes(y,x); ax.plot(x,ts.intercept+ts.slope*x,color="#2a8c4a",lw=1.6,ls="--",label="Theil–Sen")
    ax.set_ylabel("°C"); ax.grid(alpha=.25)
    ax.set_title(f"{title}: NW OLS {f.ols_slope_c_decade:+.2f}, Theil–Sen {f.theil_sen_slope_c_decade:+.2f} °C/decade")
axs[0].legend(frameon=False,ncol=3); axs[1].set_xlabel("Year")
fig.suptitle("Central Park GHCN-Daily, complete years 1991–2025",y=.995)
fig.tight_layout(); fig.savefig(ROOT/"trend_robustness.png",dpi=180,bbox_inches="tight"); plt.close(fig)
