#!/usr/bin/env python3
"""Reassemble session-layer data from a MOZA wire trace and decode TLV tier-defs.

Usage:
    tools/trace-sessions [TRACE]                    # latest trace, all sessions
    tools/trace-sessions [TRACE] --session 1        # only session 0x01
    tools/trace-sessions [TRACE] --direction h2b    # only host→wheel
    tools/trace-sessions [TRACE] --raw              # hex dump reassembled bytes
"""
import argparse
import struct
import sys
from collections import defaultdict
from pathlib import Path

sys.path.insert(0, str(Path(__file__).parent))
from moza_trace import load_trace, resolve_trace, Frame

TAG_NAMES = {
    0x00: "ENABLE_PREV",
    0x01: "TIER",
    0x03: "FLAG_BASE",
    0x04: "URL",
    0x06: "END_MARKER",
    0x07: "PROTO_VER",
}

COMP_NAMES = {
    0x00: "none",
    0x01: "speed_1",
    0x02: "rpm_1",
    0x03: "gear_1",
    0x04: "boost_1",
    0x05: "fuel_pct_1",
    0x06: "fuel_l_1",
    0x07: "water_c_1",
    0x08: "oil_c_1",
    0x09: "oil_kpa_1",
    0x0A: "lap_time_ms",
    0x0B: "lap_num_1",
    0x0C: "position_1",
    0x0D: "throttle_pct_1",
    0x0E: "brake_pct_1",
    0x0F: "clutch_pct_1",
    0x10: "steer_deg_1",
    0x11: "tc_level_1",
    0x12: "abs_level_1",
    0x13: "bb_pct_1",
    0x14: "ers_pct_1",
    0x15: "ers_mode_1",
    0x16: "tyre_pressure_1",
    0x17: "best_lap_ms",
    0x18: "last_lap_ms",
    0x19: "delta_ms",
    0x1A: "drs_1",
    0x1B: "pit_limiter_1",
    0x1C: "flag_color_1",
    0x1D: "brake_temp_1",
}


def extract_session_payload(f: Frame) -> tuple[int, int, int, bytes] | None:
    """Extract (session, stype, seq, payload) from a session-layer frame."""
    if f.session < 0 or f.stype != 0x01:
        return None

    if f.dir == 'h2b':
        if len(f.raw) <= 15:
            payload = f.raw[10:]
        else:
            payload = f.raw[10:-5]
    else:
        payload = f.raw[8:]

    return (f.session, f.stype, f.seq, payload)


def reassemble_session(frames: list[Frame], session: int, direction: str) -> list[tuple[float, int, bytes]]:
    """Reassemble data chunks for a session, returning (time, seq, payload) tuples ordered by seq."""
    chunks = []
    for f in frames:
        if f.dir != direction or f.session != session or f.stype != 0x01:
            continue
        if f.dir == 'h2b':
            payload = f.raw[10:-5] if len(f.raw) > 15 else f.raw[10:]
        else:
            payload = f.raw[8:]
        chunks.append((f.t, f.seq, payload))
    chunks.sort(key=lambda x: x[1])
    return chunks


def group_by_seq_runs(chunks: list[tuple[float, int, bytes]]) -> list[list[tuple[float, int, bytes]]]:
    """Group chunks into contiguous seq runs (messages). A gap in seq starts a new group."""
    if not chunks:
        return []
    groups = []
    current = [chunks[0]]
    for i in range(1, len(chunks)):
        prev_seq = chunks[i - 1][1]
        cur_seq = chunks[i][1]
        if cur_seq != prev_seq + 1:
            groups.append(current)
            current = [chunks[i]]
        else:
            current.append(chunks[i])
    groups.append(current)
    return groups


def decode_tlv(data: bytes) -> list[dict]:
    """Parse TLV records from reassembled tier-def bytes."""
    records = []
    pos = 0
    while pos + 5 <= len(data):
        tag = data[pos]
        size = struct.unpack_from('<I', data, pos + 1)[0]
        if pos + 5 + size > len(data):
            records.append({
                'tag': tag, 'tag_name': TAG_NAMES.get(tag, f'UNK_0x{tag:02X}'),
                'size': size, 'error': 'truncated', 'offset': pos,
                'raw': data[pos:].hex()
            })
            break
        value = data[pos + 5:pos + 5 + size]
        rec = {
            'tag': tag, 'tag_name': TAG_NAMES.get(tag, f'UNK_0x{tag:02X}'),
            'size': size, 'offset': pos, 'value_hex': value.hex()
        }

        if tag == 0x07 and size >= 4:
            rec['proto_version'] = struct.unpack_from('<I', value, 0)[0]
            if size >= 8:
                rec['proto_extra'] = struct.unpack_from('<I', value, 4)[0]

        elif tag == 0x03:
            pass

        elif tag == 0x00 and size == 1:
            rec['enable_flag'] = value[0]

        elif tag == 0x01 and size >= 1:
            rec['tier_flag'] = value[0]
            n_channels = (size - 1) // 16
            channels = []
            for ci in range(n_channels):
                off = 1 + ci * 16
                idx = struct.unpack_from('<I', value, off)[0]
                comp = struct.unpack_from('<I', value, off + 4)[0]
                bw = struct.unpack_from('<I', value, off + 8)[0]
                reserved = struct.unpack_from('<I', value, off + 12)[0]
                channels.append({
                    'index': idx,
                    'compression': comp,
                    'comp_name': COMP_NAMES.get(comp, f'0x{comp:02X}'),
                    'bit_width': bw,
                    'reserved': reserved,
                })
            rec['channels'] = channels

        elif tag == 0x06 and size == 4:
            rec['end_marker'] = struct.unpack_from('<I', value, 0)[0]

        records.append(rec)
        pos += 5 + size
    return records


def print_tlv_records(records: list[dict], indent: str = "  "):
    for rec in records:
        tag_name = rec['tag_name']
        offset = rec['offset']
        size = rec['size']

        if 'error' in rec:
            print(f"{indent}[offset={offset}] {tag_name} size={size} ** {rec['error']} ** raw={rec['raw'][:40]}...")
            continue

        if rec['tag'] == 0x07:
            ver = rec.get('proto_version', '?')
            extra = f" extra={rec['proto_extra']}" if 'proto_extra' in rec else ""
            print(f"{indent}[{offset:3d}] PROTO_VER  version={ver}{extra}")

        elif rec['tag'] == 0x03:
            print(f"{indent}[{offset:3d}] FLAG_BASE  (size=0)")

        elif rec['tag'] == 0x00:
            print(f"{indent}[{offset:3d}] ENABLE     flag=0x{rec['enable_flag']:02X}")

        elif rec['tag'] == 0x01:
            flag = rec['tier_flag']
            channels = rec.get('channels', [])
            print(f"{indent}[{offset:3d}] TIER       flag=0x{flag:02X}  channels={len(channels)}")
            for ch in channels:
                print(f"{indent}  ch: idx={ch['index']:3d}  comp={ch['comp_name']:<20s} bw={ch['bit_width']:2d}  reserved={ch['reserved']}")

        elif rec['tag'] == 0x06:
            print(f"{indent}[{offset:3d}] END_MARKER value={rec['end_marker']}")

        else:
            print(f"{indent}[{offset:3d}] {tag_name}  size={size}  hex={rec['value_hex'][:40]}")


def main():
    parser = argparse.ArgumentParser(description="Reassemble session data and decode tier-def TLVs")
    parser.add_argument("trace", nargs="?", default="latest")
    parser.add_argument("--session", "-s", type=lambda x: int(x, 0), default=None,
                        help="Filter to specific session (hex ok: 0x01)")
    parser.add_argument("--direction", "-d", choices=["h2b", "b2h"], default=None)
    parser.add_argument("--raw", action="store_true", help="Show raw hex of reassembled data")
    parser.add_argument("--no-decode", action="store_true", help="Skip TLV decode")
    args = parser.parse_args()

    path = resolve_trace(args.trace)
    frames = load_trace(path)
    print(f"Trace: {path.name}  ({len(frames)} frames)")

    sessions_seen: dict[tuple[str, int], list[Frame]] = defaultdict(list)
    for f in frames:
        if f.is_session:
            key = (f.dir, f.session)
            sessions_seen[key].append(f)

    directions = [args.direction] if args.direction else ['h2b', 'b2h']
    session_ids = [args.session] if args.session is not None else sorted(set(s for _, s in sessions_seen.keys()))

    for direction in directions:
        for sid in session_ids:
            key = (direction, sid)
            if key not in sessions_seen:
                continue

            sframes = sessions_seen[key]
            opens = [f for f in sframes if f.is_open]
            closes = [f for f in sframes if f.is_close]
            datas = [f for f in sframes if f.is_data]

            label = f"{'host→wheel' if direction == 'h2b' else 'wheel→host'} sess=0x{sid:02X}"
            print(f"\n{'='*60}")
            print(f"{label}: {len(opens)} opens, {len(closes)} closes, {len(datas)} data chunks")

            for f in opens:
                print(f"  OPEN  t={f.t:.3f}s port={f.port}")
            for f in closes:
                print(f"  CLOSE t={f.t:.3f}s")

            if not datas:
                continue

            chunks = reassemble_session(frames, sid, direction)
            groups = group_by_seq_runs(chunks)

            print(f"  {len(groups)} message group(s) from {len(chunks)} data chunks")

            for gi, group in enumerate(groups):
                first_t = group[0][0]
                first_seq = group[0][1]
                last_seq = group[-1][1]
                total_payload = sum(len(p) for _, _, p in group)
                assembled = b''.join(p for _, _, p in group)

                print(f"\n  --- Message {gi+1}: t={first_t:.3f}s  seq={first_seq}..{last_seq}  "
                      f"chunks={len(group)}  payload={total_payload}B ---")

                if args.raw:
                    hex_str = assembled.hex()
                    for i in range(0, len(hex_str), 80):
                        print(f"    {hex_str[i:i+80]}")

                if not args.no_decode and direction == 'h2b' and sid in (0x01, 1):
                    records = decode_tlv(assembled)
                    if records:
                        print_tlv_records(records, "    ")
                    else:
                        print("    (no TLV records decoded)")


if __name__ == "__main__":
    main()
