#!/usr/bin/env python3
"""
Smoking detector v2 -- universal, cross-validated, no hardcoded constants.

HOW THIS DIFFERS FROM smoke_shape.py
  1. Exertion residualisation: heart rate is regressed on recent step load and only the
     unexplained part is scored, instead of penalising movement with a hand-picked weight.
  2. Learned response shape: the template's onset, rise and decay are chosen by
     leave-one-day-out CV rather than asserted a priori.
  3. Learned operating point: the old detector emitted a fixed top-N per day, calibrated
     against a self-reported "8-12 cigarettes/day" that the 12-17 Aug labels refute
     (actual 0-4/day plus bursts). A fixed N guarantees ~9 false alarms on a 1-cigarette
     day and cannot represent a zero day at all. Here an ABSOLUTE threshold in MAD units
     is calibrated so that training-day alarms equal training-day cigarettes, and the
     number of alarms per day is then free -- including zero.
  4. Cross-device agreement: two independent sensors see the same heart. A real
     cardiovascular event appears in both; a PPG motion artifact is usually specific to
     one wrist. Whether to require agreement is fitted, not assumed.
  5. Honest nulls: chance is measured by CIRCULAR SHIFT of the alarm train, which
     preserves alarm count, spacing and clustering and destroys only alignment to labels.
     Random placement destroys that structure too and is anti-conservative for
     autocorrelated series.

Ground truth is read from data/, never embedded here.
"""
import csv
import itertools
import json
import os
import random
import sys
from collections import defaultdict
from datetime import datetime

import numpy as np

sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
from features import Channel, template                                   # noqa: E402
from hdata import load_apple_health, default_export                      # noqa: E402

ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
PRE, POST = 8, 30
SEP = 20                      # minimum minutes between two alarms
TOLS = (5, 10, 15)


# ----------------------------------------------------------------- ground truth
def load_labels():
    pts = []
    with open(os.path.join(ROOT, "data", "cigarette_log.csv")) as f:
        for row in csv.reader(f):
            if not row or row[0].startswith(("#", "datetime")):
                continue
            pts.append(datetime.strptime(row[0].strip(), "%Y-%m-%dT%H:%M"))
    ivs = []
    q = os.path.join(ROOT, "data", "label_intervals.csv")
    if os.path.exists(q):
        with open(q) as f:
            for row in csv.DictReader(r for r in f if not r.startswith("#")):
                ivs.append((datetime.strptime(row["start"], "%Y-%m-%dT%H:%M"),
                            datetime.strptime(row["end"], "%Y-%m-%dT%H:%M"),
                            row["kind"], int(row["n_min"])))
    return sorted(pts), ivs


# ----------------------------------------------------------------- data assembly
def regularise(a):
    """Put an irregularly-sampled channel on the 1-minute grid without inventing data.

    Devices report at their own cadence -- Fitbit relays every minute, Garmin every two.
    On a 1-minute grid the slower device leaves every other slot empty, so no 39-minute
    analysis window is ever complete and that device silently contributes nothing. The fix
    is to interpolate across gaps no larger than the device's OWN native cadence (measured
    here, not assumed) and to leave genuine outages as NaN. Interpolating one minute
    between two real samples two minutes apart adds no information but makes the sensor
    usable; interpolating across an hour would fabricate it, so that is refused.
    """
    idx = np.flatnonzero(np.isfinite(a))
    if idx.size < 3:
        return a
    gaps = np.diff(idx)
    cadence = float(np.median(gaps)) if gaps.size else 1.0
    max_gap = max(2.0, 1.5 * cadence)
    filled = np.interp(np.arange(len(a)), idx, a[idx])
    nearest = np.minimum(
        np.abs(np.arange(len(a)) - idx[np.clip(np.searchsorted(idx, np.arange(len(a))) - 1,
                                               0, idx.size - 1)]),
        np.abs(np.arange(len(a)) - idx[np.clip(np.searchsorted(idx, np.arange(len(a))),
                                               0, idx.size - 1)]))
    return np.where(nearest <= max_gap, filled, np.nan)


def build_grid(h, sources):
    ks = [k for s in sources for k in h.hr[s]]
    t0, t1 = min(ks), max(ks)
    minutes = np.arange(t0, t1 + 1)
    hr, steps = {}, {}
    for s in sources:
        a = np.full(len(minutes), np.nan)
        for k, v in h.hr[s].items():
            a[k - t0] = v
        hr[s] = regularise(a)
        b = np.zeros(len(minutes))
        for k, v in h.steps[s].items():
            if t0 <= k <= t1:
                b[k - t0] += v
        steps[s] = b
    return minutes, hr, steps


def score_all(chans, segs, tmpl, combine):
    """Combined score per minute; NaN where not scorable.

    `combine` selects how independent sensors are fused:
      source name / int i -- trust that one sensor alone
      "mean" -- average the sensors that are available (noise averages down)
      "min"  -- require agreement; a candidate is only as strong as its weakest
                corroborating sensor, so a motion artifact on one wrist is rejected.
                WARNING: agreement also means a minute is unscorable unless BOTH
                sensors have a complete window there. On this subject that discarded
                95% of awake time (1,352 of 29,649 minutes) and 5 of 16 labels, which
                cost far more than the artifact rejection gained.
      "max"  -- take the strongest sensor. Included for completeness, but note it
                actively selects for whichever sensor is noisiest at that instant.
    """
    names = list(chans)
    if isinstance(combine, str) and combine in names:
        combine = names.index(combine)
    per = [chans[s].match(segs[s], tmpl, PRE) for s in names]
    C = np.vstack([p[0] for p in per])
    A = np.vstack([p[1] for p in per])
    fin = np.isfinite(C) & np.isfinite(A)
    if isinstance(combine, int):
        ok = fin[combine]
        c, a = C[combine], A[combine]
    elif combine == "min":
        ok = fin.all(axis=0)
        with np.errstate(invalid="ignore"):
            c = np.min(np.where(fin, C, np.inf), axis=0)
            a = np.min(np.where(fin, A, np.inf), axis=0)
    elif combine == "mean":
        ok = fin.any(axis=0)
        n = np.maximum(fin.sum(axis=0), 1)
        c = np.where(fin, C, 0.0).sum(axis=0) / n
        a = np.where(fin, A, 0.0).sum(axis=0) / n
    else:
        ok = fin.any(axis=0)
        with np.errstate(invalid="ignore"):
            c = np.max(np.where(fin, C, -np.inf), axis=0)
            a = np.max(np.where(fin, A, -np.inf), axis=0)
    return np.where(ok & (c > 0) & (a > 0), c * a, np.nan)


def nms(minutes, score, sep=SEP):
    """Full non-maximum-suppressed candidate list, descending by score.

    Thresholding later is just a prefix of this list: NMS accepts in descending score
    order and acceptance depends only on already-accepted higher-scoring candidates, so
    lowering a threshold can add candidates but never revoke one.
    """
    idx = np.flatnonzero(np.isfinite(score))
    if idx.size == 0:
        return []
    idx = idx[np.argsort(-score[idx])]
    keep = []
    for i in idx:
        t = int(minutes[i])
        if all(abs(t - k) > sep for k, _ in keep):
            keep.append((t, float(score[i])))
    return keep


def at_threshold(cands, thr):
    return [c for c in cands if c[1] >= thr]


# ----------------------------------------------------------------- matching / nulls
def hits(alarms, labels, tol):
    """Greedy one-to-one matching: two alarms cannot both claim the same cigarette."""
    used, n = set(), 0
    for L in sorted(labels):
        best, bd = None, None
        for i, (t, _) in enumerate(alarms):
            if i in used:
                continue
            d = abs(t - L)
            if d <= tol and (bd is None or d < bd):
                best, bd = i, d
        if best is not None:
            used.add(best)
            n += 1
    return n


def shift_null(alarms, labels, tol, span, trials, rng):
    """Chance from rigidly rotating the alarm train within the day."""
    if not alarms or not labels:
        return 0.0, 1.0
    lo, hi = span
    width = max(1, hi - lo + 1)
    obs = hits(alarms, labels, tol)
    draws = []
    for _ in range(trials):
        d = rng.randrange(width)
        draws.append(hits([((t - lo + d) % width + lo, s) for t, s in alarms], labels, tol))
    exp = sum(draws) / len(draws)
    p = sum(1 for x in draws if x >= obs) / len(draws)
    return exp, p


def budget_threshold(cands_by_day, days, budget):
    """Threshold admitting at most `budget` alarms across `days`.

    The budget is the number of cigarettes actually logged on those days, so the operating
    point comes from the subject's own observed rate rather than a guess.
    """
    pool = sorted((c[1] for d in days for c in cands_by_day[d]), reverse=True)
    if not pool:
        return np.inf
    return pool[budget - 1] if 0 < budget <= len(pool) else (pool[-1] if budget else np.inf)


# ----------------------------------------------------------------- driver
def main():
    rng = random.Random(12345)
    pts, ivs = load_labels()
    h = load_apple_health(default_export())

    span_lo, span_hi = h.m(min(pts)) - 1440, h.m(max(pts)) + 1440
    sources = [s for s in h.sources()
               if sum(1 for k in h.hr[s] if span_lo <= k <= span_hi) > 500]
    minutes, hr, steps = build_grid(h, sources)
    t0 = int(minutes[0])
    print(f"sources: {sources}")
    print(f"grid: {len(minutes):,} minutes  {h.stamp(t0)} -> {h.stamp(int(minutes[-1]))}")

    labels_by_day = defaultdict(list)
    for p in pts:
        labels_by_day[p.date()].append(h.m(p))
    days = sorted(labels_by_day)
    all_days = sorted({h.stamp(int(k)).date() for k in minutes})
    day_slices = {}
    for d in all_days:
        lo = h.m(datetime(d.year, d.month, d.day)) - t0
        day_slices[d] = slice(max(0, lo), min(len(minutes), lo + 1440))
    free = [(h.m(a), h.m(b)) for a, b, k, _ in ivs if k == "free"]

    print(f"labelled days: {[str(d) for d in days]} "
          f"({sum(len(v) for v in labels_by_day.values())} cigarettes)")

    GRID = dict(bw=[30, 60, 120], st=[3, 10, 30], onset=[0, 4, 8, 12],
                tr=[1.0, 2.5], td=[6, 12, 24], cross=[False, True])
    keys = list(itertools.product(*GRID.values()))
    print(f"\nscoring {len(keys)} hyperparameter combinations ...")

    cands = {}
    for bw, st in itertools.product(GRID["bw"], GRID["st"]):
        chans = {s: Channel(minutes, hr[s], steps[s], baseline_window=bw, step_tau=st)
                 for s in sources}
        segs = {s: c.segments(PRE, POST) for s, c in chans.items()}
        for onset, tr, td, cross in itertools.product(GRID["onset"], GRID["tr"],
                                                      GRID["td"], GRID["cross"]):
            sc = score_all(chans, segs, template(onset, tr, td, PRE, POST), cross)
            cands[(bw, st, onset, tr, td, cross)] = {
                d: nms(minutes[day_slices[d]], sc[day_slices[d]]) for d in all_days}
    print(f"done: {len(cands)} score maps\n")

    def recall_over(key, tol, eval_days):
        c = cands[key]
        got = tot = 0
        for d in eval_days:
            tr_days = [x for x in days if x != d]
            thr = budget_threshold(c, tr_days, sum(len(labels_by_day[x]) for x in tr_days))
            got += hits(at_threshold(c[d], thr), labels_by_day[d], tol)
            tot += len(labels_by_day[d])
        return got, tot

    print("=" * 78)
    print("NESTED LEAVE-ONE-DAY-OUT  (hyperparameters never see the evaluated day)")
    print("=" * 78)
    results = {}
    for tol in TOLS:
        got = tot = 0
        detail = []
        for d in days:
            inner = [x for x in days if x != d]
            bkey = max(cands, key=lambda k: recall_over(k, tol, inner)[0])
            c = cands[bkey]
            thr = budget_threshold(c, inner, sum(len(labels_by_day[x]) for x in inner))
            al = at_threshold(c[d], thr)
            hgot = hits(al, labels_by_day[d], tol)
            got += hgot
            tot += len(labels_by_day[d])
            detail.append((d, bkey, len(al), hgot, len(labels_by_day[d]), thr))
        results[tol] = (got, tot, detail)
        print(f"  +/-{tol:2d} min   nested recall {got}/{tot} = {100*got/max(1,tot):.0f}%")

    print("\nper-day detail at +/-10 min:")
    for d, k, na, hg, nl, thr in results[10][2]:
        print(f"  {d}  {na:3d} alarms  {hg}/{nl} caught  thr={thr:6.2f}  "
              f"bw={k[0]} st={k[1]} onset={k[2]} tr={k[3]} td={k[4]} cross={k[5]}")

    json.dump({"sources": sources, "n_combos": len(keys),
               "nested_recall": {str(t): [results[t][0], results[t][1]] for t in TOLS}},
              open(os.path.join(ROOT, "results", "v2_meta.json"), "w"), indent=2)


if __name__ == "__main__":
    main()
