#!/usr/bin/env python3
import sys
import argparse
import numpy as np
import pandas as pd
import matplotlib.pyplot as plt

from spectrum import pburg, Periodogram
from statsmodels.tsa.stattools import pacf
from statsmodels.tsa.ar_model import AutoReg

import warnings
warnings.filterwarnings("ignore")


def load_xy(infile):
    """
    Load two-column data.

    Column 0: x-axis, e.g. k [A^-1]
    Column 1: signal, e.g. S(k), G(k), intensity, etc.
    """
    if infile.endswith(".xlsx"):
        df = pd.read_excel(infile, sheet_name=0)
        x_axis = df.iloc[:, 0].to_numpy(dtype=float)
        signal = df.iloc[:, 1].to_numpy(dtype=float)

    elif infile.endswith(".csv"):
        arr = np.loadtxt(infile, delimiter=",", skiprows=1, usecols=[0, 1])
        x_axis, signal = arr.T

    elif infile.endswith(".txt"):
        arr = np.loadtxt(infile, skiprows=1, usecols=[0, 1])
        x_axis, signal = arr.T

    else:
        raise ValueError("Unsupported file format. Please use .xlsx, .csv, or .txt.")

    mask = np.isfinite(x_axis) & np.isfinite(signal)
    x_axis = x_axis[mask]
    signal = signal[mask]

    if len(x_axis) < 4:
        raise ValueError("Not enough valid data points.")

    return x_axis, signal


def estimate_dx(x_axis, rtol=1e-3):
    """
    Estimate sampling step from x-axis.

    For diffraction/RDF-like data:
        x_axis = k [A^-1]
        dx     = dk [A^-1]

    The spectral frequency returned by spectrum is cycles/sample.
    Conversion to real-space-like coordinate is:

        R = 2*pi*f_sample/dk
    """
    diffs = np.diff(x_axis)
    dx = np.mean(diffs)

    if dx == 0:
        raise ValueError("x-axis step dx is zero.")

    rel_std = np.std(diffs) / abs(dx)

    if rel_std > rtol:
        print("WARNING: x-axis is not perfectly uniform.")
        print(f"  mean dx       = {dx:.8g}")
        print(f"  std(diff)     = {np.std(diffs):.8g}")
        print(f"  relative std  = {rel_std:.3e}")
        print("  FFT/MEM assume uniformly sampled data.")
        print("  Consider interpolating onto a uniform grid if needed.")

    return dx, rel_std


def select_ar_order(signal, preset_order=8, max_order=40, order_method="bic"):
    """
    Select AR model order.

    Recommended for this kind of spectral analysis:
        order_method = "bic"    : conservative automatic choice
        order_method = "aic"    : often gives slightly larger order
        order_method = "preset" : use user-specified preset_order

    PACF is printed only as diagnostic information.
    It is not used as the first-priority selector, because
    choosing the first significant PACF lag often gives order=1
    for strongly correlated smooth data.
    """
    n = len(signal)

    # ------------------------
    # PACF diagnostics only
    # ------------------------
    max_lag = min(max_order, n // 2)

    try:
        pacf_vals, confint = pacf(signal, nlags=max_lag, alpha=0.05)

        print(f"PACF with 95% CI for lags 1..{max_lag}:")
        significant_lags = []

        for lag in range(1, max_lag + 1):
            lower, upper = confint[lag]
            val = pacf_vals[lag]

            # Diagnostic only:
            # Significant if the confidence interval excludes zero.
            sig = (lower > 0.0) or (upper < 0.0)
            marker = "*" if sig else " "

            print(
                f"lag={lag:2d}: pacf={val: .4f} "
                f"CI=({lower: .4f},{upper: .4f}) {marker}"
            )

            if sig:
                significant_lags.append(lag)

        if significant_lags:
            print(
                "PACF significant lags, diagnostic only: "
                f"{significant_lags[:10]}"
            )
        else:
            print("No PACF spike outside zero-confidence interval.")

        print("PACF is not used as the primary order selector.\n")

    except Exception as e:
        print(f"PACF calculation failed: {e}")
        print("Continue with AIC/BIC order selection.\n")

    # ------------------------
    # AIC / BIC selection
    # ------------------------
    best_aic = np.inf
    best_bic = np.inf
    best_p_aic = None
    best_p_bic = None

    max_ar = min(max_order, n - 1)

    for p in range(1, max_ar + 1):
        try:
            model = AutoReg(signal, lags=p, old_names=False).fit()

            if model.aic < best_aic:
                best_aic = model.aic
                best_p_aic = p

            if model.bic < best_bic:
                best_bic = model.bic
                best_p_bic = p

        except Exception:
            pass

    print(f"Selected order by AIC: {best_p_aic}, AIC={best_aic:.2f}")
    print(f"Selected order by BIC: {best_p_bic}, BIC={best_bic:.2f}\n")

    # ------------------------
    # Final order selection
    # ------------------------
    if order_method == "bic":
        if best_p_bic is not None:
            order_selected = best_p_bic
        elif best_p_aic is not None:
            order_selected = best_p_aic
        else:
            order_selected = preset_order

    elif order_method == "aic":
        if best_p_aic is not None:
            order_selected = best_p_aic
        elif best_p_bic is not None:
            order_selected = best_p_bic
        else:
            order_selected = preset_order

    elif order_method == "preset":
        order_selected = preset_order

    else:
        raise ValueError(
            "order_method must be one of: 'bic', 'aic', 'preset'"
        )

    print(f"Order selection method = {order_method}")
    print(f"Using AR model order = {order_selected}\n")

    return order_selected, {
        "order_method": order_method,
        "best_p_aic": best_p_aic,
        "best_aic": best_aic,
        "best_p_bic": best_p_bic,
        "best_bic": best_bic,
    }


def main():
    parser = argparse.ArgumentParser(
        description=(
            "MEM/FFT spectral analysis with physical-axis conversion. "
            "Input column 0 is assumed to be x, e.g. k [A^-1]. "
            "The output horizontal axis is R = 2*pi*f_sample/dx."
        )
    )

    parser.add_argument(
        "infile",
        nargs="?",
        default="input_time_series.csv",
        help="Input file: .xlsx, .csv, or .txt",
    )

    parser.add_argument(
        "preset_order",
        nargs="?",
        type=int,
        default=8,
        help="Fallback AR model order",
    )

    parser.add_argument(
        "--remove-mean",
        action="store_true",
        help="Subtract mean from the signal before MEM/FFT.",
    )

    parser.add_argument(
        "--nfft",
        type=int,
        default=0,
        help="NFFT. If 0, use len(signal).",
    )

    parser.add_argument(
        "--max-order",
        type=int,
        default=40,
        help="Maximum AR order tested by PACF/AIC/BIC.",
    )

    parser.add_argument(
        "--x-unit",
        default="A^-1",
        help="Unit of input x-axis. Default: A^-1",
    )

    parser.add_argument(
        "--r-unit",
        default="A",
        help="Unit of converted output axis. Default: A",
    )

    parser.add_argument(
        "--yscale",
        choices=["linear", "log"],
        default="linear",
        help="Y scale for spectrum plot.",
    )

    parser.add_argument(
        "--order-method",
        choices=["bic", "aic", "preset"],
        default="bic",
        help="AR order selection method. Default: bic.",
    )

    args = parser.parse_args()

    infile = args.infile

    if infile == "":
        print("No input file specified. Exiting.")
        sys.exit(1)

    # ------------------------
    # Load data
    # ------------------------
    try:
        x_axis, signal_raw = load_xy(infile)
    except Exception as e:
        print(f"Error loading input file: {e}")
        sys.exit(1)

    dx, rel_std = estimate_dx(x_axis)

    signal = signal_raw.copy()

    if args.remove_mean:
        signal = signal - np.mean(signal)
        mean_removed = True
    else:
        mean_removed = False

    n_data = len(signal)
    nfft = args.nfft if args.nfft > 0 else n_data

    print("Input data summary")
    print(f"  file          = {infile}")
    print(f"  n_data        = {n_data}")
    print(f"  nfft          = {nfft}")
    print(f"  x_min         = {np.min(x_axis):.8g} [{args.x_unit}]")
    print(f"  x_max         = {np.max(x_axis):.8g} [{args.x_unit}]")
    print(f"  dx            = {dx:.8g} [{args.x_unit}]")
    print(f"  rel std dx    = {rel_std:.3e}")
    print(f"  remove mean   = {mean_removed}")
    print()

    # ------------------------
    # Select AR order
    # ------------------------
    order_selected, order_info = select_ar_order(
        signal,
        preset_order=args.preset_order,
        max_order=args.max_order,
        order_method=args.order_method,
    )

    # ------------------------
    # MEM spectrum by Burg method
    # ------------------------
    burg_spec = pburg(signal, order=order_selected, NFFT=nfft)
    mem_psd = np.asarray(burg_spec.psd)
    freqs_mem_sample = np.asarray(burg_spec.frequencies())

    # ------------------------
    # FFT-based PSD
    # ------------------------
    fft_spec = Periodogram(signal, NFFT=nfft)
    fft_psd = np.asarray(fft_spec.psd)
    freqs_fft_sample = np.asarray(fft_spec.frequencies())

    # ------------------------
    # Convert horizontal axis
    # ------------------------
    # spectrum returns normalized frequency in cycles/sample.
    #
    # If input x-axis is k [A^-1] and dx = dk [A^-1],
    #
    #     f_physical = f_sample / dk     [cycles * A]
    #     R          = 2*pi*f_sample/dk  [A]
    #
    # This is the usual conversion when Fourier kernel is exp(i k R).
    R_mem = 2.0 * np.pi * freqs_mem_sample / dx
    R_fft = 2.0 * np.pi * freqs_fft_sample / dx

    # ------------------------
    # Peak estimate, excluding near-zero component
    # ------------------------
    def find_main_peak(R, psd, rmin=1e-12):
        mask = np.isfinite(R) & np.isfinite(psd) & (R > rmin)
        if not np.any(mask):
            return np.nan, np.nan

        R2 = R[mask]
        psd2 = psd[mask]
        idx = np.argmax(psd2)
        return R2[idx], psd2[idx]

    peak_R_mem, peak_psd_mem = find_main_peak(R_mem, mem_psd)
    peak_R_fft, peak_psd_fft = find_main_peak(R_fft, fft_psd)

    print("Main peak excluding R=0")
    print(f"  MEM: R = {peak_R_mem:.8g} [{args.r_unit}], PSD = {peak_psd_mem:.8g}")
    print(f"  FFT: R = {peak_R_fft:.8g} [{args.r_unit}], PSD = {peak_psd_fft:.8g}")
    print()

    # ------------------------
    # Save results to Excel file
    # ------------------------
    filebody = infile.rsplit(".", 1)[0]
    output_filename = f"{filebody}-mem_fft_R_order={order_selected}.xlsx"

    n_out = max(len(R_mem), len(R_fft))

    def pad_array(a, n):
        a = np.asarray(a)
        out = np.full(n, np.nan)
        out[:len(a)] = a
        return out

    df_output = pd.DataFrame({
        "Frequency_MEM_cycles_per_sample": pad_array(freqs_mem_sample, n_out),
        f"R_MEM_{args.r_unit}": pad_array(R_mem, n_out),
        "PSD_MEM": pad_array(mem_psd, n_out),

        "Frequency_FFT_cycles_per_sample": pad_array(freqs_fft_sample, n_out),
        f"R_FFT_{args.r_unit}": pad_array(R_fft, n_out),
        "PSD_FFT": pad_array(fft_psd, n_out),
    })

    df_input = pd.DataFrame({
        f"x_axis_{args.x_unit}": x_axis,
        "signal_raw": signal_raw,
        "signal_used": signal,
    })

    df_summary = pd.DataFrame({
        "item": [
            "input_file",
            "n_data",
            "nfft",
            f"dx [{args.x_unit}]",
            "relative_std_dx",
            "remove_mean",
            "AR_order_selected",
            "order_method",
            "best_p_aic",
            "best_aic",
            "best_p_bic",
            "best_bic",
            f"main_peak_R_MEM [{args.r_unit}]",
            "main_peak_PSD_MEM",
            f"main_peak_R_FFT [{args.r_unit}]",
            "main_peak_PSD_FFT",
            "conversion",
        ],
        "value": [
            infile,
            n_data,
            nfft,
            dx,
            rel_std,
            mean_removed,
            order_selected,
            order_info.get("order_method", ""),
            order_info.get("best_p_aic", ""),
            order_info.get("best_aic", ""),
            order_info.get("best_p_bic", ""),
            order_info.get("best_bic", ""),
            peak_R_mem,
            peak_psd_mem,
            peak_R_fft,
            peak_psd_fft,
            "R = 2*pi*frequency_cycles_per_sample/dx",
        ],
    })

    try:
        with pd.ExcelWriter(output_filename, engine="openpyxl") as writer:
            df_input.to_excel(writer, sheet_name="Input", index=False)
            df_output.to_excel(writer, sheet_name="MEM_FFT", index=False)
            df_summary.to_excel(writer, sheet_name="Summary", index=False)

        print(f"MEM and FFT results saved to {output_filename}")

    except Exception as e:
        print(f"Error saving results to Excel file: {e}")

    # ------------------------
    # Plotting
    # ------------------------
    fig, axes = plt.subplots(nrows=2, ncols=1, figsize=(8, 8))

    axes[0].plot(x_axis, signal_raw, label="Raw signal", linewidth=1.5)

    if mean_removed:
        axes[0].plot(x_axis, signal, label="Mean-removed signal", linewidth=1.0)

    axes[0].set_xlabel(f"x [{args.x_unit}]")
    axes[0].set_ylabel("Amplitude")
    axes[0].set_title("Input data")
    axes[0].grid(True)
    axes[0].legend()

    axes[1].plot(R_mem, mem_psd, label=f"MEM / Burg, order={order_selected}", linewidth=2)
    axes[1].plot(R_fft, fft_psd, "--", label="FFT PSD", linewidth=1)

    axes[1].set_xlabel(f"R [{args.r_unit}]")
    axes[1].set_ylabel("PSD")
    axes[1].set_title("Spectral estimate with converted horizontal axis")
    axes[1].grid(True)
    axes[1].legend()

    if args.yscale == "log":
        axes[1].set_yscale("log")

    plt.tight_layout()
    plt.show()

    input("\nPress ENTER to terminate>>\n")


if __name__ == "__main__":
    main()