"""Lock-in / PRBS demodulation of a Thermal Master P3 raw stack, joined to ettelem telemetry.

temp : float32 (M,192,256) degC = thermal_u16/64.0 - 273.15   (rows 194..385 only),
       already registered and boxcar-decimated to ~2 Hz
t    : float64 (M,) host monotonic seconds -- REAL timestamps, never frame indices
seg  : int32  (M,) increments at every detected FFC/shutter event
good : bool   (M,) False for shutter epochs (+2 s) and cnt3 gaps -- drop, never interpolate
extra: (M,k)  covariates: the REALIZED launch indicator or board_w, the residual board_w
              amplitude term, camera body temperature, the ambient tab ROI
"""
import numpy as np


def design(t, freqs, harmonics=3, seg=None, seg_order=1, drift_order=2, extra=None, res_hz=None):
    """[per-FFC-segment offset AND slope | covariates | cos/sin per (freq, harmonic)].

    With seg given, NO global drift polynomial: the per-segment piecewise-linear basis
    already spans every global affine function, so adding drift^1 makes X exactly
    rank-deficient (cond ~1e16) and np.linalg.inv returns garbage without raising.
    seg_order >= 1 is mandatory: without the slope, sigma inflates 5x and a null pixel
    reads 2.7 mK against a true 3 mK channel.
    """
    cols, names = [], []
    t = np.asarray(t, float)
    if seg is None:
        u = 2 * (t - t.mean()) / max(np.ptp(t), 1e-9)
        for p in range(drift_order + 1):
            cols.append(u ** p); names.append(f"drift^{p}")
    else:
        for k in np.unique(seg):
            m = seg == k
            tc = np.zeros_like(t); tc[m] = t[m] - t[m].mean()
            sc = max(np.ptp(t[m]), 1e-9)
            for q in range(seg_order + 1):
                c = np.zeros_like(t); c[m] = (tc[m] / sc) ** q
                cols.append(c); names.append(f"seg{k}^{q}")
    if extra is not None:
        for j, col in enumerate(np.atleast_2d(np.asarray(extra, float).T)):
            cols.append(col - col.mean()); names.append(f"extra{j}")
    tones, res = {}, res_hz if res_hz else 1.0 / max(np.ptp(t), 1e-9)
    for f in freqs:
        for h in range(1, harmonics + 1):
            fh = f * h
            for g in tones:                       # key by (f,h), never by formatted name
                if abs(g - fh) < 2 * res:
                    raise ValueError(f"tone collision: {fh:g} Hz vs {g:g} Hz (res {res:.2g} Hz)")
            tones[fh] = len(cols)
            w = 2 * np.pi * fh
            cols += [np.cos(w * t), np.sin(w * t)]
            names += [f"cos{fh:.6g}", f"sin{fh:.6g}"]
    return np.column_stack(cols), names, tones


def demodulate(temp, t, freqs, good=None, seg=None, ref_mask=None, harmonics=3,
               seg_order=1, drift_order=2, extra=None, block=512, cond_max=1e8):
    """Per-pixel LSQ lock-in at several frequencies, chunked over pixels.
    Returns {freq: dict(amp[K], phase[deg lag], sigma, snr)} plus '_cond'."""
    if good is None: good = np.ones(len(t), bool)
    tg = t[good]; sg = None if seg is None else seg[good]
    eg = None if extra is None else np.asarray(extra)[good]
    X, names, tones = design(tg, freqs, harmonics, sg, seg_order, drift_order, eg)
    s = np.linalg.svd(X, compute_uv=False)
    cond = s[0] / s[-1]
    assert cond < cond_max, f"design matrix ill-conditioned: cond={cond:.3g}"
    XtXi = np.linalg.inv(X.T @ X)                 # safe only after the cond assertion
    P = np.linalg.pinv(X)
    H, W = temp.shape[1], temp.shape[2]
    npx, dof = H * W, len(tg) - X.shape[1]
    beta = np.empty((X.shape[1], npx)); sigma = np.empty(npx)
    ref = None
    if ref_mask is not None:                      # common mode: unpowered taped patch
        ref = temp[good][:, ref_mask].reshape(len(tg), -1).mean(axis=1)
    flat = temp[good].reshape(len(tg), npx)
    for i in range(0, npx, block):                # 512 px x 3600 frames = 15 MB, Pi-safe
        Y = flat[:, i:i + block].astype(np.float64)
        if ref is not None: Y -= ref[:, None]
        b = P @ Y; r = Y - X @ b
        beta[:, i:i + block] = b
        sigma[i:i + block] = np.sqrt((r * r).sum(0) / dof)
    out = {"_cond": cond}
    for fh, i in tones.items():
        c, sn = beta[i], beta[i + 1]
        sig = sigma * np.sqrt(0.5 * (XtXi[i, i] + XtXi[i + 1, i + 1]))
        amp = np.hypot(c, sn)
        out[fh] = dict(amp=amp.reshape(H, W),
                       phase=np.degrees(np.arctan2(-sn, c)).reshape(H, W),
                       sigma=sig.reshape(H, W),
                       snr=(amp / np.maximum(sig, 1e-12)).reshape(H, W))
    return out


def mseq(n, taps):
    """Maximal-length LFSR, +-1, L = 2**n-1. n=7 taps=[6]; n=9 [5]; n=10 [7].
    Periodic autocorrelation is exactly L at lag 0 and -1 elsewhere; check s.sum()==1."""
    mask = L = (1 << n) - 1; reg, out = 1, []
    for _ in range(L):
        out.append(1.0 if (reg & 1) else -1.0)
        fb = reg & 1
        for tp in taps: fb ^= (reg >> (n - tp)) & 1
        reg = ((reg >> 1) | (fb << (n - 1))) & mask
    return np.array(out)


def prbs_impulse(temp, t, seq, chip, t0, good=None, detrend=2):
    """Per-pixel thermal impulse response by periodic cross-correlation.
    Per-lag noise ~ NETD/sqrt(M) -- sqrt(2) better than single-frequency lock-in, and
    L lags instead of one amplitude. cumsum(h) is the step response; fit Foster taus to
    that. DC gain is lost with the mean -- take it from E-T13."""
    if good is None: good = np.ones(len(t), bool)
    L = len(seq); period = L * chip
    tg = t[good]
    b = np.floor(((tg - t0) % period) / chip).astype(int)
    Y = temp[good].reshape(len(tg), -1).astype(np.float64)
    V = np.vander(tg - tg.mean(), detrend + 1)
    Y = Y - V @ np.linalg.lstsq(V, Y, rcond=None)[0]
    acc = np.zeros((L, Y.shape[1])); np.add.at(acc, b, Y)
    cnt = np.bincount(b, minlength=L).astype(float)
    y = acc / np.maximum(cnt, 1)[:, None]
    y = y - y.mean(0, keepdims=True)
    S = np.fft.rfft(seq)
    h = np.fft.irfft(np.fft.rfft(y, axis=0) * np.conj(S)[:, None], n=L, axis=0) / (L + 1)
    return h.reshape(L, *temp.shape[1:])
