"""
Analyze sample timing consistency from a dual-COM-port nanoFET CSV log.

Checks:
  - Per-port inter-sample intervals (should be ~constant)
  - Inter-port offset within each cycle (COM6 vs COM7 pairing)
  - Drift and gaps over the recording
  - Overall statistics + plots

Usage:
    python analyzeSampleTiming.py [path/to/file.csv]
"""

import sys
import pandas as pd
import numpy as np
import matplotlib.pyplot as plt
import matplotlib.dates as mdates
from pathlib import Path

# ---------------------------------------------------------------------------
# Config
# ---------------------------------------------------------------------------
DEFAULT_CSV = (
    r"\\thalassa\ProjectLibrary\901805_Coastal_Biogeochemical_Sensing"
    r"\Wetlab_Sensor_Calibration\NanoFet\K0"
    r"\06_22_26_id_Nanov2_OLD_Nanov2_BOX_2SecondSampling.csv"
)

GAP_THRESHOLD_S = 10  # flag inter-sample gaps larger than this (seconds)

# ---------------------------------------------------------------------------
# Load
# ---------------------------------------------------------------------------
csv_path = Path(sys.argv[1]) if len(sys.argv) > 1 else Path(DEFAULT_CSV)
print(f"Reading: {csv_path}")

df = pd.read_csv(csv_path, parse_dates=["datetime"])
df = df.sort_values("datetime").reset_index(drop=True)

ports = sorted(df["com_port"].unique())
print(f"COM ports found: {ports}")
print(f"Total rows: {len(df)}  |  Time span: {df['datetime'].iloc[0]} to {df['datetime'].iloc[-1]}")
print(f"Duration: {(df['datetime'].iloc[-1] - df['datetime'].iloc[0]).total_seconds():.1f} s\n")

# ---------------------------------------------------------------------------
# Per-port interval analysis
# ---------------------------------------------------------------------------
print("=" * 60)
print("PER-PORT INTER-SAMPLE INTERVALS")
print("=" * 60)

port_data = {}
for port in ports:
    sub = df[df["com_port"] == port].copy()
    sub["interval_s"] = sub["datetime"].diff().dt.total_seconds()
    port_data[port] = sub

    intervals = sub["interval_s"].dropna()
    gaps = intervals[intervals > GAP_THRESHOLD_S]
    print(f"\n{port}  ({len(sub)} samples)")
    print(f"  Interval  mean={intervals.mean():.3f}s  median={intervals.median():.3f}s"
          f"  std={intervals.std():.3f}s  min={intervals.min():.3f}s  max={intervals.max():.3f}s")
    if len(gaps):
        print(f"  *** {len(gaps)} gap(s) > {GAP_THRESHOLD_S}s ***")
        for t, g in zip(sub.loc[gaps.index, "datetime"], gaps):
            print(f"      {t}  ({g:.1f} s)")
    else:
        print(f"  No gaps > {GAP_THRESHOLD_S}s detected.")

# ---------------------------------------------------------------------------
# Inter-port pairing analysis (COM6 leads COM7 each cycle)
# ---------------------------------------------------------------------------
print("\n" + "=" * 60)
print("INTER-PORT OFFSET WITHIN EACH CYCLE")
print("=" * 60)

if len(ports) == 2:
    p0, p1 = ports[0], ports[1]
    t0 = port_data[p0]["datetime"].reset_index(drop=True)
    t1 = port_data[p1]["datetime"].reset_index(drop=True)
    n_pairs = min(len(t0), len(t1))
    offsets = (t1.iloc[:n_pairs].values - t0.iloc[:n_pairs].values) / np.timedelta64(1, "s")

    print(f"\nPairing {p0} -> {p1}  ({n_pairs} pairs)")
    print(f"  Offset  mean={offsets.mean():.3f}s  median={np.median(offsets):.3f}s"
          f"  std={offsets.std():.3f}s  min={offsets.min():.3f}s  max={offsets.max():.3f}s")

    bad_pairs = np.where(offsets < 0)[0]
    if len(bad_pairs):
        print(f"  *** {len(bad_pairs)} negative offset(s) – {p1} arrived BEFORE {p0} ***")
        for idx in bad_pairs[:10]:
            print(f"      pair {idx}: {t0.iloc[idx]} vs {t1.iloc[idx]}  ({offsets[idx]:.1f}s)")
else:
    print("  Skipped – need exactly 2 ports for pairing analysis.")
    offsets = None

# ---------------------------------------------------------------------------
# Plots
# ---------------------------------------------------------------------------
n_ports = len(ports)
fig, axes = plt.subplots(n_ports + (1 if offsets is not None else 0), 2,
                         figsize=(14, 4 * (n_ports + 1)),
                         gridspec_kw={"width_ratios": [3, 1]})
fig.suptitle(f"Sample Timing Analysis\n{csv_path.name}", fontsize=11)

for row_idx, port in enumerate(ports):
    sub = port_data[port]
    intervals = sub["interval_s"].dropna()
    t_mid = sub["datetime"].iloc[1:]   # align with interval (diff shifts by 1)

    ax_ts = axes[row_idx, 0]
    ax_hist = axes[row_idx, 1]

    # Time-series of intervals
    ax_ts.plot(t_mid, intervals, lw=0.7, marker=".", ms=3, label=port)
    ax_ts.axhline(intervals.median(), color="red", lw=1, ls="--", label=f"median={intervals.median():.2f}s")
    ax_ts.axhline(GAP_THRESHOLD_S, color="orange", lw=1, ls=":", label=f"gap threshold={GAP_THRESHOLD_S}s")
    ax_ts.set_ylabel("Interval (s)")
    ax_ts.set_title(f"{port} – inter-sample interval over time")
    ax_ts.xaxis.set_major_formatter(mdates.DateFormatter("%H:%M"))
    ax_ts.legend(fontsize=8)
    ax_ts.grid(True, alpha=0.3)

    # Histogram
    ax_hist.hist(intervals, bins=40, color="steelblue", edgecolor="white", lw=0.3)
    ax_hist.axvline(intervals.median(), color="red", lw=1.5, ls="--")
    ax_hist.set_xlabel("Interval (s)")
    ax_hist.set_ylabel("Count")
    ax_hist.set_title(f"{port} – histogram")
    ax_hist.grid(True, alpha=0.3)

# Inter-port offset row
if offsets is not None and len(ports) == 2:
    row_idx = n_ports
    t_offset = t0.iloc[:n_pairs]

    ax_ts = axes[row_idx, 0]
    ax_hist = axes[row_idx, 1]

    ax_ts.plot(t_offset, offsets, lw=0.7, marker=".", ms=3, color="purple")
    ax_ts.axhline(np.median(offsets), color="red", lw=1, ls="--",
                  label=f"median={np.median(offsets):.2f}s")
    ax_ts.axhline(0, color="black", lw=0.8, ls="-")
    ax_ts.set_ylabel("Offset (s)")
    ax_ts.set_title(f"Inter-port offset ({p0} -> {p1}) over time")
    ax_ts.xaxis.set_major_formatter(mdates.DateFormatter("%H:%M"))
    ax_ts.legend(fontsize=8)
    ax_ts.grid(True, alpha=0.3)

    ax_hist.hist(offsets, bins=40, color="mediumpurple", edgecolor="white", lw=0.3)
    ax_hist.axvline(np.median(offsets), color="red", lw=1.5, ls="--")
    ax_hist.set_xlabel("Offset (s)")
    ax_hist.set_ylabel("Count")
    ax_hist.set_title("Inter-port offset – histogram")
    ax_hist.grid(True, alpha=0.3)

plt.tight_layout()

out_path = csv_path.parent / (csv_path.stem + "_timing_analysis.png")
plt.savefig(out_path, dpi=150, bbox_inches="tight")
print(f"\nPlot saved -> {out_path}")
plt.show()
