#!/usr/bin/env python3
"""
Feature construction for smoking detection (vectorised).

Everything here is scale-free and self-calibrating, so the same code runs on a new export,
a new device or a new person without edits. Concretely:

  * baselines are CAUSAL ROLLING MEDIANS, not an EMA with a hand-picked alpha
  * deviations are scaled by a rolling MAD, so "unusual" is measured in this person's own
    units on this day rather than against a literal bpm threshold
  * the exertion correction is a coefficient FITTED from the subject's own data
  * the response template is a family parameterised by (onset, tau_rise, tau_decay); which
    member to use is chosen by cross-validation in detector.py, never asserted here

WHY EXERTION RESIDUALISATION IS THE CENTRAL IDEA. The one reproducible failure mode in
this project is that walking swamps the nicotine response (REPORT.md 5: 19:24 missed with
104 concurrent steps). Steps are measured, so instead of penalising movement with a
hand-tuned weight -- which is what the original S formula did, and it carried no
information -- regress heart rate on recent step load and keep only the part of the rise
that movement does not explain.

CAUSALITY. Every statistic is computed from strictly past samples. A detector that peeks
forward cannot run tomorrow, and "works on tomorrow's datapoints" is the requirement.
"""
import numpy as np
from numpy.lib.stride_tricks import sliding_window_view


def _causal_windows(a, w):
    """(n, w) view where row t holds a[t-w:t]; rows before w are NaN."""
    pad = np.full(w, np.nan)
    return sliding_window_view(np.concatenate([pad, a]), w)[:-1]


def causal_median_mad(a, w):
    """Median and MAD of the w minutes strictly before each sample."""
    win = _causal_windows(a, w)
    with np.errstate(invalid="ignore"):
        med = np.nanmedian(win, axis=1)
        mad = np.nanmedian(np.abs(win - med[:, None]), axis=1)
    return med, mad


def step_load(steps, tau):
    """Exponentially weighted recent step count: a proxy for current metabolic demand."""
    k = float(np.exp(-1.0 / tau))
    out = np.empty_like(steps, dtype=float)
    acc = 0.0
    s = np.nan_to_num(steps, nan=0.0)
    for i in range(len(s)):
        acc = acc * k + s[i]
        out[i] = acc
    return out


def theil_sen_slope(x, y):
    """Robust slope via median of the lower/upper tertile contrast.

    Least squares would be dragged around by the exercise tail; this is not.
    """
    m = np.isfinite(x) & np.isfinite(y) & (x > 0)
    if m.sum() < 20:
        return 0.0
    xs, ys = x[m], y[m]
    o = np.argsort(xs)
    xs, ys = xs[o], ys[o]
    k = max(1, len(xs) // 3)
    dx = np.median(xs[-k:]) - np.median(xs[:k])
    dy = np.median(ys[-k:]) - np.median(ys[:k])
    return float(dy / dx) if abs(dx) > 1e-9 else 0.0


def template(onset, tau_rise, tau_decay, pre, post):
    """Unit-peak nicotine response: flat, then fast rise, then slow decay.

    `onset` shifts the response later than the logged light time. A cigarette is smoked
    over several minutes and the cardiovascular peak tracks accumulated dose, so onset = 0
    is a hypothesis rather than a given. It is fitted.
    """
    u = np.arange(-pre, post + 1, dtype=float) - onset
    t = np.where(u < 0, 0.0, (1 - np.exp(-np.clip(u, 0, None) / tau_rise))
                 * np.exp(-np.clip(u, 0, None) / tau_decay))
    pk = t.max()
    return t / pk if pk > 0 else t


class Channel:
    """One device's heart rate on a regular minute grid, residualised twice over.

    First against its own recent baseline, then against recent movement. What survives is
    the part of a heart-rate excursion that neither drift nor exertion accounts for.
    """

    def __init__(self, minutes, hr, steps, *, baseline_window, step_tau):
        self.minutes = minutes
        self.hr = hr
        med, mad = causal_median_mad(hr, baseline_window)
        self.load = step_load(steps, step_tau)
        raw = hr - med
        self.beta = theil_sen_slope(self.load, raw)
        resid = raw - self.beta * self.load
        scale = 1.4826 * mad
        with np.errstate(invalid="ignore", divide="ignore"):
            self.z = np.where(scale > 1e-6, resid / scale, np.nan)

    def segments(self, pre, post):
        """(n, pre+post+1) matrix of z-windows centred on each minute, NaN where short."""
        w = pre + post + 1
        pad_l = np.full(pre, np.nan)
        pad_r = np.full(post, np.nan)
        return sliding_window_view(np.concatenate([pad_l, self.z, pad_r]), w)

    def match(self, segs, tmpl, pre):
        """Shape correlation and response height (in MAD units) for every minute."""
        t = tmpl - tmpl.mean()
        tn = np.sqrt((t ** 2).sum())
        mu = segs.mean(axis=1)
        d = segs - mu[:, None]
        sn = np.sqrt((d ** 2).sum(axis=1))
        with np.errstate(invalid="ignore", divide="ignore"):
            corr = (d @ t) / (sn * tn)
        base = np.nanmedian(segs[:, :pre], axis=1)
        amp = np.nanmax(segs[:, pre:], axis=1) - base
        return corr, amp
