#!/usr/bin/env python
# mercury_drift_pipeline.py
"""
End-to-end pipeline:
  1. load_drift_data()  -> pandas DataFrame with ['RA','DEC','e']
  2. notch filter 6 harmonics over RA/DEC
  3. Lomb-Scargle power spectrum + window leakage plot
  4. sliding-window linear drift  -> dot_alpha, sigma
  5. Bayesian hierarchical fit of  b1  &  dGR
  6. prior sensitivity sweep
Outputs:
  ./figs/power_spectrum.png
  ./results/posterior_summary.csv
  ./results/sensitivity_summary.csv
"""
import os, json, argparse, itertools
import numpy as np, pandas as pd, matplotlib.pyplot as plt
from pathlib import Path
from scipy.signal import iirnotch, filtfilt, lombscargle, windows
from scipy.stats import linregress
import pymc as pm, arviz as az
from tqdm import tqdm

# ----------  CONFIG  ----------
SAVE_DIR_FIG = Path("figs")
SAVE_DIR_RES = Path("results")
SAVE_DIR_FIG.mkdir(exist_ok=True)
SAVE_DIR_RES.mkdir(exist_ok=True)

# ----------  UTILS ----------
def notch_filter_series(series, fs, f0, harmonics=np.arange(1, 7), bw_factor=0.0006):
    """Return notch-filtered version of `series` (numpy array)."""
    b, a = np.array([1.0]), np.array([1.0])
    for k in harmonics:
        bw   = bw_factor * k
        bi, ai = iirnotch(k * f0, k / bw, fs=fs)
        b, a   = np.convolve(b, bi), np.convolve(a, ai)
    return filtfilt(b, a, series)

# ----------  1  LOAD DATA ----------
def load_or_fetch():
    """用户如果已有 drift_utils.py，可直接 from drift_utils import load_drift_data"""
    try:
        from drift_utils import load_drift_data
        return load_drift_data()
    except ImportError:
        raise SystemExit("找不到 drift_utils.py，请先按照前述示例创建该模块。")

# ----------  2  POWER-SPECTRUM PLOT ----------
def plot_power(df, P_days=87.969):
    P  = P_days * 24*3600.0
    f0 = 1 / P

    # 采样率 (Hz)
    fs = 1 / np.median(np.diff(df.index).astype("timedelta64[s]").astype(float))

    # notch-filter
    for col in ("dRA", "dDEC"):
        df[f"{col}_filt"] = notch_filter_series(df[col].to_numpy(), fs, f0)

    # Lomb-Scargle
    t = (df.index - df.index[0]).total_seconds().values
    freq = np.linspace(1e-9, 10 * f0, 20000)
    p_raw  = lombscargle(t, df["dRA"].to_numpy(),       freq)
    p_filt = lombscargle(t, df["dRA_filt"].to_numpy(),  freq)
    p_win  = lombscargle(t, windows.hann(len(t)), freq)

    plt.figure(figsize=(8, 5))
    plt.loglog(freq / f0, p_raw,  label="raw",            lw=0.8)
    plt.loglog(freq / f0, p_filt, label="notch-filtered", lw=0.8)
    plt.loglog(freq / f0, p_win,  label="window", ls="--", lw=0.8, alpha=0.7)
    plt.xlabel(r"frequency / $f_0$")
    plt.ylabel("Lomb–Scargle power")
    plt.legend(); plt.tight_layout()
    out = SAVE_DIR_FIG / "power_spectrum.png"
    plt.savefig(out, dpi=300)
    plt.close()
    print(f"[✓] power spectrum saved → {out}")
    return f0  # 后面滑窗用

# ----------  3  滑窗年漂移 ----------
def compute_drift(df, window_days=180):
    ds = window_days
    slope, sigma = [], []
    times = df.index
    arr   = df["dRA_filt"].to_numpy()
    for i in range(len(times)):
        lo, hi = max(0, i - ds), min(len(times), i + ds + 1)
        win_series = arr[lo:hi]
        t_win = (times[lo:hi] - times[lo]).total_seconds().astype(float) / 86400.0
        res = linregress(t_win, win_series)
        slope.append(res.slope)
        sigma.append(res.stderr)
    df["dot_alpha"] = np.array(slope) * 365.25   # arcsec/yr
    df["sigma"]     = np.array(sigma) * 365.25
    return df.dropna(subset=["dot_alpha"])

# ----------  4  BAYESIAN  MODEL ----------
def bayes_fit(df, prior_sigma_b1=0.5, prior_sigma_gr=0.1, seed=2025, tune=1000, draws=2000):
    f_e = 1 / np.power(1 - df["e"].to_numpy()**2, 1.5)
    y   = df["dot_alpha"].to_numpy()
    yerr= df["sigma"].to_numpy()
    with pm.Model() as mdl:
        b1  = pm.Normal("b1", 0, prior_sigma_b1)
        dGR = pm.Normal("dGR", 0.43, prior_sigma_gr)
        mu  = b1 * f_e + dGR
        pm.Normal("obs", mu, yerr, observed=y)
        trace = pm.sample(draws, tune=tune, target_accept=0.95,
                          chains=2, cores=2,
                          random_seed=seed, progressbar=False)
    return az.summary(trace, var_names=["b1", "dGR"])

# ----------  5  PRIOR SENSITIVITY SWEEP ----------
def sensitivity_table(df):
    configs = [
        ("N(0,0.5) + N(0.43,0.1)", 0.5, 0.1),
        ("N(0,1.0) + N(0.43,0.2)", 1.0, 0.2),
        ("N(0,0.2) + N(0.43,0.05)",0.2, 0.05),
    ]
    frames=[]
    for label, sig_b, sig_gr in tqdm(configs, desc="sensitivity"):
        summ = bayes_fit(df, sig_b, sig_gr)
        summ["prior"] = label
        frames.append(summ)
    return pd.concat(frames, keys=[c[0] for c in configs])

# ----------  MAIN ----------
def main():
    parser = argparse.ArgumentParser(description="Mercury drift analysis pipeline")
    parser.add_argument("--no-bayes", action="store_true",
                        help="skip Bayesian inference (only do notch & plot)")
    args = parser.parse_args()

    df = load_or_fetch()
    # baseline demean
    for col in ("RA_deg", "DEC_deg"):
        centered = df[col] - df[col].mean()
        df["d"+col.split('_')[0]] = centered

    f0 = plot_power(df)
    df = compute_drift(df)

    if args.no_bayes:
        return

    # baseline fit
    baseline = bayes_fit(df)
    baseline.to_csv(SAVE_DIR_RES / "posterior_summary.csv")
    print(f"[✓] posterior summary → {SAVE_DIR_RES/'posterior_summary.csv'}")

    # sensitivity sweep
    sens = sensitivity_table(df)
    sens.to_csv(SAVE_DIR_RES / "sensitivity_summary.csv")
    print(f"[✓] sensitivity table → {SAVE_DIR_RES/'sensitivity_summary.csv'}")

if __name__ == "__main__":
    main()
