#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
Robust Savitzky-Golay smoothing for data with spike-like outliers.

Workflow:
  1. Build a moderately smoothed reference curve with a Savitzky-Golay filter.
  2. Compute residuals: raw - reference.
  3. Estimate residual scatter by MAD: sigma = 1.4826 * median(|r - median(r)|).
  4. Detect outliers where |residual - median(residual)| > threshold * sigma.
  5. Replace only outlier points by interpolation from non-outlier data.
  6. Apply a final, shorter Savitzky-Golay filter for analysis/visualization.

Example:
  python smoothing_robust_sg_outlier.py xrd-lowSN.xlsx --sheet xrd --xcol x --ycol "y(raw)"
"""

from __future__ import annotations

import argparse
from pathlib import Path

import numpy as np
import pandas as pd
import matplotlib.pyplot as plt
from scipy.signal import savgol_filter


def make_odd_window(window: int, n: int, polyorder: int) -> int:
    """Return a valid odd SG window length for n data points."""
    if window < polyorder + 2:
        window = polyorder + 2
    if window % 2 == 0:
        window += 1
    if window > n:
        window = n if n % 2 == 1 else n - 1
    min_window = polyorder + 2
    if min_window % 2 == 0:
        min_window += 1
    if window < min_window:
        raise ValueError(f"Too few data points: need at least {min_window}, got {n}")
    return window


def robust_sigma_mad(residual: np.ndarray) -> tuple[float, float]:
    """Return robust center and robust sigma estimated from MAD."""
    center = np.nanmedian(residual)
    mad = np.nanmedian(np.abs(residual - center))
    sigma = 1.4826 * mad

    # Fallback for perfectly flat or tiny residual sets.
    if not np.isfinite(sigma) or sigma <= 0:
        sigma = np.nanstd(residual)
    if not np.isfinite(sigma) or sigma <= 0:
        sigma = 1.0
    return center, sigma


def detect_and_smooth(
    x: np.ndarray,
    y: np.ndarray,
    detect_window: int = 15,
    final_window: int = 7,
    polyorder: int = 2,
    threshold: float = 4.0,
    max_iter: int = 1,
) -> dict[str, np.ndarray | float | int]:
    """Detect spike outliers and return corrected/final-smoothed arrays."""
    x = np.asarray(x, dtype=float)
    y = np.asarray(y, dtype=float)

    good = np.isfinite(x) & np.isfinite(y)
    if good.sum() < polyorder + 3:
        raise ValueError("Not enough finite x/y data points.")

    # Sort by x to avoid interpolation surprises.
    order = np.argsort(x[good])
    xg = x[good][order]
    yg = y[good][order]

    n = len(yg)
    detect_window = make_odd_window(detect_window, n, polyorder)
    final_window = make_odd_window(final_window, n, polyorder)

    y_work = yg.copy()
    outlier = np.zeros(n, dtype=bool)
    ref = savgol_filter(y_work, detect_window, polyorder, mode="interp")
    residual = y_work - ref
    center, sigma = robust_sigma_mad(residual)

    for _ in range(max_iter):
        ref = savgol_filter(y_work, detect_window, polyorder, mode="interp")
        residual = y_work - ref
        center, sigma = robust_sigma_mad(residual)
        new_outlier = np.abs(residual - center) > threshold * sigma
        outlier |= new_outlier

        keep = ~outlier
        if keep.sum() < 2:
            raise ValueError("Too many outliers detected; try a larger threshold.")
        y_work[outlier] = np.interp(xg[outlier], xg[keep], yg[keep])

    y_corrected = y_work
    y_final = savgol_filter(y_corrected, final_window, polyorder, mode="interp")
    final_residual = yg - ref

    return {
        "x": xg,
        "y_raw": yg,
        "y_reference": ref,
        "residual": final_residual,
        "outlier": outlier,
        "y_corrected": y_corrected,
        "y_smooth": y_final,
        "sigma": float(sigma),
        "residual_center": float(center),
        "detect_window": int(detect_window),
        "final_window": int(final_window),
        "polyorder": int(polyorder),
        "threshold": float(threshold),
    }


def choose_xy_columns(df: pd.DataFrame, xcol: str | None, ycol: str | None) -> tuple[str, str]:
    """Use specified columns, or infer the first two numeric columns."""
    if xcol is not None and ycol is not None:
        return xcol, ycol

    numeric_cols = [c for c in df.columns if pd.api.types.is_numeric_dtype(df[c])]
    if len(numeric_cols) < 2:
        raise ValueError("Could not infer x/y columns: fewer than two numeric columns found.")

    return xcol or numeric_cols[0], ycol or numeric_cols[1]


def plot_result(result: dict, out_png: Path = None, title: str = "Robust SG outlier smoothing") -> None:
    x = result["x"]
    y = result["y_raw"]
    ref = result["y_reference"]
    smooth = result["y_smooth"]
    outlier = result["outlier"]
    residual = result["residual"]
    center = result["residual_center"]
    sigma = result["sigma"]
    threshold = result["threshold"]

    fig, axes = plt.subplots(2, 1, figsize=(11, 8), sharex=True, height_ratios=[3, 1])

    axes[0].plot(x, y, linewidth=0.8, label="raw")
    axes[0].plot(x, ref, linewidth=1.0, label="detection SG reference")
    axes[0].plot(x, smooth, linewidth=1.3, label="corrected + final SG")
    if np.any(outlier):
        axes[0].scatter(x[outlier], y[outlier], s=24, marker="x", label="detected outlier")
    axes[0].set_ylabel("Intensity")
    axes[0].set_title(title)
    axes[0].legend()
    axes[0].grid(True, alpha=0.3)

    axes[1].plot(x, residual, linewidth=0.8, label="raw - detection reference")
    axes[1].axhline(center, linewidth=0.8)
    axes[1].axhline(center + threshold * sigma, linestyle="--", linewidth=0.8, label=f"±{threshold:g}σ by MAD")
    axes[1].axhline(center - threshold * sigma, linestyle="--", linewidth=0.8)
    if np.any(outlier):
        axes[1].scatter(x[outlier], residual[outlier], s=24, marker="x")
    axes[1].set_xlabel("x")
    axes[1].set_ylabel("Residual")
    axes[1].legend()
    axes[1].grid(True, alpha=0.3)

    fig.tight_layout()

    if out_png:
        fig.savefig(out_png, dpi=200)
        plt.close(fig)
    else:
        plt.pause(0.1)
        input("\nPress ENTER to terminate>>\n")
        


def main() -> None:
    parser = argparse.ArgumentParser(description="Robust SG smoothing with MAD-based outlier detection.")
    parser.add_argument("infile", help="Input Excel/CSV file")
    parser.add_argument("--sheet", default=0, help="Excel sheet name or index. Default: first sheet")
    parser.add_argument("--xcol", default=None, help="x column name. Default: first numeric column")
    parser.add_argument("--ycol", default=None, help="y column name. Default: second numeric column")
    parser.add_argument("--detect-window", type=int, default=15, help="SG window for outlier detection, odd integer, e.g. 11-17")
    parser.add_argument("--final-window", type=int, default=7, help="Final SG window, odd integer, e.g. 7")
    parser.add_argument("--polyorder", type=int, default=2, help="SG polynomial order")
    parser.add_argument("--threshold", type=float, default=4.0, help="Outlier threshold in robust sigma units, e.g. 3-5")
    parser.add_argument("--max-iter", type=int, default=1, help="Number of detection/correction iterations")
    parser.add_argument("--out-prefix", default=None, help="Output prefix. Default: input stem")
    parser.add_argument("--show", type=int, default=1, help="flag to show graph. Default: 1")
    args = parser.parse_args()

    infile = Path(args.infile)
    suffix = infile.suffix.lower()
    if suffix in [".xlsx", ".xlsm", ".xls"]:
        sheet = int(args.sheet) if str(args.sheet).isdigit() else args.sheet
        df = pd.read_excel(infile, sheet_name=sheet)
    elif suffix in [".csv", ".txt", ".dat"]:
        df = pd.read_csv(infile)
    else:
        raise ValueError(f"Unsupported file type: {infile.suffix}")

    xcol, ycol = choose_xy_columns(df, args.xcol, args.ycol)
    result = detect_and_smooth(
        df[xcol].to_numpy(),
        df[ycol].to_numpy(),
        detect_window=args.detect_window,
        final_window=args.final_window,
        polyorder=args.polyorder,
        threshold=args.threshold,
        max_iter=args.max_iter,
    )

    out_prefix = Path(args.out_prefix) if args.out_prefix else infile.with_suffix("")
    out_xlsx = out_prefix.with_name(out_prefix.name + "_sg_smoothed.xlsx")
    if args.show:
        out_png = None
    else:
        out_png = out_prefix.with_name(out_prefix.name + "_sg_smoothed.png")

    out_df = pd.DataFrame({
        xcol: result["x"],
        ycol: result["y_raw"],
        "sg_reference_for_detection": result["y_reference"],
        "residual_raw_minus_reference": result["residual"],
        "is_outlier": result["outlier"],
        "y_corrected_outlier_only": result["y_corrected"],
        "y_sg_final": result["y_smooth"],
    })

    with pd.ExcelWriter(out_xlsx) as writer:
        out_df.to_excel(writer, sheet_name="smoothed", index=False)
        pd.DataFrame({
            "parameter": ["x column", "y column", "detect_window", "final_window", "polyorder", "threshold", "MAD sigma", "residual center", "outlier count"],
            "value": [xcol, ycol, result["detect_window"], result["final_window"], result["polyorder"], result["threshold"], result["sigma"], result["residual_center"], int(np.sum(result["outlier"]))],
        }).to_excel(writer, sheet_name="parameters", index=False)

    plot_result(result, out_png, title=f"{infile.name}: robust SG smoothing")

    print(f"x column: {xcol}")
    print(f"y column: {ycol}")
    print(f"outliers: {int(np.sum(result['outlier']))} / {len(result['x'])}")
    print(f"MAD sigma: {result['sigma']:.6g}")
    print(f"saved: {out_xlsx}")
    print(f"saved: {out_png}")


if __name__ == "__main__":
    main()
