"""Per-subject features from the sit-to-stand recordings of the PhysioNet
"Cerebral Vasoregulation in Diabetes" database (v1.0.0, files S####DC).

For each subject and each standing bout (1 = eyes open, 2 = eyes closed):
  - posture from the force plate's vertical channel (fz), not from event markers;
  - a sitting window (the SIT_S seconds ending GAP_S before standing begins) and a
    standing window (from SETTLE_S after standing begins to END_TRIM_S before it ends);
  - beat-to-beat systolic pressure (SBP) and mean cerebral blood flow velocity (CBFV)
    from the finger arterial pressure and middle cerebral artery Doppler waveforms;
  - early-warning statistics of the linearly detrended beat series: lag-1
    autocorrelation (AR1) and variance;
  - sway while standing, from horizontal ground reaction forces, and magnitude-squared
    coherence between anteroposterior sway and SBP in the 0.05-0.15 Hz band, with a phase-randomised
    surrogate z-score.

Usage: python code/features.py <directory holding the S####DC.hea/.dat files> <out.csv>
"""
import csv
import math
import pathlib
import sys

import numpy as np
import wfdb
from scipy import signal

SIT_S = 240.0
GAP_S = 10.0
SETTLE_S = 30.0
END_TRIM_S = 5.0
MIN_BOUT_S = 60.0
MIN_BEATS = 100
MAX_BAD_FRACTION = 0.10
FS_RESAMPLE = 4.0
BAND = (0.05, 0.15)
NPERSEG_S = 64.0
N_SURROGATES = 200
SEED = 20261007


def standing_bouts(fz, fs):
    """Contiguous runs where the vertical force exceeds half its loaded level."""
    step = int(fs)  # 1 s medians
    med = np.array([np.nanmedian(fz[i:i + step]) for i in range(0, len(fz) - step, step)])
    loaded = np.nanpercentile(med, 95)
    unloaded = np.nanpercentile(med, 5)
    if loaded - unloaded < 50:  # no clear weight on the plate
        return []
    up = med > unloaded + 0.5 * (loaded - unloaded)
    bouts, start = [], None
    for i, u in enumerate(np.append(up, False)):
        if u and start is None:
            start = i
        elif not u and start is not None:
            if i - start >= MIN_BOUT_S:
                bouts.append((float(start), float(i)))
            start = None
    return bouts


def beats(abp, fs, t0, t1):
    """Systolic peaks and per-beat values within [t0, t1) seconds."""
    a, b = int(t0 * fs), int(t1 * fs)
    x = abp[a:b]
    if len(x) < fs * 30 or np.isnan(x).mean() > 0.1:
        return None
    x = np.nan_to_num(x, nan=np.nanmedian(x))
    sos = signal.butter(4, 10, fs=fs, output="sos")
    xf = signal.sosfiltfilt(sos, x)
    iqr = np.subtract(*np.percentile(xf, [90, 10]))
    peaks, _ = signal.find_peaks(xf, distance=int(0.33 * fs), prominence=max(0.3 * iqr, 5))
    if len(peaks) < 3:
        return None
    sbp = xf[peaks]
    times = (peaks + a) / fs
    # artefact rule: implausible values, or more than 30% away from the 9-beat running median
    run = np.array([np.median(sbp[max(0, i - 4):i + 5]) for i in range(len(sbp))])
    good = (sbp > 60) & (sbp < 250) & (np.abs(sbp - run) <= 0.3 * run)
    return times, sbp, good


def per_beat_mean(sig, fs, times):
    """Mean of a waveform between consecutive beat times."""
    idx = (times * fs).astype(int)
    out = np.full(len(times), np.nan)
    for i in range(len(idx) - 1):
        seg = sig[idx[i]:idx[i + 1]]
        if len(seg) and not np.isnan(seg).all():
            out[i] = np.nanmean(seg)
    return out


def ews(values):
    """AR1 and variance of the linearly detrended series."""
    v = np.asarray(values, float)
    v = v[~np.isnan(v)]
    if len(v) < MIN_BEATS:
        return float("nan"), float("nan")
    r = signal.detrend(v)
    ar1 = float(np.corrcoef(r[:-1], r[1:])[0, 1])
    return ar1, float(np.var(r, ddof=1))


def window_stats(rec, t0, t1, cbfv_name):
    fs = rec.fs
    names = rec.sig_name
    abp = rec.p_signal[:, names.index("abp")]
    got = beats(abp, fs, t0, t1)
    if got is None:
        return None
    times, sbp, good = got
    bad_fraction = 1 - good.mean()
    if bad_fraction > MAX_BAD_FRACTION or good.sum() < MIN_BEATS:
        return {"ok": False, "bad_fraction": bad_fraction, "beats": int(good.sum())}
    sbp_g = sbp[good]
    out = {"ok": True, "bad_fraction": bad_fraction, "beats": int(good.sum()),
           "sbp_mean": float(np.mean(sbp_g))}
    out["sbp_ar1"], out["sbp_var"] = ews(sbp_g)
    if cbfv_name:
        cb = per_beat_mean(rec.p_signal[:, names.index(cbfv_name)], fs, times)[good]
        if np.isnan(cb).mean() < 0.1:
            out["cbfv_mean_v"] = float(np.nanmean(cb))
            out["cbfv_ar1"], var = ews(cb)
            out["cbfv_cv"] = float(math.sqrt(var) / np.nanmean(cb)) if np.isfinite(var) else float("nan")
    out["_times"], out["_sbp"] = times[good], sbp_g
    return out


def choose_cbfv(rec, t0, t1):
    """The middle cerebral artery channel with the larger pulsatile amplitude while sitting.
    The Doppler channels are stored uncalibrated (volts), so only scale-free statistics
    (AR1, coefficient of variation) are computed from them; a channel whose standard
    deviation is below 0.05 V is treated as not recorded."""
    best, best_v = None, 0.05
    for name in ("mcar", "mcal"):
        if name in rec.sig_name:
            x = rec.p_signal[int(t0 * rec.fs):int(t1 * rec.fs), rec.sig_name.index(name)]
            v = np.nanstd(x)
            if np.isfinite(v) and v > best_v and np.isnan(x).mean() < 0.1:
                best, best_v = name, v
    return best


def sway_and_coherence(rec, t0, t1, beat_times, sbp, rng):
    fs = rec.fs
    names = rec.sig_name
    a, b = int(t0 * fs), int(t1 * fs)
    # The centre-of-pressure channels (px, py) are corrupted in this database (runs of exact
    # zeros and spikes of thousands of mm), so sway is measured by the horizontal ground
    # reaction forces fy (labelled anteroposterior) and fx (mediolateral), divided by the
    # mean vertical force: by Newton's second law, the horizontal acceleration of the body's
    # centre of mass in units of g.
    fz = np.nanmean(rec.p_signal[a:b, names.index("fz")])
    ap = rec.p_signal[a:b, names.index("fy")] / fz
    ml = rec.p_signal[a:b, names.index("fx")] / fz
    if np.isnan(ap).mean() > 0.1 or np.nanstd(ap) == 0:
        return {}
    step = int(fs / FS_RESAMPLE)
    sos = signal.butter(4, 1.0, fs=fs, output="sos")
    ap4 = signal.sosfiltfilt(sos, np.nan_to_num(ap, nan=np.nanmedian(ap)))[::step]
    ml4 = signal.sosfiltfilt(sos, np.nan_to_num(ml, nan=np.nanmedian(ml)))[::step]
    out = {"sway_ap_sd_g": float(np.std(signal.detrend(ap4))),
           "sway_ml_sd_g": float(np.std(signal.detrend(ml4)))}
    grid = t0 + np.arange(len(ap4)) / FS_RESAMPLE
    keep = (grid >= beat_times[0]) & (grid <= beat_times[-1])
    sbp4 = np.interp(grid[keep], beat_times, sbp)
    x, y = signal.detrend(ap4[keep]), signal.detrend(sbp4)
    nper = int(NPERSEG_S * FS_RESAMPLE)
    if len(x) < 2 * nper:
        return out
    f, c = signal.coherence(x, y, fs=FS_RESAMPLE, nperseg=nper)
    band = (f >= BAND[0]) & (f <= BAND[1])
    observed = float(c[band].mean())
    sur = []
    Y = np.fft.rfft(y)
    for _ in range(N_SURROGATES):
        ph = np.exp(1j * rng.uniform(0, 2 * np.pi, len(Y)))
        ph[0] = 1
        ys = np.fft.irfft(np.abs(Y) * ph, n=len(y))
        _, cs = signal.coherence(x, ys, fs=FS_RESAMPLE, nperseg=nper)
        sur.append(cs[band].mean())
    sur = np.array(sur)
    out["coh_sway_sbp"] = observed
    out["coh_sway_sbp_z"] = float((observed - sur.mean()) / sur.std(ddof=1))
    return out


def subject_rows(path, rng):
    rec = wfdb.rdrecord(str(path))
    subject = path.name[:5].upper()
    fz = rec.p_signal[:, rec.sig_name.index("fz")]
    bouts = standing_bouts(fz, rec.fs)
    rows = []
    for k, (s0, s1) in enumerate(bouts[:2], start=1):
        row = {"subject": subject, "bout": k, "stand_onset_s": s0, "stand_end_s": s1}
        sit0, sit1 = s0 - GAP_S - SIT_S, s0 - GAP_S
        st0, st1 = s0 + SETTLE_S, s1 - END_TRIM_S
        if sit0 < 0 or st1 - st0 < 90:
            row["status"] = "window_too_short"
            rows.append(row)
            continue
        cb = choose_cbfv(rec, sit0, sit1)
        row["cbfv_channel"] = cb or ""
        sit = window_stats(rec, sit0, sit1, cb)
        st = window_stats(rec, st0, st1, cb)
        if not sit or not st or not sit["ok"] or not st["ok"]:
            row["status"] = "bp_quality"
            for tag, w in (("sit", sit), ("stand", st)):
                if w:
                    row[f"{tag}_bad_fraction"] = round(w["bad_fraction"], 4)
            rows.append(row)
            continue
        row["status"] = "ok"
        for tag, w in (("sit", sit), ("stand", st)):
            for key, val in w.items():
                if not key.startswith("_") and key != "ok":
                    row[f"{tag}_{key}"] = val
        row.update(sway_and_coherence(rec, st0, st1, st["_times"], st["_sbp"], rng))
        rows.append(row)
    if not bouts:
        rows.append({"subject": subject, "bout": 0, "status": "no_standing_detected"})
    return rows


def main():
    src, dst = pathlib.Path(sys.argv[1]), pathlib.Path(sys.argv[2])
    rng = np.random.default_rng(SEED)
    rows = []
    for hea in sorted(src.glob("*DC.hea")):
        rows += subject_rows(hea.with_suffix(""), rng)
        print(rows[-1]["subject"], [r.get("status") for r in rows if r["subject"] == rows[-1]["subject"]], flush=True)
    keys = []
    for r in rows:
        keys += [k for k in r if k not in keys]
    dst.parent.mkdir(parents=True, exist_ok=True)
    with open(dst, "w", newline="") as f:
        w = csv.DictWriter(f, fieldnames=keys)
        w.writeheader()
        for r in rows:
            w.writerow({k: (round(v, 6) if isinstance(v, float) and math.isfinite(v) else v) for k, v in r.items()})


if __name__ == "__main__":
    main()
