#!/usr/bin/env python3
"""Decode AB9 session JSONL captures.

Usage:
    ab9-decode-session <subcommand> [path] [...]

Subcommands:
    inventory   — total frames, group/cmd/sub-cmd counts + rate-per-second window
    init        — extract early-session frames (FFB init handshake)
    slots       — slot-ID → period scatter for 0x0A 0x05 (resolves Phase 1.C)
    pulse       — 0x0B 0x02/03 16-bit field timeline (resolves Phase 1.A)
    trig5       — 0x0D 0x05 rate model (resolves Phase 1.B)
    lowrate     — 0x08 0x04/06 occurrences + payload (resolves Phase 1.D)
    dump        — pretty frame dump within a time window
    h2b-unique  — unique h2b frame shapes by (grp, sub-cmd), with sample payloads

The path argument defaults to the most recent sim/logs/ab9-*.jsonl.
"""
from __future__ import annotations

import argparse
import sys
from collections import Counter, defaultdict
from pathlib import Path

# Allow running from anywhere — tools/ on sys.path so we can import ab9_session.
sys.path.insert(0, str(Path(__file__).resolve().parent))

from ab9_session import iter_frames, resolve_log, sub_label, AFrame  # noqa: E402


# ────────────────────────────────────────────────────────────────────────────
# inventory
# ────────────────────────────────────────────────────────────────────────────

def cmd_inventory(args: argparse.Namespace) -> int:
    path = resolve_log(args.path)
    print(f"Inventory of {path}\n")

    total = 0
    h2b_count = 0
    b2h_count = 0
    grp_counts: Counter = Counter()           # (dir, grp)
    sub_counts: Counter = Counter()           # (dir, grp, sub_hi, sub_lo) for FFB
    payload_len_hist: dict[tuple[str, int, int, int], Counter] = defaultdict(Counter)
    t_first: float | None = None
    t_last: float | None = None
    sub_first: dict[tuple[str, int, int, int], float] = {}

    for f in iter_frames(path):
        total += 1
        if t_first is None:
            t_first = f.t
        t_last = f.t
        if f.is_h2b:
            h2b_count += 1
        else:
            b2h_count += 1
        grp_counts[(f.dir, f.grp)] += 1

        if f.is_ffb and len(f.payload) >= 1:
            sub_hi = f.payload[0]
            sub_lo = f.payload[1] if len(f.payload) > 1 else -1
            key = (f.dir, f.grp, sub_hi, sub_lo)
            sub_counts[key] += 1
            payload_len_hist[key][len(f.payload)] += 1
            sub_first.setdefault(key, f.t_rel)

    duration = (t_last - t_first) if (t_first and t_last) else 0.0

    print(f"  total frames : {total:,}")
    print(f"  h2b          : {h2b_count:,}")
    print(f"  b2h          : {b2h_count:,}")
    print(f"  duration     : {duration:.1f} s ({duration/60:.1f} min)")
    print()

    print("Group/direction frequencies (top 20):")
    print(f"  {'dir':<4} {'grp':<6} {'count':>10} {'rate Hz':>10}")
    for (d, g), c in grp_counts.most_common(20):
        rate = c / duration if duration > 0 else 0.0
        print(f"  {d:<4} 0x{g:02x}   {c:>10,} {rate:>10.2f}")
    print()

    if sub_counts:
        print("Group 0x20 sub-commands (label, len-hist, first-seen t_rel):")
        print(f"  {'dir':<4} {'sub':<6} {'label':<18} {'count':>10} {'rate Hz':>10}  first    payload-lens")
        for (d, g, sub_hi, sub_lo), c in sub_counts.most_common():
            lab = sub_label(sub_hi, sub_lo)
            rate = c / duration if duration > 0 else 0.0
            sub_repr = f"{sub_hi:02x}/{sub_lo:02x}" if sub_lo != -1 else f"{sub_hi:02x}/--"
            t0 = sub_first[(d, g, sub_hi, sub_lo)]
            lens = payload_len_hist[(d, g, sub_hi, sub_lo)]
            len_str = ", ".join(f"{L}B×{cnt}" for L, cnt in sorted(lens.items()))
            print(f"  {d:<4} {sub_repr:<6} {lab:<18} {c:>10,} {rate:>10.2f}  {t0:>7.2f}  {len_str}")

    return 0


# ────────────────────────────────────────────────────────────────────────────
# init — early-session frames + FFB init handshake isolation
# ────────────────────────────────────────────────────────────────────────────

def cmd_init(args: argparse.Namespace) -> int:
    path = resolve_log(args.path)
    print(f"Early-session frames of {path}\n")

    cap_seconds = args.seconds
    seen_ffb_first = False
    ffb_first_t: float | None = None

    for f in iter_frames(path):
        if f.t_rel > cap_seconds and seen_ffb_first:
            break
        # Always print the first chunk to capture connect-time
        if f.t_rel > cap_seconds and not seen_ffb_first:
            break
        if f.is_ffb and ffb_first_t is None:
            ffb_first_t = f.t_rel
            seen_ffb_first = True
        # Skip pure heartbeats during init dump unless requested
        if f.grp in (0x00, 0x80) and not args.include_heartbeats:
            continue

        sub_str = ""
        if f.is_ffb and len(f.payload) >= 2:
            sub_hi, sub_lo = f.payload[0], f.payload[1]
            sub_str = f"  [{sub_label(sub_hi, sub_lo)}]"

        print(f"  {f.t_rel:>8.4f}  {f.dir}  grp=0x{f.grp:02x} dev=0x{f.dev:02x} "
              f"len={len(f.payload):2d}  {f.payload.hex()}{sub_str}")

    print(f"\nFirst FFB frame at t_rel = {ffb_first_t}")
    return 0


# ────────────────────────────────────────────────────────────────────────────
# slots — slot-ID vs period scatter for 0x0A 0x05
# ────────────────────────────────────────────────────────────────────────────

def cmd_slots(args: argparse.Namespace) -> int:
    path = resolve_log(args.path)
    print(f"Slot/period scatter of 0x0A 0x05 from {path}\n")

    # 0x0A 0x05 wire format:
    #   payload = 0A 05 [SS SS] [00 00 00 00 00 00 00] [PP PP PP] 04 [00 00 00 00]
    # Slot ID at payload[2..4] (BE 16), period at payload[11..14] (BE 24).
    # Per ab9-shifter.md (line 87-91).
    by_slot: dict[int, Counter] = defaultdict(Counter)   # slot_id → period histogram
    slot_total: Counter = Counter()
    period_total: Counter = Counter()
    first_seen: dict[int, float] = {}
    last_seen: dict[int, float] = {}

    for f in iter_frames(path):
        if not f.is_h2b or not f.is_ffb:
            continue
        if len(f.payload) < 19:
            continue
        if f.payload[0] != 0x0A or f.payload[1] != 0x05:
            continue
        slot = (f.payload[2] << 8) | f.payload[3]
        period = (f.payload[11] << 16) | (f.payload[12] << 8) | f.payload[13]
        by_slot[slot][period] += 1
        slot_total[slot] += 1
        period_total[period] += 1
        first_seen.setdefault(slot, f.t_rel)
        last_seen[slot] = f.t_rel

    print(f"  unique slots : {len(by_slot)}")
    print(f"  unique periods: {len(period_total)}")
    print()
    print(f"  {'slot':<8} {'count':>10} {'first':>9} {'last':>9}  period-range (min..max)")
    for slot, c in slot_total.most_common():
        periods = sorted(by_slot[slot].keys())
        rng = f"{periods[0]:#08x} .. {periods[-1]:#08x}" if periods else ""
        first_t = first_seen.get(slot, 0.0)
        last_t = last_seen.get(slot, 0.0)
        print(f"  0x{slot:04x}  {c:>10,} {first_t:>9.2f} {last_t:>9.2f}  {rng}")

    if args.top_periods:
        print("\n  Top period values across all slots:")
        for period, c in period_total.most_common(15):
            print(f"    {period:#08x}  ×{c:,}")
    return 0


# ────────────────────────────────────────────────────────────────────────────
# pulse — 0x0B 0x02/03 16-bit field timeline (Phase 1.A)
# ────────────────────────────────────────────────────────────────────────────

def cmd_pulse(args: argparse.Namespace) -> int:
    path = resolve_log(args.path)
    print(f"0x0B 0x02/03 engine-pulse fields from {path}\n")

    # 0x0B sub-cmds: 22-byte wire frames per ab9-shifter.md (line 68).
    # Payload schema unknown; doc mentions "16-bit field at payload offset 4-7"
    # that ranges 0..65535. Let's dump everything past sub-cmd for sample frames
    # plus extract any candidate 16-bit fields from each fixed offset.

    pulse_on_payloads: list[tuple[float, bytes]] = []
    pulse_off_payloads: list[tuple[float, bytes]] = []
    max_dump = args.dump

    for f in iter_frames(path):
        if not f.is_h2b or not f.is_ffb or len(f.payload) < 2:
            continue
        if f.payload[0] != 0x0B:
            continue
        sub_lo = f.payload[1]
        if sub_lo == 0x02:
            if len(pulse_on_payloads) < max_dump:
                pulse_on_payloads.append((f.t_rel, f.payload))
        elif sub_lo == 0x03:
            if len(pulse_off_payloads) < max_dump:
                pulse_off_payloads.append((f.t_rel, f.payload))
        if len(pulse_on_payloads) >= max_dump and len(pulse_off_payloads) >= max_dump:
            if not args.full_scan:
                break

    def dump(label: str, samples: list[tuple[float, bytes]]) -> None:
        print(f"  ── {label} ({len(samples)} samples) ──")
        for t, p in samples:
            # Args after the 2-byte sub-cmd
            args_bytes = p[2:]
            # Try each 2-byte BE offset, useful for spotting the "varying" field
            print(f"    t={t:>9.4f}  payload={p.hex()}  args={args_bytes.hex()}")
            for off in range(0, len(args_bytes) - 1):
                be = (args_bytes[off] << 8) | args_bytes[off + 1]
                # Only print non-zero
                if be:
                    print(f"        offset {off:2d}-{off+1:2d}: 0x{be:04x} ({be})")
        print()

    dump("0x0B 0x02 (engine-pulse-ON)", pulse_on_payloads)
    dump("0x0B 0x03 (engine-pulse-OFF)", pulse_off_payloads)
    return 0


# ────────────────────────────────────────────────────────────────────────────
# trig5 — 0x0D 0x05 rate timeline
# ────────────────────────────────────────────────────────────────────────────

def cmd_trig5(args: argparse.Namespace) -> int:
    path = resolve_log(args.path)
    print(f"0x0D 0x05 rate over time from {path}\n")

    bucket_s = args.bucket
    buckets: Counter = Counter()
    for f in iter_frames(path):
        if not f.is_h2b or not f.is_ffb or len(f.payload) < 2:
            continue
        if f.payload[0] != 0x0D or f.payload[1] != 0x05:
            continue
        b = int(f.t_rel // bucket_s)
        buckets[b] += 1

    print(f"  bucket = {bucket_s} s, {len(buckets)} non-empty buckets\n")
    print(f"  {'t_start':>9}  {'count':>6}  {'Hz':>6}  histogram")
    for b in sorted(buckets):
        c = buckets[b]
        rate = c / bucket_s
        bar = "#" * min(int(rate), 80)
        print(f"  {b*bucket_s:>9.1f}  {c:>6}  {rate:>6.2f}  {bar}")
    return 0


# ────────────────────────────────────────────────────────────────────────────
# lowrate — 0x08 0x04/06 occurrence + payload
# ────────────────────────────────────────────────────────────────────────────

def cmd_lowrate(args: argparse.Namespace) -> int:
    path = resolve_log(args.path)
    print(f"0x08 0x04/06 frames from {path}\n")

    for f in iter_frames(path):
        if not f.is_h2b or not f.is_ffb or len(f.payload) < 2:
            continue
        if f.payload[0] != 0x08:
            continue
        sub_lo = f.payload[1]
        if sub_lo not in (0x04, 0x06):
            continue
        sub_str = sub_label(0x08, sub_lo)
        args_bytes = f.payload[2:]
        print(f"  t={f.t_rel:>9.4f}  {sub_str:<14}  payload={f.payload.hex()}  args={args_bytes.hex()}")
    return 0


# ────────────────────────────────────────────────────────────────────────────
# dump — pretty frame dump in a time window
# ────────────────────────────────────────────────────────────────────────────

def cmd_dump(args: argparse.Namespace) -> int:
    path = resolve_log(args.path)
    t_start = args.start
    t_end = args.end
    only_ffb = args.only_ffb
    only_h2b = args.only_h2b
    only_b2h = args.only_b2h
    no_heartbeat = args.no_heartbeat

    print(f"Frame dump from {path} (t_rel [{t_start}, {t_end}])\n")
    for f in iter_frames(path):
        if f.t_rel < t_start:
            continue
        if f.t_rel > t_end:
            break
        if only_h2b and not f.is_h2b:
            continue
        if only_b2h and not f.is_b2h:
            continue
        if only_ffb and not f.is_ffb:
            continue
        if no_heartbeat and f.grp in (0x00, 0x80):
            continue

        sub_str = ""
        if f.is_ffb and len(f.payload) >= 2:
            sub_hi, sub_lo = f.payload[0], f.payload[1]
            sub_str = f"  [{sub_label(sub_hi, sub_lo)}]"

        print(f"  {f.t_rel:>9.4f}  {f.dir}  grp=0x{f.grp:02x} dev=0x{f.dev:02x} "
              f"len={len(f.payload):2d}  {f.payload.hex()}{sub_str}")
    return 0


# ────────────────────────────────────────────────────────────────────────────
# h2b-unique — unique h2b frame shapes by (grp, sub-cmd)
# ────────────────────────────────────────────────────────────────────────────

def cmd_h2b_unique(args: argparse.Namespace) -> int:
    """For each (grp, sub-cmd-2-byte-or-marker) shape, show a few sample payloads.

    Useful for catching frame types we missed in the protocol doc.
    """
    path = resolve_log(args.path)
    print(f"Unique h2b frame shapes from {path}\n")

    # Group by (grp, sub_hi, sub_lo) for FFB; just (grp, payload[0..1] if any) otherwise.
    samples: dict[tuple[int, int, int], list[tuple[float, bytes]]] = defaultdict(list)
    counts: Counter = Counter()
    max_samples = args.samples

    for f in iter_frames(path):
        if not f.is_h2b:
            continue
        sub_hi = f.payload[0] if len(f.payload) > 0 else -1
        sub_lo = f.payload[1] if len(f.payload) > 1 else -1
        key = (f.grp, sub_hi, sub_lo)
        counts[key] += 1
        if len(samples[key]) < max_samples:
            samples[key].append((f.t_rel, f.payload))

    print(f"  {'grp':<6} {'sub':<8} {'count':>10}  samples (t_rel, payload)")
    for key, c in counts.most_common():
        grp, sub_hi, sub_lo = key
        sub_repr = f"{sub_hi:02x}/{sub_lo:02x}" if sub_lo != -1 else f"{sub_hi:02x}/--"
        if sub_hi == -1:
            sub_repr = "--"
        print(f"  0x{grp:02x}   {sub_repr:<8} {c:>10,}")
        for t, p in samples[key]:
            print(f"      t={t:>9.4f}  {p.hex()}")
    return 0


# ────────────────────────────────────────────────────────────────────────────
# CLI
# ────────────────────────────────────────────────────────────────────────────

def build_parser() -> argparse.ArgumentParser:
    parser = argparse.ArgumentParser(prog="ab9-decode-session",
                                     description=__doc__,
                                     formatter_class=argparse.RawDescriptionHelpFormatter)
    sub = parser.add_subparsers(dest="cmd", required=True)

    p_inv = sub.add_parser("inventory", help="frame counts, sub-cmd histogram, rates")
    p_inv.add_argument("path", nargs="?", help="JSONL path (default: most recent)")
    p_inv.set_defaults(func=cmd_inventory)

    p_init = sub.add_parser("init", help="early-session frames + first FFB frame timestamp")
    p_init.add_argument("path", nargs="?")
    p_init.add_argument("--seconds", type=float, default=2.0,
                        help="early-window seconds (default 2.0)")
    p_init.add_argument("--include-heartbeats", action="store_true")
    p_init.set_defaults(func=cmd_init)

    p_slots = sub.add_parser("slots", help="slot-ID/period scatter for 0x0A 0x05")
    p_slots.add_argument("path", nargs="?")
    p_slots.add_argument("--top-periods", action="store_true",
                         help="also show top-15 period values")
    p_slots.set_defaults(func=cmd_slots)

    p_pulse = sub.add_parser("pulse", help="0x0B 0x02/03 16-bit field samples")
    p_pulse.add_argument("path", nargs="?")
    p_pulse.add_argument("--dump", type=int, default=8, help="sample count per sub")
    p_pulse.add_argument("--full-scan", action="store_true",
                         help="scan entire file (don't stop at --dump samples)")
    p_pulse.set_defaults(func=cmd_pulse)

    p_t5 = sub.add_parser("trig5", help="0x0D 0x05 rate over time")
    p_t5.add_argument("path", nargs="?")
    p_t5.add_argument("--bucket", type=float, default=10.0,
                      help="bucket size in seconds (default 10)")
    p_t5.set_defaults(func=cmd_trig5)

    p_lr = sub.add_parser("lowrate", help="0x08 0x04/06 frames")
    p_lr.add_argument("path", nargs="?")
    p_lr.set_defaults(func=cmd_lowrate)

    p_dump = sub.add_parser("dump", help="pretty frame dump in a time window")
    p_dump.add_argument("path", nargs="?")
    p_dump.add_argument("--start", type=float, default=0.0)
    p_dump.add_argument("--end", type=float, default=5.0)
    p_dump.add_argument("--only-ffb", action="store_true")
    p_dump.add_argument("--only-h2b", action="store_true")
    p_dump.add_argument("--only-b2h", action="store_true")
    p_dump.add_argument("--no-heartbeat", action="store_true")
    p_dump.set_defaults(func=cmd_dump)

    p_uniq = sub.add_parser("h2b-unique", help="unique h2b frame shapes")
    p_uniq.add_argument("path", nargs="?")
    p_uniq.add_argument("--samples", type=int, default=3,
                        help="samples per shape (default 3)")
    p_uniq.set_defaults(func=cmd_h2b_unique)

    return parser


def main(argv: list[str] | None = None) -> int:
    parser = build_parser()
    args = parser.parse_args(argv)
    return args.func(args)


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