#!/usr/bin/env python3
"""Reusable wire trace analysis toolkit for Moza protocol captures.

Usage:
  python3 /tmp/trace-tools.py <command> <trace.jsonl> [options]

Commands:
  summary     - Overview of traffic (h2b/b2h counts, sessions, frame types)
  sessions    - Session-layer breakdown (per-session chunk counts, types)
  catalog     - Extract wheel catalog from b2h tag=0x04 records
  value       - Analyze value frames (flag bytes, emission rate)
  timeline    - Show key protocol events in chronological order
  hexdump     - Dump raw frames matching filters
"""
import json, sys, os
from collections import Counter, defaultdict

def load_trace(path):
    records = []
    with open(path) as f:
        for line in f:
            try:
                rec = json.loads(line)
                rec['raw'] = bytes.fromhex(rec.get('hex', ''))
                rec['ts'] = rec.get('t', 0)
                records.append(rec)
            except:
                continue
    return records

def parse_h2b_frame(raw):
    """Parse host→wheel frame: 7E [N] [grp] [dev] [payload...] [chk]"""
    if len(raw) < 4 or raw[0] != 0x7E:
        return None
    return {
        'n': raw[1],
        'grp': raw[2],
        'dev': raw[3],
        'body': raw[4:-1] if len(raw) > 5 else b'',
        'full': raw,
    }

def parse_b2h_frame(raw):
    """Parse wheel→host frame: C3 71 [sub...] or other groups.
    Session frames: C3 71 7C 00 [session] [type] [seq_lo] [seq_hi] [payload...]
    FC ack frames:  C3 71 FC 00 [session] [ack_lo] [ack_hi]
    """
    if len(raw) < 2:
        return None
    info = {'grp': raw[0], 'dev': raw[1], 'full': raw}
    if raw[0] == 0xC3 and raw[1] == 0x71:
        if len(raw) >= 8 and raw[2] == 0x7C and raw[3] == 0x00:
            info['session'] = raw[4]
            info['type'] = raw[5]
            info['seq'] = raw[6] | (raw[7] << 8)
            info['payload'] = raw[8:]
            info['is_session'] = True
        elif len(raw) >= 7 and raw[2] == 0xFC and raw[3] == 0x00:
            info['session'] = raw[4]
            info['ack_seq'] = raw[5] | (raw[6] << 8)
            info['is_ack'] = True
        else:
            info['subcmd'] = raw[2:4].hex() if len(raw) >= 4 else ''
    return info

def cmd_summary(records, args):
    h2b_count = sum(1 for r in records if r.get('dir') == 'h2b')
    b2h_count = sum(1 for r in records if r.get('dir') == 'b2h')

    # h2b breakdown
    h2b_groups = Counter()
    h2b_7d23 = 0
    h2b_sess_chunks = Counter()
    h2b_enable = 0
    h2b_seqctr = 0
    flag_bytes = set()

    for r in records:
        if r.get('dir') != 'h2b':
            continue
        raw = r['raw']
        f = parse_h2b_frame(raw)
        if not f:
            continue
        h2b_groups[f"0x{f['grp']:02x}:0x{f['dev']:02x}"] += 1

        if f['grp'] == 0x43 and f['dev'] == 0x17 and len(raw) >= 8:
            if raw[4] == 0x7D and raw[5] == 0x23:
                h2b_7d23 += 1
                if len(raw) > 10:
                    flag_bytes.add(raw[10])
            elif raw[4] == 0x7C and raw[5] == 0x00:
                sess = raw[6]
                typ = raw[7]
                h2b_sess_chunks[f"sess=0x{sess:02x} type=0x{typ:02x}"] += 1
        if f['grp'] == 0x41 and len(raw) >= 6 and raw[4] == 0xFD and raw[5] == 0xDE:
            h2b_enable += 1
        if f['grp'] == 0x2D and len(raw) >= 6 and raw[4] == 0xF5 and raw[5] == 0x31:
            h2b_seqctr += 1

    # b2h breakdown
    b2h_sessions = Counter()
    b2h_session_types = Counter()
    b2h_acks = 0
    b2h_other = Counter()

    for r in records:
        if r.get('dir') != 'b2h':
            continue
        raw = r['raw']
        f = parse_b2h_frame(raw)
        if not f:
            continue
        if f.get('is_session'):
            b2h_sessions[f"sess=0x{f['session']:02x}"] += 1
            b2h_session_types[f"sess=0x{f['session']:02x} type=0x{f['type']:02x}"] += 1
        elif f.get('is_ack'):
            b2h_acks += 1
        else:
            key = f"0x{f['grp']:02x}:0x{f['dev']:02x}"
            if 'subcmd' in f:
                key += f" {f['subcmd']}"
            b2h_other[key] += 1

    print(f"=== TRACE SUMMARY ===")
    print(f"Total frames: {len(records)} (h2b={h2b_count}, b2h={b2h_count})")
    print()

    print(f"--- h2b (host → wheel) ---")
    print(f"  Value frames (7d23): {h2b_7d23}")
    print(f"  Flag bytes: {sorted(flag_bytes)}")
    print(f"  FFB enable: {h2b_enable}")
    print(f"  Sequence counter: {h2b_seqctr}")
    print(f"  Session chunks:")
    for k, v in sorted(h2b_sess_chunks.items()):
        print(f"    {k}: {v}")
    print(f"  Group breakdown:")
    for k, v in h2b_groups.most_common(15):
        print(f"    {k}: {v}")
    print()

    print(f"--- b2h (wheel → host) ---")
    print(f"  Session data/open/close (C3 71 7C 00):")
    for k, v in sorted(b2h_sessions.items()):
        print(f"    {k}: {v}")
    for k, v in sorted(b2h_session_types.items()):
        print(f"      {k}: {v}")
    print(f"  FC acks: {b2h_acks}")
    print(f"  Other groups:")
    for k, v in b2h_other.most_common(15):
        print(f"    {k}: {v}")

def cmd_sessions(records, args):
    print("=== SESSION BREAKDOWN ===")
    # h2b sessions
    h2b = defaultdict(lambda: Counter())
    for r in records:
        if r.get('dir') != 'h2b':
            continue
        raw = r['raw']
        if len(raw) >= 8 and raw[0] == 0x7E and raw[2] == 0x43 and raw[4] == 0x7C and raw[5] == 0x00:
            sess = raw[6]
            typ = raw[7]
            h2b[sess][typ] += 1

    print("h2b session chunks:")
    for sess in sorted(h2b):
        total = sum(h2b[sess].values())
        parts = ", ".join(f"type=0x{t:02x}: {c}" for t, c in sorted(h2b[sess].items()))
        print(f"  sess=0x{sess:02x}: {total} ({parts})")

    # b2h sessions
    b2h = defaultdict(lambda: Counter())
    for r in records:
        if r.get('dir') != 'b2h':
            continue
        raw = r['raw']
        f = parse_b2h_frame(raw)
        if f and f.get('is_session'):
            b2h[f['session']][f['type']] += 1

    print("\nb2h session chunks:")
    for sess in sorted(b2h):
        total = sum(b2h[sess].values())
        parts = ", ".join(f"type=0x{t:02x}: {c}" for t, c in sorted(b2h[sess].items()))
        print(f"  sess=0x{sess:02x}: {total} ({parts})")

    # b2h acks per session
    b2h_acks = Counter()
    for r in records:
        if r.get('dir') != 'b2h':
            continue
        raw = r['raw']
        f = parse_b2h_frame(raw)
        if f and f.get('is_ack'):
            b2h_acks[f['session']] += 1

    print("\nb2h FC acks per session:")
    for sess, count in sorted(b2h_acks.items()):
        print(f"  sess=0x{sess:02x}: {count}")

def cmd_catalog(records, args):
    """Extract catalog from b2h tag=0x04 records."""
    print("=== CATALOG EXTRACTION ===")
    # Catalog arrives on b2h session chunks. Scan all b2h session data for tag=0x04.
    catalog = {}
    for r in records:
        if r.get('dir') != 'b2h':
            continue
        raw = r['raw']
        f = parse_b2h_frame(raw)
        if not f or not f.get('is_session') or f['type'] != 0x01:
            continue
        payload = f['payload']
        # Scan for tag=0x04 records: 04 [len_lo] [len_hi?] [00 00] [idx_lo] [idx_hi?] ...
        for i in range(len(payload)):
            if payload[i] == 0x04 and i + 5 < len(payload):
                tag_len = payload[i+1]
                idx = payload[i+5] if i+5 < len(payload) else -1
                # Try to extract URL string after the index bytes
                if i + 7 < len(payload):
                    url_start = i + 7  # skip tag(1) + len(1) + 00 00(2) + idx(1) + 00(1) + type(1)
                    # Look for ASCII URL
                    url_bytes = []
                    for j in range(i+6, min(i+1+tag_len, len(payload))):
                        if 0x20 <= payload[j] <= 0x7E:
                            url_bytes.append(payload[j])
                        elif url_bytes:
                            break
                    if url_bytes:
                        url = bytes(url_bytes).decode('ascii', errors='replace')
                        catalog[idx] = url

    if catalog:
        print(f"Found {len(catalog)} catalog entries:")
        for idx in sorted(catalog):
            print(f"  [{idx}] {catalog[idx]}")
    else:
        print("No catalog entries found in b2h session data.")

def cmd_value(records, args):
    """Analyze value frame emission patterns."""
    print("=== VALUE FRAME ANALYSIS ===")
    flag_times = defaultdict(list)
    total = 0

    for r in records:
        if r.get('dir') != 'h2b':
            continue
        raw = r['raw']
        if len(raw) >= 12 and raw[0] == 0x7E and raw[2] == 0x43 and raw[3] == 0x17:
            if raw[4] == 0x7D and raw[5] == 0x23:
                total += 1
                flag = raw[10]
                ts = r.get('ts', 0)
                flag_times[flag].append(ts)

    print(f"Total value frames: {total}")
    print(f"Unique flag bytes: {len(flag_times)}")
    print()

    for flag in sorted(flag_times):
        times = flag_times[flag]
        count = len(times)
        if count >= 2:
            duration = times[-1] - times[0]
            rate = count / duration if duration > 0 else 0
            print(f"  flag=0x{flag:02x} ({flag:3d}): {count} frames, first={times[0]:.3f}s last={times[-1]:.3f}s rate={rate:.1f}/s")
        else:
            print(f"  flag=0x{flag:02x} ({flag:3d}): {count} frame(s)")

def cmd_timeline(records, args):
    """Show key protocol events chronologically."""
    print("=== PROTOCOL TIMELINE ===")
    t0 = records[0].get('ts', 0) if records else 0

    for r in records:
        raw = r['raw']
        ts = r.get('ts', 0) - t0
        d = r.get('dir', '?')

        event = None
        if d == 'h2b':
            f = parse_h2b_frame(raw)
            if not f:
                continue
            # Session open (type=0x81)
            if f['grp'] == 0x43 and len(raw) >= 8 and raw[4] == 0x7C and raw[5] == 0x00 and raw[7] == 0x81:
                sess = raw[6]
                event = f"h2b OPEN sess=0x{sess:02x}"
            # Session close (type=0x00)
            elif f['grp'] == 0x43 and len(raw) >= 8 and raw[4] == 0x7C and raw[5] == 0x00 and raw[7] == 0x00:
                sess = raw[6]
                event = f"h2b CLOSE sess=0x{sess:02x}"
            # FF records on sess02 (kind detection)
            elif f['grp'] == 0x43 and len(raw) >= 10 and raw[4] == 0x7C and raw[5] == 0x00 and raw[6] == 0x02 and raw[7] == 0x01:
                payload = raw[8:-1]
                if len(payload) >= 5 and payload[0] == 0xFF:
                    kind = payload[4]
                    event = f"h2b FF kind={kind} on sess=02 ({len(payload)}B)"
            # First value frame
            elif f['grp'] == 0x43 and len(raw) >= 12 and raw[4] == 0x7D and raw[5] == 0x23:
                flag = raw[10]
                event = f"h2b VALUE flag=0x{flag:02x} ({len(raw)}B)"
        elif d == 'b2h':
            bf = parse_b2h_frame(raw)
            if not bf:
                continue
            if bf.get('is_session'):
                if bf['type'] == 0x81:
                    event = f"b2h OPEN sess=0x{bf['session']:02x} seq={bf['seq']}"
                elif bf['type'] == 0x00:
                    event = f"b2h CLOSE sess=0x{bf['session']:02x}"
                elif bf['type'] == 0x01 and bf['session'] in (0x09, 0x0A):
                    event = f"b2h DATA sess=0x{bf['session']:02x} seq={bf['seq']} ({len(bf.get('payload', b''))}B)"
            elif bf.get('is_ack'):
                if bf['session'] in (0x01, 0x09):
                    event = f"b2h ACK sess=0x{bf['session']:02x} ack_seq={bf['ack_seq']}"

        if event:
            # Deduplicate consecutive value frames
            if 'VALUE' in (event or ''):
                continue  # skip value frames in timeline (too many)
            print(f"  {ts:8.3f}s  {event}")

def cmd_hexdump(records, args):
    """Dump frames matching filter. Usage: hexdump <trace> [dir=h2b|b2h] [grp=0xNN] [sess=0xNN] [limit=N]"""
    filt_dir = None
    filt_grp = None
    filt_sess = None
    limit = 20
    for a in args:
        if a.startswith('dir='):
            filt_dir = a.split('=')[1]
        elif a.startswith('grp='):
            filt_grp = int(a.split('=')[1], 0)
        elif a.startswith('sess='):
            filt_sess = int(a.split('=')[1], 0)
        elif a.startswith('limit='):
            limit = int(a.split('=')[1])

    count = 0
    for r in records:
        if filt_dir and r.get('dir') != filt_dir:
            continue
        raw = r['raw']
        if filt_grp is not None:
            if r.get('dir') == 'h2b' and (len(raw) < 3 or raw[2] != filt_grp):
                continue
            if r.get('dir') == 'b2h' and (len(raw) < 1 or raw[0] != filt_grp):
                continue
        if filt_sess is not None:
            if r.get('dir') == 'h2b':
                if len(raw) < 7 or raw[6] != filt_sess:
                    continue
            elif r.get('dir') == 'b2h':
                f = parse_b2h_frame(raw)
                if not f or f.get('session') != filt_sess:
                    continue

        ts = r.get('ts', 0)
        d = r.get('dir', '?')
        hexstr = raw.hex()
        print(f"  {ts:.3f}s [{d}] {hexstr}")
        count += 1
        if count >= limit:
            print(f"  ... (limit={limit} reached)")
            break

COMMANDS = {
    'summary': cmd_summary,
    'sessions': cmd_sessions,
    'catalog': cmd_catalog,
    'value': cmd_value,
    'timeline': cmd_timeline,
    'hexdump': cmd_hexdump,
}

def main():
    if len(sys.argv) < 3:
        print(__doc__)
        print("Commands:", ", ".join(COMMANDS))
        sys.exit(1)

    cmd = sys.argv[1]
    trace_path = sys.argv[2]
    extra_args = sys.argv[3:]

    if cmd not in COMMANDS:
        print(f"Unknown command: {cmd}. Available: {', '.join(COMMANDS)}")
        sys.exit(1)

    records = load_trace(trace_path)
    print(f"Loaded {len(records)} records from {os.path.basename(trace_path)}")
    print()
    COMMANDS[cmd](records, extra_args)

if __name__ == '__main__':
    main()
