#!/usr/bin/env python3
"""Correlate the AB9 host-rendered engine-vibration stream against ground-truth
RPM telemetry captured in the SAME USB pcapng/JSONL.

Background
----------
The capture `usb-capture/AB9/ab9-pithouse-engine-vibration-intensity-2.pcapng`
records PitHouse driving a real AB9 (group 0x20, dev 0x12) WHILE the wheel
dashboard receives live telemetry (group 0x43, dev 0x17, cmd 0x7D 0x23).
Because the freq slider was held at 100 Hz, the oscillator period satisfies
    period = K / (rpm * freq)   =>   period * (rpm/maxRpm) = const
which lets us validate the RPM field empirically and recover PitHouse's actual
encoding laws instead of inferring them.

Findings reproduced by this tool (2026-05-31, intensity-2 capture):
  * Engine-vib stream sub-cmd is 0x0A 0x04 in this capture (was 0x0A 0x05 in the
    2026-05-13/24 captures); pulse pair is 0x0B 0x01 (ON, amp16=0x2328) +
    0x0B 0x02 (OFF, amp16=0x0000) (was 0x0B 0x02/0x03).
  * INTENSITY is encoded in the 16-bit "slot" field (0x0A 0x04 payload off 2-3),
    LINEARLY:  slot = round(intensity_percent * 65.5)  (100%->0x1996=6550,
    60%->0x0F5A=3930, 40%->0x0A3C=2620). The prior "DirectInput effect handle"
    interpretation was wrong — it only ever observed 0x1996 (100%) / 0x0000 (0%).
  * Pulse-pair emission rate is CONSTANT ~48 Hz, independent of rpm AND intensity.
    amp16 is constant 0x2328. Intensity is NOT carried by pulse rate.
  * Period (off 11-13, 24-bit BE) = K/(rpm*freq); period*(rpm/maxRpm) is constant.
    K is recoverable as period*(rpm/maxRpm)*maxRpm*freq once maxRpm is known.

The tier-30 telemetry payload (flag 0x3d, 4 data bytes) for this capture's
dashboard carries a 10-bit NORMALIZED rpm channel at bits 0-9: fraction =
raw/1000 = rpm/maxRpm. (Validated here: it minimises CV(period*fraction).)

Usage
-----
    tools/ab9-rpm-correlate <jsonl> [--maxrpm N] [--freq HZ] [--tier-flag 0x3d]

If --maxrpm is given, absolute K and absolute RPM are reported; otherwise the
rpm-fraction-relative results (which need no redline) are reported.
"""
from __future__ import annotations
import argparse, bisect, json, sys
from collections import Counter
from statistics import mean, median, pstdev


def load(path):
    """Return dict of decoded streams keyed by name."""
    vib = []      # (t, period, slot)
    pulse_on = [] # (t, phase16)
    frac = []     # (t, fraction)  rpm/maxRpm from tier-30 bits 0-9
    tier_flag = None
    with open(path) as fh:
        for line in fh:
            r = json.loads(line)
            b = bytes.fromhex(r["hex"])
            if len(b) < 4:
                continue
            grp, dev, t, d = b[2], b[3], r["t"], r["dir"]
            if d == "h2b" and grp == 0x20 and dev == 0x12:
                pl = b[4:-1]
                if len(pl) >= 15 and pl[0] == 0x0a and pl[1] == 0x04:
                    slot = (pl[2] << 8) | pl[3]
                    period = (pl[11] << 16) | (pl[12] << 8) | pl[13]
                    vib.append((t, period, slot))
                elif len(pl) >= 9 and pl[0] == 0x0b and pl[1] == 0x01:
                    phase = (pl[5] << 8) | pl[6]
                    pulse_on.append((t, phase))
            elif d == "h2b" and grp == 0x43 and dev == 0x17 and len(b) >= 6 \
                    and b[4] == 0x7d and b[5] == 0x23:
                pl = b[4:-1]   # 7d 23 32 00 23 32 [flag] 20 [data]
                if len(pl) >= 8 and len(pl[8:]) == 4:
                    flag = pl[6]
                    # tier-30 is the high-rate 4-byte tier; pick the dominant one
                    raw10 = pl[8] | ((pl[9] & 0x03) << 8)
                    frac.append((t, flag, raw10 / 1000.0))
    return vib, pulse_on, frac


def dominant_flag(frac):
    c = Counter(f for _, f, _ in frac)
    return c.most_common(1)[0][0] if c else None


def main(argv=None):
    ap = argparse.ArgumentParser(description=__doc__,
                                 formatter_class=argparse.RawDescriptionHelpFormatter)
    ap.add_argument("jsonl")
    ap.add_argument("--maxrpm", type=float, default=None,
                    help="car redline (SimHub MaxRpm) to report absolute K / RPM")
    ap.add_argument("--freq", type=float, default=100.0,
                    help="engine-vib frequency slider Hz during capture (default 100)")
    ap.add_argument("--tier-flag", type=lambda s: int(s, 0), default=None,
                    help="telemetry tier flag byte (default: most common)")
    args = ap.parse_args(argv)

    vib, pulse_on, frac_all = load(args.jsonl)
    if not vib:
        print("no 0x0A 0x04 engine-vib frames found", file=sys.stderr)
        return 1
    flag = args.tier_flag if args.tier_flag is not None else dominant_flag(frac_all)
    frac = [(t, f) for t, fl, f in frac_all if fl == flag]
    print(f"engine-vib(0a04)={len(vib)}  pulse-on(0b01)={len(pulse_on)}  "
          f"tier-30 telemetry(flag={flag:#04x})={len(frac)}")
    if not frac:
        print("no telemetry frames for tier flag", file=sys.stderr)
        return 1

    ft = [t for t, _ in frac]
    fv = [f for _, f in frac]

    def frac_at(t, tol=0.1):
        i = bisect.bisect_left(ft, t)
        cands = [j for j in (i - 1, i) if 0 <= j < len(ft)]
        if not cands:
            return None
        j = min(cands, key=lambda k: abs(ft[k] - t))
        return fv[j] if abs(ft[j] - t) <= tol else None

    # ---- Intensity = slot field ----
    print("\n== Intensity encoding (0a04 slot, off 2-3) ==")
    slots = Counter(s for _, _, s in vib)
    for s, c in slots.most_common(8):
        print(f"  slot={s:#06x}={s:5d}  intensity={s/65.5:6.2f}%  n={c}")
    print("  => slot = round(intensity_percent * 65.5)")

    # ---- Period law: period * fraction (= K/(maxRpm*freq)) ----
    prods = []
    for t, period, _ in vib:
        f = frac_at(t)
        if f and f > 0.1:
            prods.append(period * f)
    med = median(prods)
    print("\n== Oscillator period law: period = K/(rpm*freq) ==")
    print(f"  period*fraction: median={med:,.0f}  CV={pstdev(prods)/mean(prods):.4f}  n={len(prods)}")
    print(f"  (low CV confirms period proportional to 1/rpm)")
    if args.maxrpm:
        K = med * args.maxrpm * args.freq
        print(f"  maxRpm={args.maxrpm:.0f}, freq={args.freq:.0f}Hz => K = {K:.4e}")
    else:
        print(f"  K = {med:,.0f} * maxRpm * freq   (supply --maxrpm to resolve)")
        for mx in (7000, 7500, 7700, 8000, 8500):
            print(f"      maxRpm={mx}: K = {med*mx*args.freq:.3e}")

    # ---- Pulse-pair rate (constant?) ----
    pon = sorted(t for t, _ in pulse_on)
    dts = [pon[i + 1] - pon[i] for i in range(len(pon) - 1)
           if 0 < pon[i + 1] - pon[i] < 0.2]
    if dts:
        print("\n== Pulse-pair emission rate (0b01/0b02) ==")
        print(f"  median inter-pair dt={median(dts)*1000:.1f}ms  => {1/median(dts):.1f} Hz")
        # rate vs fraction
        bins = {}
        po = [t for t, _ in pulse_on]
        po.sort()
        for i in range(len(po) - 1):
            dt = po[i + 1] - po[i]
            if not (0 < dt < 0.2):
                continue
            f = frac_at(po[i])
            if f is None:
                continue
            bins.setdefault(round(f * 10) / 10, []).append(dt)
        print("  rate vs rpm-fraction (flat => rpm-independent):")
        for k in sorted(bins):
            print(f"    frac~{k:.1f}: {1/median(bins[k]):5.1f} Hz  n={len(bins[k])}")

    # ---- Pulse phase-counter step vs rpm ----
    if len(pulse_on) > 2:
        steps = {}
        for i in range(len(pulse_on) - 1):
            t0p, p0 = pulse_on[i]
            t1p, p1 = pulse_on[i + 1]
            if t1p - t0p > 0.2:
                continue
            step = (p1 - p0) & 0xFFFF
            if step > 0x8000:
                step -= 0x10000
            f = frac_at(t0p)
            if f is None:
                continue
            steps.setdefault(round(f * 10) / 10, []).append(step)
        if steps:
            print("\n== Pulse phase-counter step per pair vs rpm-fraction ==")
            for k in sorted(steps):
                print(f"    frac~{k:.1f}: step median={median(steps[k]):7.0f}  n={len(steps[k])}")

    return 0


if __name__ == "__main__":
    sys.exit(main())
