#!/usr/bin/env python3
"""
Does the accelerometer channel actually separate cigarettes from ordinary life?

Tested BEFORE building any detector on it. With 8 labels it is trivially easy to find a
pattern that is not real, so every feature is compared against matched control windows
with a permutation test.
"""
import ast
import csv
import os
import random
import statistics as st
from datetime import datetime, timedelta

random.seed(11)
DS = os.path.join(os.path.dirname(os.path.abspath(__file__)), "..", "claude-science", "datasets")
LOG = os.path.join(os.path.dirname(os.path.abspath(__file__)), "..", "data", "cigarette_log.csv")
EPOCH = datetime(2026, 8, 10)
mins = lambda d: (d - EPOCH).total_seconds() / 60


def load(name, conv):
    out = []
    for r in csv.DictReader(open(os.path.join(DS, name))):
        t = datetime.strptime(r["datetime_local"], "%Y-%m-%dT%H:%M:%S")
        out.append((mins(t), conv(r)))
    out.sort()
    return out


motion = load("motion_1s.csv", lambda r: float(r["motion"]))
orient = load("motion_orientation_30s.csv",
              lambda r: (float(r["mean_x"]), float(r["mean_y"]), float(r["mean_z"])))
steps = load("steps_1min.csv", lambda r: float(r["steps"]))
still = load("stillness.csv", lambda r: r["still"].strip().lower() == "true")

cigs = []
for r in csv.reader(open(LOG)):
    if not r or r[0].startswith(("datetime", "#")):
        continue
    cigs.append(mins(datetime.strptime(r[0], "%Y-%m-%dT%H:%M")))

WIN_A, WIN_B = -2.0, 8.0        # window around onset, minutes


def slice_(arr, t, a=WIN_A, b=WIN_B):
    return [v for m, v in arr if t + a <= m <= t + b]


def autocorr_peak(x, lo=20, hi=90):
    """strongest autocorrelation at lags 20-90 s — the cadence of a repeated puff"""
    n = len(x)
    if n < hi * 2:
        return 0.0
    mu = sum(x) / n
    d = [v - mu for v in x]
    den = sum(v * v for v in d)
    if den < 1e-9:
        return 0.0
    best = 0.0
    for lag in range(lo, min(hi, n // 2)):
        num = sum(d[i] * d[i + lag] for i in range(n - lag))
        best = max(best, num / den)
    return best


def features(t):
    m = slice_(motion, t)
    if len(m) < 300:
        return None
    o = slice_(orient, t)
    s = slice_(steps, t)
    sl = slice_(still, t)
    zs = [z for _, _, z in o] or [0]
    xs = [x for x, _, _ in o] or [0]
    ys = [y for _, y, _ in o] or [0]
    return {
        "motion_mean": st.mean(m),
        "motion_std": st.pstdev(m),
        "periodicity": autocorr_peak(m),
        "orient_z_range": max(zs) - min(zs),
        "orient_spread": st.pstdev(xs) + st.pstdev(ys) + st.pstdev(zs),
        "steps": sum(s),
        "still_frac": (sum(1 for v in sl if v) / len(sl)) if sl else 0.0,
    }


# candidate control windows: any waking minute with motion coverage, not near a cigarette
lo_m, hi_m = motion[0][0], motion[-1][0]
controls = []
for t in range(int(lo_m) + 3, int(hi_m) - 9):
    if any(abs(t - c) < 20 for c in cigs):
        continue
    if 1479 <= t <= 1960:            # asleep 00:39-08:40 on 11 Aug
        continue
    controls.append(t)

cig_f = [f for f in (features(c) for c in cigs) if f]
ctl_f = [f for f in (features(c) for c in controls) if f]
print(f"cigarettes with motion coverage: {len(cig_f)}/{len(cigs)}   control windows: {len(ctl_f)}\n")

KEYS = ["motion_mean", "motion_std", "periodicity", "orient_z_range",
        "orient_spread", "steps", "still_frac"]
print(f"{'feature':>16} {'cigarettes':>12} {'controls':>12} {'p (perm)':>10}  verdict")
TRIALS = 20000
for k in KEYS:
    cv = [f[k] for f in cig_f]
    av = [f[k] for f in ctl_f]
    obs = st.mean(cv)
    ctl_mean = st.mean(av)
    # two-sided permutation: how often do random draws of the same size beat the observed gap?
    hits = 0
    n = len(cv)
    for _ in range(TRIALS):
        samp = st.mean(random.sample(av, n))
        if abs(samp - ctl_mean) >= abs(obs - ctl_mean):
            hits += 1
    p = hits / TRIALS
    verdict = "*** separates" if p < 0.05 else ("~ weak" if p < 0.15 else "no")
    print(f"{k:>16} {obs:>12.2f} {ctl_mean:>12.2f} {p:>10.3f}  {verdict}")

print("\nPer-cigarette detail:")
print(f"  {'time':>16} {'motion':>8} {'period':>8} {'z-range':>8} {'steps':>7} {'still%':>7}")
for c, f in zip([c for c in cigs if features(c)], cig_f):
    h = EPOCH + timedelta(minutes=c)
    print(f"  {h.strftime('%d %b %H:%M'):>16} {f['motion_mean']:>8.1f} {f['periodicity']:>8.2f} "
          f"{f['orient_z_range']:>8.0f} {f['steps']:>7.0f} {100*f['still_frac']:>7.0f}")
