#!/usr/bin/env python3
"""Show detailed traffic timeline around dashboard switches in a bridge capture.

Prints all telemetry-group traffic (session data, value frames, commands, FC)
in a ±N second window around each FF-record kind=4 (DASH_SWITCH) event, plus
any session open/close events and flag transitions.

Usage:
    tools/bridge-switch [CAPTURE] [--window N]
    tools/bridge-switch sim/logs/bridge-20260429-163951.jsonl
    tools/bridge-switch sim/logs/bridge-20260429-163951.jsonl --window 30
"""
import argparse
import struct
import sys
from collections import defaultdict
from pathlib import Path

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


def decode_cmd(payload: bytes) -> str:
    if not payload:
        return "?"
    cmd = payload[0]
    rest = payload[1:].hex() if len(payload) > 1 else ""
    return f"cmd=0x{cmd:02X} {rest}".strip()


def describe_frame(f: BFrame) -> str:
    """One-line description of a telemetry-group frame."""
    prefix = f"{'→' if f.dir == 'h2b' else '←'}"

    if f.is_value_frame:
        flag = f.vf_flag
        data = f.vf_data
        return f"{prefix} VALUE flag=0x{flag:02X} data={len(data)}B {data[:8].hex()}"

    if f.is_session_data:
        sid = f.sess_id
        stype = f.sess_type
        if stype == 0x81:
            port_lo = f.payload[4] if len(f.payload) > 4 else 0
            port_hi = f.payload[5] if len(f.payload) > 5 else 0
            port = port_lo | (port_hi << 8)
            return f"{prefix} SESS-OPEN sess=0x{sid:02X} port={port}"
        elif stype == 0x00:
            return f"{prefix} SESS-CLOSE sess=0x{sid:02X} seq={f.sess_seq}"
        elif stype == 0x01:
            data = f.sess_data
            ff = ""
            if f.ff_kind >= 0:
                ff = f" FF-kind={f.ff_kind} size={f.ff_size}"
            return (f"{prefix} SESS-DATA sess=0x{sid:02X} seq={f.sess_seq} "
                    f"data={len(data)}B{ff} {data[:16].hex()}")
        elif stype == 0x80:
            return f"{prefix} SESS-OPEN-ACK sess=0x{sid:02X}"
        else:
            return f"{prefix} SESS-?? sess=0x{sid:02X} stype=0x{stype:02X}"

    if f.is_flow_control:
        return f"{prefix} FC sess=0x{f.fc_session:02X} seq={f.fc_seq}"

    if f.is_command:
        return f"{prefix} CMD {decode_cmd(f.payload)}"

    return f"{prefix} RAW {f.payload[:16].hex()}"


def main():
    ap = argparse.ArgumentParser(description=__doc__,
                                 formatter_class=argparse.RawDescriptionHelpFormatter)
    ap.add_argument("capture", nargs="?", help="Bridge JSONL file (default: latest)")
    ap.add_argument("--window", "-w", type=float, default=10,
                    help="Seconds before/after each switch to show (default: 10)")
    ap.add_argument("--all-frames", "-a", action="store_true",
                    help="Show all frame types (default: skip high-volume cmd=0x00/0x80)")
    args = ap.parse_args()

    path = resolve_bridge(args.capture)
    frames = load_bridge(path)
    if not frames:
        print("Empty capture.")
        return

    print(f"Capture: {path.name}  ({len(frames)} frames, {frames[-1].t_rel:.1f}s)")
    print()

    # Find all KIND4 events
    switches = [f for f in frames if f.ff_kind == 4]
    if not switches:
        print("No FF-record kind=4 (DASH_SWITCH) found.")
        # Fall back: look for flag transitions
        print("\nLooking for value-frame flag transitions instead...")
        vf = [f for f in frames if f.is_value_frame]
        if not vf:
            print("No value frames either.")
            return
        prev_flag = -1
        for f in vf:
            if f.vf_flag != prev_flag:
                print(f"  t={f.t_rel:8.3f}s  flag → 0x{f.vf_flag:02X}  data={len(f.vf_data)}B")
                prev_flag = f.vf_flag
        return

    # Find all session opens/closes for context
    sess_events = [(f.t_rel, f) for f in frames
                   if f.is_session_data and f.sess_type in (0x81, 0x00, 0x80)]

    print(f"Found {len(switches)} DASH_SWITCH events\n")

    # Also gather value frame flag stats by time bucket
    vf_frames = [f for f in frames if f.is_value_frame]

    for si, sw in enumerate(switches):
        t_switch = sw.t_rel
        t_lo = t_switch - args.window
        t_hi = t_switch + args.window

        print(f"{'='*72}")
        print(f"SWITCH #{si+1} at t={t_switch:.3f}s")
        print(f"Window: [{t_lo:.1f}s .. {t_hi:.1f}s]")
        print(f"{'='*72}")

        # Value frame flag summary in this window
        window_vf = [f for f in vf_frames if t_lo <= f.t_rel <= t_hi]
        pre_vf = [f for f in window_vf if f.t_rel < t_switch]
        post_vf = [f for f in window_vf if f.t_rel >= t_switch]

        if pre_vf:
            pre_flags = defaultdict(lambda: {"count": 0, "sizes": set()})
            for f in pre_vf:
                pre_flags[f.vf_flag]["count"] += 1
                pre_flags[f.vf_flag]["sizes"].add(len(f.vf_data))
            print(f"\n  Pre-switch VFs ({len(pre_vf)}):")
            for flag in sorted(pre_flags):
                info = pre_flags[flag]
                print(f"    flag=0x{flag:02X}: {info['count']}× data_sizes={sorted(info['sizes'])}")

        if post_vf:
            post_flags = defaultdict(lambda: {"count": 0, "sizes": set()})
            for f in post_vf:
                post_flags[f.vf_flag]["count"] += 1
                post_flags[f.vf_flag]["sizes"].add(len(f.vf_data))
            print(f"\n  Post-switch VFs ({len(post_vf)}):")
            for flag in sorted(post_flags):
                info = post_flags[flag]
                print(f"    flag=0x{flag:02X}: {info['count']}× data_sizes={sorted(info['sizes'])}")

        print(f"\n  Timeline:")

        # Gather all interesting frames in the window
        window_frames = []
        for f in frames:
            if f.t_rel < t_lo or f.t_rel > t_hi:
                continue
            if not f.is_telemetry:
                continue
            # Skip high-volume noise unless --all-frames
            if not args.all_frames:
                if f.is_command and f.payload and f.payload[0] in (0x00, 0x80):
                    continue
            window_frames.append(f)

        prev_t = None
        for f in window_frames:
            gap = ""
            if prev_t is not None and (f.t_rel - prev_t) > 0.1:
                gap = f"\n  {'':>12s}  --- gap {f.t_rel - prev_t:.3f}s ---\n"
            marker = " *** SWITCH ***" if f is sw else ""
            if gap:
                print(gap, end="")
            print(f"  t={f.t_rel:10.3f}s  {describe_frame(f)}{marker}")
            prev_t = f.t_rel

        print()


if __name__ == "__main__":
    main()
