"""Signal measurements for the AIME secondary analysis (PT Technologies Research, R-001).

Every metric here is gain-invariant, so no loudness normalisation is applied:
scaling a signal changes none of them.

Each file is analysed at its own sample rate. All results are in absolute hertz, so
nothing is lost: a file simply has no content above its Nyquist frequency. (Resampling
everything to one rate first was tried and rejected - a polyphase filter short enough to
run fast leaves images a few tens of dB down, which reads as bandwidth the source does
not have.)

Per track:
  1. decode the WAV bytes as stored in the dataset (native rate, channels, bit depth)
  2. take the 10 s window of highest RMS (the rule AIME itself used for its listening clips)
  3. STFT: Hann window of ~46 ms (2048 points at 44.1 kHz; the nearest power of two at
     other rates), hop 1/4 window
  4. spectral metrics per channel (so a stereo image cannot cancel itself out of the
     spectrum), then averaged; stereo metrics on the two channels
"""
import io
import numpy as np
from scipy.io import wavfile
from scipy.signal import get_window


def decode(wav_bytes):
    sr, x = wavfile.read(io.BytesIO(wav_bytes))
    bits = x.dtype.itemsize * 8
    if x.dtype.kind == "i":
        x = x.astype(np.float64) / float(np.iinfo(x.dtype).max)
    elif x.dtype.kind == "u":
        x = (x.astype(np.float64) - 128.0) / 128.0
    else:
        x = x.astype(np.float64)
    if x.ndim == 1:
        x = x[:, None]
    return sr, x, bits


def loudest_window(x, sr, seconds=10.0):
    n = int(seconds * sr)
    if len(x) <= n:
        return x
    e = np.sum(x ** 2, axis=1)
    c = np.concatenate([[0.0], np.cumsum(e)])
    starts = np.arange(0, len(x) - n + 1, sr // 10)  # 0.1 s steps
    s = int(starts[int(np.argmax(c[starts + n] - c[starts]))])
    return x[s:s + n]


def analyse(wav_bytes):
    sr, x, bits = decode(wav_bytes)
    out = {"sr": sr, "channels": x.shape[1], "bits": bits, "seconds": round(len(x) / sr, 2)}
    seg = loudest_window(x, sr)
    stereo = seg.shape[1] >= 2
    dual_mono = bool(stereo and np.allclose(seg[:, 0], seg[:, 1], atol=1e-6))
    out["dual_mono"] = dual_mono
    nyq = sr / 2

    nfft = int(2 ** round(np.log2(sr * 2048 / 44100)))
    hop = nfft // 4
    freqs = np.fft.rfftfreq(nfft, 1 / sr)
    win = get_window("hann", nfft)

    def power_stft(ch):
        frames = 1 + (len(ch) - nfft) // hop
        idx = np.arange(nfft)[None, :] + hop * np.arange(frames)[:, None]
        return np.abs(np.fft.rfft(ch[idx] * win, axis=1)) ** 2

    chans = [power_stft(seg[:, c]) for c in range(1 if dual_mono else min(2, seg.shape[1]))]
    P = np.mean(chans, axis=0)
    frame_e = P.sum(axis=1)
    live = frame_e > frame_e.max() * 1e-6  # drop frames more than 60 dB below the loudest
    P = P[live]
    chans = [c[live] for c in chans]

    # 1. spectral roll-off: the frequency below which 85 % (and 99 %) of a frame's energy lies
    cum = np.cumsum(P, axis=1)
    for g in (0.85, 0.99):
        k = np.argmax(cum >= g * cum[:, -1:], axis=1)
        out[f"rolloff{int(g * 100)}"] = float(np.mean(freqs[k]))

    # 2. long-term average spectrum and the bandwidth it spans: the highest frequency at which
    #    it is within 60 dB (and 80 dB) of its own peak, after ~100 Hz smoothing
    ltas = P.mean(axis=0)
    ldb = 10 * np.log10(ltas + 1e-30)
    w = max(3, int(round(100 / (freqs[1] - freqs[0]))) | 1)
    ldb_s = np.convolve(np.pad(ldb, w // 2, mode="edge"), np.ones(w) / w, mode="valid")
    peak = ldb_s.max()
    for depth in (60, 80):
        above = np.nonzero(ldb_s > peak - depth)[0]
        out[f"bw{depth}"] = float(freqs[above[-1]]) if len(above) else 0.0
    # is the band edge simply the file's Nyquist frequency?
    out["nyquist_limited"] = bool(out["bw60"] >= nyq - 400)
    # steepness of the edge: mean level 250-750 Hz below the -60 dB edge minus 250-750 Hz above it
    f60 = out["bw60"]
    lo = ldb_s[(freqs > f60 - 750) & (freqs < f60 - 250)]
    hi = ldb_s[(freqs > f60 + 250) & (freqs < f60 + 750)]
    out["edge_db_per_khz"] = float(lo.mean() - hi.mean()) if len(lo) and len(hi) and not out["nyquist_limited"] else None
    # share of energy above 16 kHz (None when the file cannot hold any)
    hf = ltas[freqs >= 16000].sum()
    out["hf16_db"] = float(10 * np.log10(hf / ltas.sum())) if hf > 0 else None

    # 3. spectral flatness (Gray & Markel): geometric over arithmetic mean of the power
    #    spectrum, per channel then averaged, in a band every source holds (50 Hz - 7.5 kHz)
    #    and, where the file reaches it, in 16-20 kHz
    def flat(band):
        vals = []
        for C in chans:
            b = C[:, band] + 1e-20
            vals.append(np.mean(np.exp(np.mean(np.log(b), axis=1)) / np.mean(b, axis=1)))
        return float(np.mean(vals))
    out["flatness_low"] = flat((freqs >= 50) & (freqs <= 7500))
    out["flatness_16_20k"] = flat((freqs >= 16000) & (freqs <= 20000)) if out["bw60"] > 20000 else None

    # 4. stereo: zero-lag inter-channel correlation over the window, its short-term
    #    distribution (46 ms frames), and the level lost when L and R are summed to mono
    if stereo and not dual_mono:
        L, R = seg[:, 0], seg[:, 1]
        out["rho"] = float(np.sum(L * R) / np.sqrt(np.sum(L * L) * np.sum(R * R) + 1e-30))
        fw, fh = nfft, nfft // 2
        n = 1 + (len(L) - fw) // fh
        idx = np.arange(fw)[None, :] + fh * np.arange(n)[:, None]
        l, r = L[idx], R[idx]
        el, er = np.sum(l * l, axis=1), np.sum(r * r, axis=1)
        ok = (el + er) > (el + er).max() * 1e-6
        rho_t = np.sum(l * r, axis=1)[ok] / np.sqrt(el[ok] * er[ok] + 1e-30)
        out["rho_frames_lt_0_2"] = float(np.mean(rho_t < 0.2))
        out["rho_frames_lt_0"] = float(np.mean(rho_t < 0.0))
        m = 0.5 * (L + R)
        out["mono_sum_db"] = float(10 * np.log10(np.mean(m * m) / (0.5 * (np.mean(L * L) + np.mean(R * R))) + 1e-30))
    return out
