#!/usr/bin/env python3
from __future__ import annotations

import argparse
import json
import sys
from pathlib import Path
import tkinter as tk
from tkinter import filedialog, messagebox

import matplotlib.pyplot as plt
import pandas as pd
from scipy.stats import mannwhitneyu
from statsmodels.nonparametric.smoothers_lowess import lowess


# ---------------------------------------------------------------------
# Defaults
# ---------------------------------------------------------------------

DATE_DEFAULT = "2026-04-02"
BAND_LABEL_DEFAULT = "20m"
BAND_MIN_HZ_DEFAULT = 14_095_000
BAND_MAX_HZ_DEFAULT = 14_099_000
WINDOW_DEFAULT_MIN = 15
LOWESS_FRAC_DEFAULT = 0.30
TOP_RX_PREVIEW = 10
MAIN_CLUSTER_HALF_WIDTH_HZ = 50


# ---------------------------------------------------------------------
# CLI
# ---------------------------------------------------------------------

def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(
        description=(
            "Analyse WSPR reports from a PSK Reporter JSON dump and build "
            "windowed median-offset / IQR trend outputs. "
            "Run with --gui for an interactive picker."
        )
    )
    parser.add_argument(
        "json_file",
        nargs="?",
        default="psk_m7sqi.json",
        help="Input JSON file (default: psk_m7sqi.json in current directory)",
    )
    parser.add_argument(
        "--gui",
        action="store_true",
        help="Launch tkinter GUI instead of running directly from CLI.",
    )
    parser.add_argument(
        "--start-utc",
        default=f"{DATE_DEFAULT} 00:00",
        help=f"UTC start datetime in 'YYYY-MM-DD HH:MM' form "
             f"(default: {DATE_DEFAULT} 00:00)",
    )
    parser.add_argument(
        "--end-utc",
        default=f"{DATE_DEFAULT} 23:59",
        help=f"UTC end datetime in 'YYYY-MM-DD HH:MM' form "
             f"(default: {DATE_DEFAULT} 23:59)",
    )
    parser.add_argument(
        "--band-label",
        default=BAND_LABEL_DEFAULT,
        help=f"Band label for titles/output (default: {BAND_LABEL_DEFAULT})",
    )
    parser.add_argument(
        "--band-min-hz",
        type=int,
        default=BAND_MIN_HZ_DEFAULT,
        help=f"Minimum frequency in Hz (default: {BAND_MIN_HZ_DEFAULT})",
    )
    parser.add_argument(
        "--band-max-hz",
        type=int,
        default=BAND_MAX_HZ_DEFAULT,
        help=f"Maximum frequency in Hz (default: {BAND_MAX_HZ_DEFAULT})",
    )
    parser.add_argument(
        "--window-minutes",
        type=int,
        default=WINDOW_DEFAULT_MIN,
        help=f"Window size in minutes (default: {WINDOW_DEFAULT_MIN})",
    )
    parser.add_argument(
        "--lowess-frac",
        type=float,
        default=LOWESS_FRAC_DEFAULT,
        help=f"LOWESS smoothing fraction (default: {LOWESS_FRAC_DEFAULT})",
    )
    parser.add_argument(
        "--top-rx",
        type=int,
        default=0,
        help=(
            "If > 0, also run a second analysis using only the N most frequent "
            "reporting stations for the selected time/frequency slice."
        ),
    )
    parser.add_argument(
        "--outdir",
        default="atmos_wobble_outputs",
        help="Output directory (default: atmos_wobble_outputs)",
    )
    return parser.parse_args()


# ---------------------------------------------------------------------
# Data loading / filtering
# ---------------------------------------------------------------------

def load_json(json_path: Path) -> pd.DataFrame:
    with json_path.open("r", encoding="utf-8") as fh:
        raw = json.load(fh)

    if not isinstance(raw, list):
        raise ValueError("JSON root must be a list of spot dictionaries.")

    df = pd.DataFrame(raw)
    required = {"ts", "rxCall", "freqHz", "mode"}
    missing = required - set(df.columns)
    if missing:
        raise ValueError(f"JSON is missing required fields: {sorted(missing)}")

    df = df.copy()
    df["time_utc"] = pd.to_datetime(df["ts"], unit="s", utc=True)
    df["freqHz"] = pd.to_numeric(df["freqHz"], errors="coerce")
    df["snr"] = pd.to_numeric(df.get("snr"), errors="coerce")
    df["rxCall"] = df["rxCall"].astype(str)
    df["mode"] = df["mode"].astype(str)
    df = df.dropna(subset=["freqHz", "time_utc"])
    return df.sort_values("time_utc").reset_index(drop=True)


def parse_utc_datetime(text: str) -> pd.Timestamp:
    ts = pd.Timestamp(text)
    if ts.tzinfo is None:
        ts = ts.tz_localize("UTC")
    else:
        ts = ts.tz_convert("UTC")
    return ts


def filter_wspr_band_timerange(
    df: pd.DataFrame,
    start_utc: str,
    end_utc: str,
    band_min_hz: int,
    band_max_hz: int,
) -> pd.DataFrame:
    start_ts = parse_utc_datetime(start_utc)
    end_ts = parse_utc_datetime(end_utc)

    if end_ts <= start_ts:
        raise ValueError("End time must be later than start time.")

    mask = (
        (df["mode"].str.upper() == "WSPR")
        & (df["freqHz"] >= band_min_hz)
        & (df["freqHz"] <= band_max_hz)
        & (df["time_utc"] >= start_ts)
        & (df["time_utc"] <= end_ts)
    )
    out = df.loc[mask].copy()
    if out.empty:
        raise ValueError(
            "No WSPR records found in the chosen time/frequency range."
        )
    return out


# ---------------------------------------------------------------------
# Analysis logic
# ---------------------------------------------------------------------

def build_windowed(
    df: pd.DataFrame,
    window_minutes: int,
    lowess_frac: float,
) -> tuple[pd.DataFrame, float]:
    # Keep ORIGINAL analysis logic:
    # 1) filter to main tone cluster using dataset median ±50 Hz
    # 2) define offsets relative to filtered median
    # 3) resample into time windows
    main_freq = df["freqHz"].median()
    df = df[
        (df["freqHz"] > main_freq - MAIN_CLUSTER_HALF_WIDTH_HZ)
        & (df["freqHz"] < main_freq + MAIN_CLUSTER_HALF_WIDTH_HZ)
    ].copy()

    if df.empty:
        raise ValueError("No points remain after main tone cluster filtering.")

    reference_hz = float(df["freqHz"].median())
    work = df.copy()
    work["offset_hz"] = work["freqHz"] - reference_hz
    work = work.set_index("time_utc")

    window = f"{window_minutes}min"
    grouped = work["offset_hz"].resample(window)

    windowed = pd.DataFrame(
        {
            "median_offset_hz": grouped.median(),
            "q25_hz": grouped.quantile(0.25),
            "q75_hz": grouped.quantile(0.75),
            "n": grouped.count(),
            "mean_offset_hz": grouped.mean(),
            "std_offset_hz": grouped.std(),
        }
    )
    windowed["iqr_hz"] = windowed["q75_hz"] - windowed["q25_hz"]
    windowed = windowed[windowed["n"] > 0].copy()

    if len(windowed) >= 3:
        valid = windowed["median_offset_hz"].notna()
        x = windowed.index.view("int64")[valid]
        y = windowed.loc[valid, "median_offset_hz"].to_numpy()
        smoothed = lowess(y, x, frac=lowess_frac, return_sorted=False)
        windowed.loc[valid, "lowess_median_offset_hz"] = smoothed
    else:
        windowed["lowess_median_offset_hz"] = windowed["median_offset_hz"]

    windowed["rolling3_median_offset_hz"] = (
        windowed["median_offset_hz"]
        .rolling(3, center=True, min_periods=1)
        .mean()
    )

    return windowed.reset_index(), reference_hz


def summary_stats(windowed: pd.DataFrame) -> dict[str, float | int | None]:
    out: dict[str, float | int | None] = {}
    out["windows"] = int(len(windowed))
    out["points"] = int(windowed["n"].sum())
    out["median_n_per_window"] = float(windowed["n"].median())
    out["mean_n_per_window"] = float(windowed["n"].mean())
    out["median_iqr_hz"] = float(windowed["iqr_hz"].median())
    out["max_iqr_hz"] = float(windowed["iqr_hz"].max())
    out["overall_median_of_window_medians_hz"] = float(
        windowed["median_offset_hz"].median()
    )

    morning = windowed[
        (windowed["time_utc"].dt.hour >= 10) & (windowed["time_utc"].dt.hour < 13)
    ]["median_offset_hz"].dropna()
    afternoon = windowed[
        (windowed["time_utc"].dt.hour >= 14) & (windowed["time_utc"].dt.hour < 18)
    ]["median_offset_hz"].dropna()

    out["morning_windows"] = int(len(morning))
    out["afternoon_windows"] = int(len(afternoon))
    out["morning_mean_hz"] = float(morning.mean()) if len(morning) else None
    out["afternoon_mean_hz"] = float(afternoon.mean()) if len(afternoon) else None

    if len(morning) >= 2 and len(afternoon) >= 2:
        stat = mannwhitneyu(morning, afternoon, alternative="two-sided")
        out["mannwhitney_u"] = float(stat.statistic)
        out["mannwhitney_p"] = float(stat.pvalue)
    else:
        out["mannwhitney_u"] = None
        out["mannwhitney_p"] = None

    if len(windowed) >= 2:
        out["lag1_autocorr"] = float(windowed["median_offset_hz"].autocorr(lag=1))
    else:
        out["lag1_autocorr"] = None

    return out


def reporters_table(df: pd.DataFrame) -> pd.DataFrame:
    rep = (
        df.groupby("rxCall", dropna=False)
        .agg(
            spots=("rxCall", "size"),
            first_seen_utc=("time_utc", "min"),
            last_seen_utc=("time_utc", "max"),
            median_freq_hz=("freqHz", "median"),
            median_snr_db=("snr", "median"),
        )
        .sort_values(["spots", "rxCall"], ascending=[False, True])
        .reset_index()
    )
    return rep


# ---------------------------------------------------------------------
# Plotting / outputs
# ---------------------------------------------------------------------

def plot_windowed(
    windowed: pd.DataFrame,
    stats: dict[str, float | int | None],
    out_png: Path,
    title_suffix: str,
    lowess_frac: float,
) -> None:
    fig, (ax1, ax2) = plt.subplots(
        2, 1, figsize=(13, 8), sharex=True, constrained_layout=True
    )

    ax1.scatter(
        windowed["time_utc"],
        windowed["median_offset_hz"],
        s=30,
        label="Window median offset",
        zorder=3,
    )
    ax1.plot(
        windowed["time_utc"],
        windowed["rolling3_median_offset_hz"],
        linewidth=1.6,
        label="Rolling mean (3 windows)",
    )
    ax1.plot(
        windowed["time_utc"],
        windowed["lowess_median_offset_hz"],
        linewidth=2.2,
        label=f"LOWESS (frac={lowess_frac:.2f})",
    )
    ax1.axhline(0, linestyle="--", linewidth=1)
    ax1.set_ylabel("Median offset (Hz)")
    ax1.set_title(f"Median offset vs daily reference — {title_suffix}")
    ax1.legend(loc="best")
    ax1.grid(True, alpha=0.3)

    ax2.bar(windowed["time_utc"], windowed["iqr_hz"], width=0.008, label="IQR")
    ax2.plot(windowed["time_utc"], windowed["n"], linewidth=1.5, label="n per window")
    ax2.set_ylabel("IQR / n")
    ax2.set_title(
        "Window spread and sample size"
        f" | windows={stats['windows']} points={stats['points']} "
        f"median n={stats['median_n_per_window']:.1f}"
    )
    ax2.legend(loc="best")
    ax2.grid(True, alpha=0.3)

    fig.savefig(out_png, dpi=150)
    plt.close(fig)


def write_summary(
    out_txt: Path,
    json_path: Path,
    filtered: pd.DataFrame,
    windowed: pd.DataFrame,
    stats: dict[str, float | int | None],
    reporters: pd.DataFrame,
    reference_hz: float,
    args_like: argparse.Namespace,
    title: str,
) -> None:
    with out_txt.open("w", encoding="utf-8") as fh:
        fh.write(f"Analysis: {title}\n")
        fh.write(f"Input file: {json_path}\n")
        fh.write(f"UTC start: {args_like.start_utc}\n")
        fh.write(f"UTC end: {args_like.end_utc}\n")
        fh.write(
            f"Band: {args_like.band_label} "
            f"({args_like.band_min_hz} to {args_like.band_max_hz} Hz)\n"
        )
        fh.write(f"Window: {args_like.window_minutes} minutes\n")
        fh.write(f"LOWESS frac: {args_like.lowess_frac}\n")
        fh.write(
            f"Main tone cluster filter: median ±{MAIN_CLUSTER_HALF_WIDTH_HZ} Hz\n"
        )
        fh.write(f"Daily median reference frequency: {reference_hz:.3f} Hz\n\n")

        fh.write("Top reporters by spot count:\n")
        fh.write(reporters.head(TOP_RX_PREVIEW).to_string(index=False))
        fh.write("\n\n")

        fh.write("Headline statistics:\n")
        for key, value in stats.items():
            fh.write(f"- {key}: {value}\n")
        fh.write("\n")

        fh.write("Interpretation notes:\n")
        fh.write(
            "- median_offset_hz is relative to the filtered median frequency "
            "of the selected WSPR records.\n"
        )
        fh.write(
            "- iqr_hz measures the within-window spread; large IQR means more "
            "frequency scatter.\n"
        )
        fh.write(
            "- LOWESS is the smoothed line through the window medians.\n"
        )
        fh.write(
            "- morning/afternoon comparison uses a Mann-Whitney U test because "
            "small samples are common.\n\n"
        )

        fh.write("Windowed data preview:\n")
        fh.write(windowed.head(20).to_string(index=False))
        fh.write("\n")


# ---------------------------------------------------------------------
# Runner
# ---------------------------------------------------------------------

def run_one_analysis(
    df: pd.DataFrame,
    outdir: Path,
    base_name: str,
    args_like: argparse.Namespace,
    title_suffix: str,
    json_path: Path,
) -> None:
    filtered = df.copy()
    reporters = reporters_table(filtered)
    windowed, reference_hz = build_windowed(
        filtered,
        args_like.window_minutes,
        args_like.lowess_frac,
    )

    stats = summary_stats(windowed)

    csv_path = outdir / f"{base_name}_windowed.csv"
    reporters_csv = outdir / f"{base_name}_reporters.csv"
    plot_path = outdir / f"{base_name}_trend.png"
    summary_path = outdir / f"{base_name}_summary.txt"

    windowed.to_csv(csv_path, index=False)
    reporters.to_csv(reporters_csv, index=False)
    plot_windowed(windowed, stats, plot_path, title_suffix, args_like.lowess_frac)
    write_summary(
        summary_path,
        json_path.resolve(),
        filtered,
        windowed,
        stats,
        reporters,
        reference_hz,
        args_like,
        title_suffix,
    )

    print(f"\n[{title_suffix}]")
    print(f"Filtered spots: {len(filtered)}")
    print(f"Distinct reporters: {filtered['rxCall'].nunique()}")
    print(f"Filtered median reference frequency: {reference_hz:.3f} Hz")
    print(f"Window CSV:      {csv_path}")
    print(f"Reporters CSV:   {reporters_csv}")
    print(f"Trend plot:      {plot_path}")
    print(f"Summary text:    {summary_path}")
    print("Top reporters:")
    print(reporters.head(min(TOP_RX_PREVIEW, len(reporters))).to_string(index=False))


# ---------------------------------------------------------------------
# GUI
# ---------------------------------------------------------------------

class GuiArgs:
    """Simple args-like object so we can reuse the same analysis functions."""
    def __init__(
        self,
        json_file: str,
        start_utc: str,
        end_utc: str,
        band_label: str,
        band_min_hz: int,
        band_max_hz: int,
        window_minutes: int,
        lowess_frac: float,
        top_rx: int,
        outdir: str,
    ) -> None:
        self.json_file = json_file
        self.start_utc = start_utc
        self.end_utc = end_utc
        self.band_label = band_label
        self.band_min_hz = band_min_hz
        self.band_max_hz = band_max_hz
        self.window_minutes = window_minutes
        self.lowess_frac = lowess_frac
        self.top_rx = top_rx
        self.outdir = outdir


def launch_gui() -> None:
    root = tk.Tk()
    root.title("Atmospheric Wobble Analysis")

    # Vars
    json_var = tk.StringVar(value="psk_m7sqi.json")
    outdir_var = tk.StringVar(value="atmos_wobble_outputs")
    start_var = tk.StringVar(value=f"{DATE_DEFAULT} 00:00")
    end_var = tk.StringVar(value=f"{DATE_DEFAULT} 23:59")
    band_label_var = tk.StringVar(value=BAND_LABEL_DEFAULT)
    band_min_var = tk.StringVar(value=str(BAND_MIN_HZ_DEFAULT / 1e6))
    band_max_var = tk.StringVar(value=str(BAND_MAX_HZ_DEFAULT / 1e6))
    window_var = tk.StringVar(value=str(WINDOW_DEFAULT_MIN))
    lowess_var = tk.StringVar(value=str(LOWESS_FRAC_DEFAULT))
    top_rx_var = tk.StringVar(value="0")

    row = 0

    def add_label_entry(label: str, var: tk.StringVar, width: int = 40) -> None:
        nonlocal row
        tk.Label(root, text=label, anchor="w").grid(row=row, column=0, sticky="w", padx=6, pady=4)
        tk.Entry(root, textvariable=var, width=width).grid(row=row, column=1, sticky="we", padx=6, pady=4)
        row += 1

    def browse_json() -> None:
        path = filedialog.askopenfilename(
            title="Select input JSON",
            filetypes=[("JSON files", "*.json"), ("All files", "*.*")],
        )
        if path:
            json_var.set(path)

    def browse_outdir() -> None:
        path = filedialog.askdirectory(title="Select output directory")
        if path:
            outdir_var.set(path)

    tk.Label(root, text="Input JSON", anchor="w").grid(row=row, column=0, sticky="w", padx=6, pady=4)
    tk.Entry(root, textvariable=json_var, width=40).grid(row=row, column=1, sticky="we", padx=6, pady=4)
    tk.Button(root, text="Browse...", command=browse_json).grid(row=row, column=2, padx=6, pady=4)
    row += 1

    tk.Label(root, text="Output directory", anchor="w").grid(row=row, column=0, sticky="w", padx=6, pady=4)
    tk.Entry(root, textvariable=outdir_var, width=40).grid(row=row, column=1, sticky="we", padx=6, pady=4)
    tk.Button(root, text="Browse...", command=browse_outdir).grid(row=row, column=2, padx=6, pady=4)
    row += 1

    add_label_entry("UTC start (YYYY-MM-DD HH:MM)", start_var)
    add_label_entry("UTC end (YYYY-MM-DD HH:MM)", end_var)
    add_label_entry("Band label", band_label_var)
    add_label_entry("Band min MHz", band_min_var)
    add_label_entry("Band max MHz", band_max_var)
    add_label_entry("Window minutes", window_var)
    add_label_entry("LOWESS frac", lowess_var)
    add_label_entry("Top RX (0 = off)", top_rx_var)

    status = tk.StringVar(value="Ready.")

    def run_analysis() -> None:
        try:
            args_like = GuiArgs(
                json_file=json_var.get().strip(),
                start_utc=start_var.get().strip(),
                end_utc=end_var.get().strip(),
                band_label=band_label_var.get().strip(),
                band_min_hz=int(float(band_min_var.get().strip()) * 1e6),
                band_max_hz=int(float(band_max_var.get().strip()) * 1e6),
                window_minutes=int(window_var.get().strip()),
                lowess_frac=float(lowess_var.get().strip()),
                top_rx=int(top_rx_var.get().strip()),
                outdir=outdir_var.get().strip(),
            )

            json_path = Path(args_like.json_file).expanduser().resolve()
            outdir = Path(args_like.outdir).expanduser().resolve()
            outdir.mkdir(parents=True, exist_ok=True)

            status.set("Loading JSON...")
            root.update_idletasks()
            df = load_json(json_path)

            status.set("Filtering selected range...")
            root.update_idletasks()
            filtered = filter_wspr_band_timerange(
                df,
                args_like.start_utc,
                args_like.end_utc,
                args_like.band_min_hz,
                args_like.band_max_hz,
            )

            # Safe-ish base name
            start_tag = args_like.start_utc.replace(":", "").replace(" ", "_")
            end_tag = args_like.end_utc.replace(":", "").replace(" ", "_")
            base_name = f"{args_like.band_label}_{start_tag}_to_{end_tag}"

            status.set("Running main analysis...")
            root.update_idletasks()
            run_one_analysis(
                filtered,
                outdir,
                base_name=base_name,
                args_like=args_like,
                title_suffix=(
                    f"{args_like.band_label} WSPR reports | "
                    f"{args_like.start_utc} to {args_like.end_utc} UTC"
                ),
                json_path=json_path,
            )

            if args_like.top_rx > 0:
                reps = reporters_table(filtered)
                chosen = reps.head(args_like.top_rx)["rxCall"].tolist()
                top_df = filtered[filtered["rxCall"].isin(chosen)].copy()

                status.set("Running top-RX subset analysis...")
                root.update_idletasks()
                run_one_analysis(
                    top_df,
                    outdir,
                    base_name=f"{base_name}_top_{args_like.top_rx}_rx",
                    args_like=args_like,
                    title_suffix=(
                        f"Top {args_like.top_rx} reporters only | "
                        f"{args_like.band_label} | "
                        f"{args_like.start_utc} to {args_like.end_utc} UTC"
                    ),
                    json_path=json_path,
                )

            status.set(f"Done. Outputs in: {outdir}")
            messagebox.showinfo("Analysis complete", f"Outputs written to:\n{outdir}")

        except Exception as exc:
            status.set("Error.")
            messagebox.showerror("Error", str(exc))

    tk.Button(root, text="Run analysis", command=run_analysis).grid(
        row=row, column=0, columnspan=3, pady=10
    )
    row += 1

    tk.Label(root, textvariable=status, anchor="w", fg="blue").grid(
        row=row, column=0, columnspan=3, sticky="w", padx=6, pady=4
    )

    root.columnconfigure(1, weight=1)
    root.mainloop()


# ---------------------------------------------------------------------
# Main
# ---------------------------------------------------------------------

def main() -> int:
    args = parse_args()

    if args.gui:
        launch_gui()
        return 0

    json_path = Path(args.json_file).expanduser().resolve()
    outdir = Path(args.outdir).expanduser().resolve()
    outdir.mkdir(parents=True, exist_ok=True)

    try:
        df = load_json(json_path)
        filtered = filter_wspr_band_timerange(
            df,
            args.start_utc,
            args.end_utc,
            args.band_min_hz,
            args.band_max_hz,
        )
    except Exception as exc:
        print(f"Error: {exc}", file=sys.stderr)
        return 1

    start_tag = args.start_utc.replace(":", "").replace(" ", "_")
    end_tag = args.end_utc.replace(":", "").replace(" ", "_")
    base_name = f"{args.band_label}_{start_tag}_to_{end_tag}"

    run_one_analysis(
        filtered,
        outdir,
        base_name=base_name,
        args_like=args,
        title_suffix=(
            f"{args.band_label} WSPR reports | "
            f"{args.start_utc} to {args.end_utc} UTC"
        ),
        json_path=json_path,
    )

    if args.top_rx and args.top_rx > 0:
        reps = reporters_table(filtered)
        chosen = reps.head(args.top_rx)["rxCall"].tolist()
        top_df = filtered[filtered["rxCall"].isin(chosen)].copy()
        run_one_analysis(
            top_df,
            outdir,
            base_name=f"{base_name}_top_{args.top_rx}_rx",
            args_like=args,
            title_suffix=(
                f"Top {args.top_rx} reporters only | "
                f"{args.band_label} | "
                f"{args.start_utc} to {args.end_utc} UTC"
            ),
            json_path=json_path,
        )

    return 0


if __name__ == "__main__":
    raise SystemExit(main())
