#!/usr/bin/env python3
"""Decode grp=0x0E register-read values across captures to characterize
what each register means.

For each reg_id, show:
  * Value bytes across captures (BE u32, LE u32, signed)
  * Whether constant or varies
  * Whether it matches a known wheel-setting value (brightness 0..100,
    color, RPM range, etc.)
  * Group reg_ids by category (high byte) to suggest banks
"""
import sys, struct
from pathlib import Path
from collections import defaultdict

sys.path.insert(0, str(Path(__file__).parent))
from moza_bridge import load_bridge


def collect_replies(path):
    """Returns: reg_id -> [value_bytes] over the capture's first 15s"""
    frames = load_bridge(path)
    if not frames: return {}
    # anchor at first h2b sess=01 open
    anchor = None
    for f in frames:
        if f.dir == 'h2b' and f.is_session_data and f.sess_id == 1 and f.sess_type == 0x81:
            anchor = f.t_rel; break
    if anchor is None: return {}
    out = defaultdict(list)
    for f in frames:
        if f.t_rel - anchor < 0 or f.t_rel - anchor > 15: continue
        if f.dir != 'b2h' or f.grp != 0x8E or f.dev != 0x21: continue
        if len(f.payload) < 7 or f.payload[0] != 0x00: continue
        reg = f.payload[1] | (f.payload[2] << 8)
        val = bytes(f.payload[3:7])
        out[reg].append(val)
    return out


def interpret_value(val_bytes):
    """Return a list of plausible interpretations for a 4-byte value."""
    be_u32 = struct.unpack('>I', val_bytes)[0]
    le_u32 = struct.unpack('<I', val_bytes)[0]
    interps = []
    if val_bytes == b'\x00\x00\x80\x00':
        interps.append("SENTINEL_0x8000 (unset?)")
    elif val_bytes == b'\x00\x00\x00\x00':
        interps.append("ZERO")
    elif val_bytes == b'\xff\xff\xff\xff':
        interps.append("UNDEFINED (-1)")
    else:
        # u16 candidates (only low 2 bytes when val is `00 00 XX YY`)
        if val_bytes[0] == 0 and val_bytes[1] == 0:
            v16_be = (val_bytes[2] << 8) | val_bytes[3]
            v16_le = val_bytes[2] | (val_bytes[3] << 8)
            tags = []
            if 0 <= v16_be <= 100: tags.append(f"u8?={v16_be} (range 0-100, e.g. brightness/percent)")
            if v16_be < 360: tags.append(f"u16_be={v16_be} (could be angle/range)")
            if v16_be < 12000: tags.append(f"u16_be={v16_be} (could be RPM/ms)")
            interps.append(f"BE2={v16_be:5d} (0x{v16_be:04x})  LE2={v16_le:5d}{' [' + ', '.join(tags) + ']' if tags else ''}")
        else:
            interps.append(f"BE={be_u32:11d} (0x{be_u32:08x})  LE={le_u32:11d}")
    return interps


def main():
    base = Path("/home/rorth/src/moza-simhub-plugin/sim/logs")
    paths = sorted(base.glob("bridge-*.jsonl"))[:6]

    all_data = {}
    for p in paths:
        if p.stat().st_size == 0: continue
        all_data[p.name] = collect_replies(p)

    # Get union of all reg_ids
    all_regs = set()
    for cap_data in all_data.values():
        all_regs.update(cap_data.keys())

    # Group by category (high byte of reg_id)
    by_cat = defaultdict(list)
    for r in all_regs:
        cat = (r >> 8) & 0xFF
        idx = r & 0xFF
        by_cat[idx].append((cat, r))
    for idx in by_cat:
        by_cat[idx].sort()

    print(f"Captures analyzed: {len(all_data)}")
    print(f"Distinct registers seen: {len(all_regs)}")

    for idx in sorted(by_cat):
        regs_in_idx = by_cat[idx]
        print(f"\n  ── BANK index={idx} (reg low byte) — {len(regs_in_idx)} registers, categories {hex(regs_in_idx[0][0])}..{hex(regs_in_idx[-1][0])}")
        for cat, reg in regs_in_idx:
            # First value seen in each capture for this reg
            vals = []
            for cap_name, cap_data in all_data.items():
                if reg in cap_data:
                    vals.append((cap_name[:25], cap_data[reg][0]))
            if not vals: continue
            unique_vals = set(v.hex() for _, v in vals)
            marker = "VARIES" if len(unique_vals) > 1 else "const "
            sample = vals[0][1]
            interp = '; '.join(interpret_value(sample))
            print(f"    reg=0x{reg:04x} cat=0x{cat:02x}  {marker} sample={sample.hex()}  → {interp}")
            if marker == "VARIES":
                # Show the variety
                for cap_name, v in vals:
                    print(f"        {cap_name:<28} {v.hex()}")


if __name__ == '__main__':
    main()
