#!/usr/bin/env python3
"""
ts_diag.py -- HLS/MPEG-TS segment diagnostic tool.

Run this ON YOUR SERVER (where your key/segments already live) and share
the TEXT REPORT it prints -- never the key or the raw/decrypted segments
themselves. This tool decrypts locally, analyzes locally, and only
produces a plain-text summary of findings.

What it checks, and why each one matters for LG/Samsung-specific
playback failures that don't show up in ExoPlayer/VLC:

  1. Ciphertext block alignment (AES-CBC requires an exact multiple of
     16 bytes -- a malformed length here is a real corruption signal,
     independent of whether the key is right).
  2. TS sync-byte integrity + continuity-counter errors per PID.
  3. PAT/PMT parsing -- confirms video/audio PID discovery and the
     declared stream_types.
  4. Leading NAL type at the very start of the segment's video stream:
     IDR (type 5) vs non-IDR (type 1). A segment that doesn't open on
     a real IDR is a mid-GOP cut -- strict decoders (native LG/Samsung
     HLS stacks) can refuse a fragment like that outright, while
     ExoPlayer/VLC often just wait/skip ahead silently.
  5. Full NAL type histogram for the segment.
  6. First/last video and audio PTS/DTS, and the audio-video start-time
     gap -- broken/anchored-wrong timestamps are a common root cause
     already confirmed elsewhere in this pipeline.
  7. PCR presence/cadence.
  8. EXTINF vs measured duration, if you pass the manifest.

Requirements: python3 stdlib only, plus the `openssl` CLI on PATH
(same tool used everywhere else in this debugging process).

USAGE
-----
Basic (positional list of segments, in the order you want them
reported -- ideally the same consecutive run you're seeing fail):

    python3 ts_diag.py --key enc.key \
        --iv 2aa96af110f99fd7cfccddee74f52e86 \
        --segments index7957.ts index7958.ts index7959.ts \
                   index7960.ts index7961.ts \
        --manifest index.m3u8 \
        > ts_diag_report.txt

If you don't have a raw hex IV handy but do have the .keyinfo file
(3rd line is the IV), pass that instead:

    python3 ts_diag.py --key enc.key --keyinfo file.keyinfo \
        --segments index7957.ts index7958.ts ... \
        > ts_diag_report.txt

Share the resulting ts_diag_report.txt (plain text, no key material,
no raw stream bytes) back for review.
"""

import argparse
import re
import struct
import subprocess
import sys
import tempfile
import os


# --------------------------------------------------------------------------
# Decryption
# --------------------------------------------------------------------------

def read_iv(args):
    if args.iv:
        return args.iv.strip().lower().lstrip("0x")
    if args.keyinfo:
        with open(args.keyinfo, "r") as f:
            lines = [l.strip() for l in f if l.strip()]
        if len(lines) >= 3:
            return lines[2].lower().lstrip("0x")
        raise SystemExit(
            f"--keyinfo {args.keyinfo} doesn't have a 3rd (IV) line; "
            f"pass --iv explicitly instead.")
    return None  # no fixed IV given -- rely on sequence-derived fallback only


def seq_num_from_filename(path):
    """Extract the trailing integer from a segment filename, e.g.
    'index12496.ts' -> 12496. Returns None if no digits found."""
    base = os.path.basename(path)
    m = re.search(r'(\d+)(?=\.\w+$)', base)
    return int(m.group(1)) if m else None


def seq_derived_iv_hex(seq_num):
    """Per HLS spec (RFC 8216 4.3.2.4): when EXT-X-KEY has no explicit
    IV, the IV is the segment's Media Sequence Number, treated as a
    128-bit big-endian integer (i.e. 16 bytes, sequence number in the
    low-order bytes, zero-padded on the left)."""
    return f"{seq_num:032x}"


def find_key_line_in_manifest(manifest_path):
    if not manifest_path:
        return None
    with open(manifest_path, "r", errors="ignore") as f:
        for line in f:
            if line.strip().startswith("#EXT-X-KEY"):
                return line.strip()
    return None


def read_key_hex(key_path):
    with open(key_path, "rb") as f:
        raw = f.read()
    if len(raw) == 16:
        return raw.hex()
    # maybe it's already a hex-text file
    text = raw.decode("ascii", "ignore").strip()
    if re.fullmatch(r"[0-9a-fA-F]{32}", text):
        return text.lower()
    raise SystemExit(
        f"{key_path} is {len(raw)} bytes -- expected a 16-byte raw AES "
        f"key file (or a 32-char hex text file).")


def try_decrypt_with_iv(path, key_hex, iv_hex):
    """Try normal (padded) decrypt first; if that fails, try truncating
    to the nearest lower multiple of 16 and decrypting with -nopad.
    Returns (data_bytes_or_None, mode_str)."""
    size = os.path.getsize(path)
    with tempfile.NamedTemporaryFile(delete=False) as tmp_out:
        out_path = tmp_out.name
    try:
        r = subprocess.run(
            ["openssl", "enc", "-d", "-aes-128-cbc",
             "-K", key_hex, "-iv", iv_hex,
             "-in", path, "-out", out_path],
            capture_output=True, text=True)
        if r.returncode == 0:
            with open(out_path, "rb") as f:
                data = f.read()
            return data, "padded-ok"
        trunc = (size // 16) * 16
        with open(path, "rb") as f:
            raw = f.read(trunc)
        with tempfile.NamedTemporaryFile(delete=False) as tmp_in:
            tmp_in.write(raw)
            in_path = tmp_in.name
        try:
            r2 = subprocess.run(
                ["openssl", "enc", "-d", "-aes-128-cbc", "-nopad",
                 "-K", key_hex, "-iv", iv_hex,
                 "-in", in_path, "-out", out_path],
                capture_output=True, text=True)
            if r2.returncode == 0:
                with open(out_path, "rb") as f:
                    data = f.read()
                return data, f"nopad-truncated(dropped {size-trunc}B)"
            return None, f"DECRYPT FAILED: {r2.stderr.strip()[-200:]}"
        finally:
            os.unlink(in_path)
    finally:
        os.unlink(out_path)


def decrypt_segment(path, key_hex, iv_hex, seq_num):
    """Tries, in order: the given fixed IV (if any), then the HLS-spec
    sequence-number-derived IV (if a sequence number could be pulled
    from the filename). Returns (data, mode, orig_size, mod16,
    iv_used_label)."""
    orig_size = os.path.getsize(path)
    mod16 = orig_size % 16
    candidates = []
    if iv_hex:
        candidates.append(("given IV", iv_hex))
    if seq_num is not None:
        candidates.append(("sequence-derived IV", seq_derived_iv_hex(seq_num)))
    if not candidates:
        return None, "NO IV AVAILABLE (pass --iv/--keyinfo, or ensure " \
                      "filename has a trailing sequence number)", orig_size, mod16, None

    last_mode = None
    for label, cand_iv in candidates:
        data, mode = try_decrypt_with_iv(path, key_hex, cand_iv)
        last_mode = mode
        if data is not None and data[:1] == b'\x47':
            return data, mode, orig_size, mod16, f"{label} ({cand_iv})"
        # keep the bytes around in case NONE of the candidates produce
        # a valid sync byte, so we can still report what we got with
        # the primary candidate
        if label == candidates[0][0]:
            first_attempt = (data, mode)
    # nothing produced a valid TS sync byte -- return the first
    # candidate's result so the report shows what was actually tried
    data, mode = first_attempt
    return data, f"{mode} [tried {len(candidates)} IV candidate(s), none produced a valid TS sync byte]", \
           orig_size, mod16, None


# --------------------------------------------------------------------------
# MPEG-TS / PSI parsing
# --------------------------------------------------------------------------

TS_SZ = 188


def iter_ts_packets(data):
    n = len(data) // TS_SZ
    for i in range(n):
        yield data[i*TS_SZ:(i+1)*TS_SZ]


def parse_pat(pkt):
    payload = pkt[4:]
    pointer = payload[0]
    sec = payload[1+pointer:]
    if len(sec) < 3:
        return None
    sec_len = ((sec[1] & 0x0f) << 8) | sec[2]
    section = sec[:3+sec_len]
    prog_data = section[8:-4]
    pmt_pid = None
    i = 0
    while i + 3 < len(prog_data):
        prog_num = (prog_data[i] << 8) | prog_data[i+1]
        pid = ((prog_data[i+2] & 0x1f) << 8) | prog_data[i+3]
        if prog_num != 0:
            pmt_pid = pid
        i += 4
    return pmt_pid


def parse_pmt(pkt):
    payload = pkt[4:]
    pointer = payload[0]
    sec = payload[1+pointer:]
    if len(sec) < 3:
        return None, None, []
    sec_len = ((sec[1] & 0x0f) << 8) | sec[2]
    section = sec[:3+sec_len]
    if len(section) < 12:
        return None, None, []
    prog_info_len = ((section[10] & 0x0f) << 8) | section[11]
    i = 12 + prog_info_len
    video_pid = None
    streams = []
    while i < len(section) - 4:
        stream_type = section[i]
        pid = ((section[i+1] & 0x1f) << 8) | section[i+2]
        es_info_len = ((section[i+3] & 0x0f) << 8) | section[i+4]
        streams.append((stream_type, pid))
        i += 5 + es_info_len
    return streams


def find_pat_pmt(pkts):
    pat_pkt = None
    for p in pkts:
        pid = ((p[1] & 0x1f) << 8) | p[2]
        if pid == 0:
            pat_pkt = p
            break
    if not pat_pkt:
        return None, []
    pmt_pid = parse_pat(pat_pkt)
    if pmt_pid is None:
        return None, []
    pmt_pkt = None
    for p in pkts:
        pid = ((p[1] & 0x1f) << 8) | p[2]
        if pid == pmt_pid:
            pmt_pkt = p
            break
    if not pmt_pkt:
        return pmt_pid, []
    streams = parse_pmt(pmt_pkt)
    return pmt_pid, streams


def cc_and_sync_check(pkts):
    cc = {}
    cc_errors = 0
    sync_errors = 0
    tei_errors = 0
    pcr_count = 0
    for p in pkts:
        if p[0] != 0x47:
            sync_errors += 1
            continue
        tei = (p[1] >> 7) & 1
        if tei:
            tei_errors += 1
        pid = ((p[1] & 0x1f) << 8) | p[2]
        afc = (p[3] >> 4) & 0x3
        cur_cc = p[3] & 0xf
        has_payload = afc in (1, 3)
        has_adapt = afc in (2, 3)
        if has_adapt and p[4] > 0:
            adapt_flags = p[5]
            if adapt_flags & 0x10:
                pcr_count += 1
        if pid == 0x1fff:
            continue
        if pid in cc and has_payload:
            expected = (cc[pid] + 1) & 0xf
            if cur_cc != expected:
                cc_errors += 1
        cc[pid] = cur_cc
    return sync_errors, tei_errors, cc_errors, pcr_count


def extract_es(pkts, pid):
    es = bytearray()
    for p in pkts:
        p_pid = ((p[1] & 0x1f) << 8) | p[2]
        if p_pid != pid:
            continue
        afc = (p[3] >> 4) & 0x3
        off = 4
        if afc in (2, 3):
            adapt_len = p[4]
            off = 5 + adapt_len
        if afc in (1, 3):
            es += p[off:TS_SZ]
    return bytes(es)


def find_nal_units(es):
    """Returns list of (nal_type, offset) for every NAL start code."""
    out = []
    for m in re.finditer(b'\x00\x00\x01', es):
        s = m.start()
        if s + 3 >= len(es):
            continue
        nal_hdr = es[s+3]
        nal_type = nal_hdr & 0x1f
        if nal_type == 0:
            continue  # artifact of 4-byte start codes, not a real NAL
        out.append((nal_type, s))
    return out


PES_STREAM_IDS_VIDEO = set(range(0xE0, 0xF0))
PES_STREAM_IDS_AUDIO = set(range(0xC0, 0xE0))


def find_pes_pts_dts(es, max_scan=40):
    """Scan the first few PES headers in this elementary stream and
    return a list of (pts, dts_or_None) tuples, 90kHz ticks."""
    out = []
    i = 0
    n = len(es)
    scanned = 0
    while i < n - 9 and scanned < max_scan:
        if es[i] == 0 and es[i+1] == 0 and es[i+2] == 1:
            stream_id = es[i+3]
            if (stream_id in PES_STREAM_IDS_VIDEO or
                    stream_id in PES_STREAM_IDS_AUDIO):
                if i + 9 > n:
                    break
                pts_dts_flags = (es[i+7] >> 6) & 0x3
                hdr_len = es[i+8]
                pts = dts = None
                if pts_dts_flags in (2, 3) and i + 9 + 5 <= n:
                    b = es[i+9:i+14]
                    pts = (((b[0] >> 1) & 0x07) << 30) | \
                          (b[1] << 22) | ((b[2] >> 1) << 15) | \
                          (b[3] << 7) | (b[4] >> 1)
                if pts_dts_flags == 3 and i + 14 + 5 <= n:
                    b = es[i+14:i+19]
                    dts = (((b[0] >> 1) & 0x07) << 30) | \
                          (b[1] << 22) | ((b[2] >> 1) << 15) | \
                          (b[3] << 7) | (b[4] >> 1)
                if pts is not None:
                    out.append((pts, dts))
                    scanned += 1
                i += 9 + hdr_len
                continue
        i += 1
    return out


def ebsp_to_rbsp(b):
    out = bytearray()
    i = 0
    zero_count = 0
    while i < len(b):
        if zero_count >= 2 and b[i] == 0x03:
            zero_count = 0
            i += 1
            continue
        out.append(b[i])
        zero_count = zero_count + 1 if b[i] == 0 else 0
        i += 1
    return bytes(out)


class BitReader:
    def __init__(self, data):
        self.data = data
        self.pos = 0

    def u(self, n):
        v = 0
        for _ in range(n):
            byte_i = self.pos // 8
            if byte_i >= len(self.data):
                raise IndexError("ran off the end of SPS/PPS RBSP")
            byte = self.data[byte_i]
            bit = (byte >> (7 - (self.pos % 8))) & 1
            v = (v << 1) | bit
            self.pos += 1
        return v

    def ue(self):
        lz = 0
        while self.u(1) == 0:
            lz += 1
            if lz > 32:
                break
        return (1 << lz) - 1 + self.u(lz)


# H.264 Annex A Table A-1: level_idc*10 -> (MaxMBPS, MaxFS, MaxDpbMbs)
LEVEL_LIMITS = {
    10: (1485, 99, 396), 11: (3000, 396, 900), 12: (6000, 396, 2376),
    13: (11880, 396, 2376), 20: (11880, 396, 2376), 21: (19800, 792, 4752),
    22: (20250, 1620, 8100), 30: (40500, 1620, 8100), 31: (108000, 3600, 18000),
    32: (216000, 5120, 20480), 40: (245760, 8192, 32768),
    41: (245760, 8192, 32768), 42: (522240, 8704, 34816),
    50: (589824, 22080, 110400), 51: (983040, 36864, 184320),
}

PROFILE_NAMES = {66: "Baseline", 77: "Main", 88: "Extended", 100: "High",
                  110: "High 10", 122: "High 4:2:2", 244: "High 4:4:4"}

HIGH_FAMILY = {100, 110, 122, 244, 44, 83, 86, 118, 128, 138, 139, 134, 135}


def parse_sps(sps_nal):
    """sps_nal includes the 1-byte NAL header. Returns a dict of
    findings, or {'error': ...} if parsing failed partway."""
    out = {}
    try:
        rbsp = ebsp_to_rbsp(sps_nal[1:])
        br = BitReader(rbsp)
        profile_idc = br.u(8)
        constraint = br.u(8)
        level_idc = br.u(8)
        br.ue()  # sps_id
        out["profile_idc"] = profile_idc
        out["profile_name"] = PROFILE_NAMES.get(profile_idc, f"unknown({profile_idc})")
        out["level_idc"] = level_idc
        out["level"] = f"{level_idc/10:.1f}"
        out["constraint_flags"] = f"0x{constraint:02x}"

        chroma_format_idc = 1
        if profile_idc in HIGH_FAMILY:
            chroma_format_idc = br.ue()
            if chroma_format_idc == 3:
                br.u(1)
            br.ue()  # bit_depth_luma_minus8
            br.ue()  # bit_depth_chroma_minus8
            br.u(1)  # qpprime_y_zero_transform_bypass_flag
            seq_scaling = br.u(1)
            out["seq_scaling_matrix_present"] = bool(seq_scaling)
            if seq_scaling:
                out["note"] = ("SPS has scaling matrices -- not parsed "
                                "further, remaining fields unavailable")
                return out
        out["chroma_format_idc"] = chroma_format_idc

        br.ue()  # log2_max_frame_num_minus4
        poc_type = br.ue()
        if poc_type == 0:
            br.ue()
        elif poc_type == 1:
            out["note"] = ("pic_order_cnt_type==1 (rare) -- not parsed "
                            "further past this point, remaining fields "
                            "unavailable for this SPS")
            return out

        max_ref_frames = br.ue()
        br.u(1)  # gaps_in_frame_num_value_allowed_flag
        pic_width_mbs = br.ue() + 1
        pic_height_map_units = br.ue() + 1
        frame_mbs_only = br.u(1)
        out["max_num_ref_frames"] = max_ref_frames
        out["frame_mbs_only_flag"] = frame_mbs_only
        out["width"] = pic_width_mbs * 16
        if not frame_mbs_only:
            br.u(1)  # mb_adaptive_frame_field_flag
        out["height"] = (2 - frame_mbs_only) * pic_height_map_units * 16

        br.u(1)  # direct_8x8_inference_flag
        crop = br.u(1)
        if crop:
            br.ue(); br.ue(); br.ue(); br.ue()

        vui = br.u(1)
        out["vui_present"] = bool(vui)
        if vui:
            ar = br.u(1)
            if ar:
                idc = br.u(8)
                out["aspect_ratio_idc"] = idc
                if idc == 255:
                    br.u(16); br.u(16)
            overscan = br.u(1)
            if overscan:
                br.u(1)
            vsig = br.u(1)
            if vsig:
                br.u(3); br.u(1)
                cd = br.u(1)
                if cd:
                    br.u(8); br.u(8); br.u(8)
            cl = br.u(1)
            if cl:
                br.ue(); br.ue()
            ti = br.u(1)
            out["timing_info_present"] = bool(ti)
            vui_fps = None
            if ti:
                num_units = br.u(32)
                time_scale = br.u(32)
                fixed_fr = br.u(1)
                vui_fps = time_scale / (2 * num_units) if num_units else None
                out["vui_fps"] = vui_fps
                out["fixed_frame_rate_flag"] = bool(fixed_fr)

            # Level margin check (Annex A) -- the layer that differs
            # between lenient decoders (ExoPlayer/VLC, don't enforce
            # this) and strict ones (many native Smart TV HLS stacks
            # do). Uses the SPS's own declared fps (from its VUI), so
            # this is a self-consistency check of the stream against
            # its own signaled level, not an external guess.
            if level_idc in LEVEL_LIMITS and "width" in out:
                max_mbps, max_fs, _ = LEVEL_LIMITS[level_idc]
                frame_mbs = (out["width"] // 16) * (out["height"] // 16)
                out["frame_macroblocks"] = frame_mbs
                out["level_max_fs"] = max_fs
                out["level_max_mbps"] = max_mbps
                if vui_fps:
                    needed_mbps = frame_mbs * vui_fps
                    out["needed_mbps_at_stream_fps"] = round(needed_mbps, 1)
                    headroom_pct = 100.0 * (max_mbps - needed_mbps) / max_mbps
                    out["mbps_headroom_pct"] = round(headroom_pct, 1)
                    if needed_mbps > max_mbps:
                        out["LEVEL_VIOLATION"] = (
                            f"Content needs {needed_mbps:.0f} MB/s but "
                            f"Level {level_idc/10:.1f} only allows "
                            f"{max_mbps} MB/s -- OVER the limit, a strict "
                            f"decoder can refuse this outright.")
                    elif headroom_pct < 5.0:
                        out["LEVEL_MARGIN_WARNING"] = (
                            f"Content needs {needed_mbps:.0f} MB/s against "
                            f"a Level {level_idc/10:.1f} ceiling of "
                            f"{max_mbps} MB/s -- only {headroom_pct:.1f}% "
                            f"headroom. Minor PCR/encoder jitter can push "
                            f"this over the limit on a strict decoder even "
                            f"though it nominally fits.")
                if frame_mbs > max_fs:
                    out["FS_VIOLATION"] = (
                        f"Frame size {frame_mbs} MBs exceeds Level "
                        f"{level_idc/10:.1f}'s MaxFS of {max_fs} MBs.")
        return out
    except Exception as e:
        out["error"] = f"SPS parse stopped early: {e}"
        return out


def parse_pps(pps_nal):
    out = {}
    try:
        rbsp = ebsp_to_rbsp(pps_nal[1:])
        br = BitReader(rbsp)
        br.ue()  # pps_id
        br.ue()  # sps_id
        entropy = br.u(1)
        out["entropy_coding_mode"] = "CABAC" if entropy else "CAVLC"
        bottom_field = br.u(1)
        num_slice_groups = br.ue() + 1
        out["num_slice_groups"] = num_slice_groups
        return out
    except Exception as e:
        out["error"] = f"PPS parse stopped early: {e}"
        return out


def find_first_nal(es, nal_type):
    for m in re.finditer(b'\x00\x00\x01', es):
        s = m.start()
        if s + 3 >= len(es):
            continue
        hdr = es[s+3]
        if (hdr & 0x1f) == nal_type:
            # find end: next start code
            end = len(es)
            for m2 in re.finditer(b'\x00\x00\x01', es[s+3:]):
                if m2.start() > 0:
                    end = s + 3 + m2.start()
                    break
            return es[s+3:end]
    return None


# --------------------------------------------------------------------------
# Manifest (EXTINF) parsing
# --------------------------------------------------------------------------

def parse_manifest_extinf(manifest_path):
    """Returns {segment_filename: extinf_seconds}."""
    out = {}
    if not manifest_path:
        return out
    with open(manifest_path, "r", errors="ignore") as f:
        lines = [l.strip() for l in f]
    pending = None
    for line in lines:
        if line.startswith("#EXTINF:"):
            try:
                pending = float(line[len("#EXTINF:"):].split(",")[0])
            except ValueError:
                pending = None
        elif line and not line.startswith("#"):
            if pending is not None:
                out[os.path.basename(line)] = pending
                pending = None
    return out


# --------------------------------------------------------------------------
# Report
# --------------------------------------------------------------------------

def analyze_one(path, key_hex, iv_hex, extinf_map):
    name = os.path.basename(path)
    print(f"\n{'='*70}")
    print(f"SEGMENT: {name}")
    print(f"{'='*70}")

    seq_num = seq_num_from_filename(path)
    print(f"  sequence number (from filename): {seq_num}")

    data, mode, orig_size, mod16, iv_used = decrypt_segment(path, key_hex, iv_hex, seq_num)
    print(f"  ciphertext size       : {orig_size} bytes  (mod 16 = {mod16})")
    if mod16 != 0:
        print(f"  *** ciphertext is NOT a multiple of 16 bytes (see note "
              f"below on whether this turns out to matter once decrypted "
              f"with the right IV).")
    print(f"  decrypt mode          : {mode}")
    if iv_used:
        print(f"  IV that worked        : {iv_used}")
    if data is None or data[:1] != b'\x47':
        print("  -> no candidate IV produced a valid TS sync byte (0x47) "
              "for this segment -- skipping further analysis.")
        return

    pkts = list(iter_ts_packets(data))
    print(f"  TS packets            : {len(pkts)}")

    sync_err, tei_err, cc_err, pcr_count = cc_and_sync_check(pkts)
    print(f"  sync errors           : {sync_err}")
    print(f"  TEI errors            : {tei_err}")
    print(f"  continuity errors     : {cc_err}")
    print(f"  PCR packets           : {pcr_count}")

    pmt_pid, streams = find_pat_pmt(pkts)
    if not streams:
        print("  *** Could not parse PAT/PMT -- no PID info available.")
        return
    print(f"  PMT pid               : 0x{pmt_pid:x}" if pmt_pid else "  PMT pid: none")
    video_pid = None
    audio_pid = None
    for st, pid in streams:
        tag = ""
        if st in (0x1b, 0x24) and video_pid is None:
            video_pid = pid
            tag = " (video)"
        elif st in (0x03, 0x04, 0x0f, 0x11, 0x06, 0x81, 0x87) and audio_pid is None:
            audio_pid = pid
            tag = " (audio)"
        print(f"    stream_type=0x{st:02x} pid=0x{pid:x}{tag}")

    if video_pid:
        ves = extract_es(pkts, video_pid)
        nals = find_nal_units(ves)
        hist = {}
        for t, _ in nals:
            hist[t] = hist.get(t, 0) + 1
        NAL_NAMES = {1: "P/B-slice", 5: "IDR", 6: "SEI", 7: "SPS",
                     8: "PPS", 9: "AUD"}
        hist_str = ", ".join(
            f"{NAL_NAMES.get(t, f'type{t}')}={c}" for t, c in sorted(hist.items()))
        print(f"  video NAL histogram   : {hist_str}")
        first_slice = next((t for t, _ in nals if t in (1, 5)), None)
        if first_slice == 5:
            print("  leading slice type    : IDR (5) -- segment opens cleanly")
        elif first_slice == 1:
            print("  *** leading slice type: NON-IDR (1) -- MID-GOP CUT. "
                  "A strict decoder has nothing to anchor to at the start "
                  "of this fragment.")
        else:
            print("  *** no slice NAL (type 1 or 5) found at all in this "
                  "segment's video stream.")

        vpts = find_pes_pts_dts(ves, max_scan=10)
        if vpts:
            first_pts = vpts[0][0]
            last_pts = vpts[-1][0]
            print(f"  video first PTS (90k) : {first_pts}  "
                  f"({first_pts/90000:.3f}s)")
            print(f"  video last PTS sampled: {last_pts}  "
                  f"({last_pts/90000:.3f}s)  [first {len(vpts)} PES headers only]")
        else:
            print("  *** could not extract any video PTS from PES headers.")

        # SPS/PPS bitstream-level compliance check -- this is the layer
        # that most often differs between lenient decoders (ExoPlayer,
        # VLC) and strict native Smart TV HLS stacks. A segment can pass
        # every container/timestamp check above and still fail here.
        sps_nal = find_first_nal(ves, 7)
        pps_nal = find_first_nal(ves, 8)
        if sps_nal:
            sps = parse_sps(sps_nal)
            if "error" in sps:
                print(f"  SPS parse             : {sps['error']}")
            else:
                print(f"  SPS profile/level     : {sps.get('profile_name')} "
                      f"({sps.get('profile_idc')}) / Level {sps.get('level')}  "
                      f"constraint_flags={sps.get('constraint_flags')}")
                print(f"  SPS resolution        : {sps.get('width')}x{sps.get('height')}  "
                      f"frame_mbs_only_flag={sps.get('frame_mbs_only_flag')}  "
                      f"max_ref_frames={sps.get('max_num_ref_frames')}")
                if sps.get("vui_present"):
                    print(f"  SPS VUI               : fps={sps.get('vui_fps')}  "
                          f"fixed_frame_rate={sps.get('fixed_frame_rate_flag')}  "
                          f"aspect_ratio_idc={sps.get('aspect_ratio_idc')}")
                else:
                    print("  SPS VUI               : not present")
                if "needed_mbps_at_stream_fps" in sps:
                    print(f"  Level margin check    : needs "
                          f"{sps['needed_mbps_at_stream_fps']} MB/s, Level "
                          f"{sps.get('level')} allows {sps.get('level_max_mbps')} "
                          f"MB/s -> {sps.get('mbps_headroom_pct')}% headroom")
                for key in ("LEVEL_VIOLATION", "LEVEL_MARGIN_WARNING", "FS_VIOLATION"):
                    if key in sps:
                        print(f"  *** {key}: {sps[key]}")
                if "note" in sps:
                    print(f"  (SPS note: {sps['note']})")
        else:
            print("  *** no SPS NAL found in this segment.")
        if pps_nal:
            pps = parse_pps(pps_nal)
            if "error" in pps:
                print(f"  PPS parse             : {pps['error']}")
            else:
                print(f"  PPS                   : entropy={pps.get('entropy_coding_mode')}  "
                      f"num_slice_groups={pps.get('num_slice_groups')}")
        else:
            print("  *** no PPS NAL found in this segment.")
    else:
        print("  *** no video PID found in PMT.")

    if audio_pid:
        aes = extract_es(pkts, audio_pid)
        apts = find_pes_pts_dts(aes, max_scan=5)
        if apts:
            first_apts = apts[0][0]
            print(f"  audio first PTS (90k) : {first_apts}  "
                  f"({first_apts/90000:.3f}s)")
            if video_pid and vpts:
                gap = (first_apts - vpts[0][0]) / 90000.0
                print(f"  audio-video start gap : {gap:+.3f}s "
                      f"(audio minus video; ~0 is healthy, seconds-scale "
                      f"values indicate a PTS anchor problem)")
        else:
            print("  *** could not extract any audio PTS from PES headers.")
    else:
        print("  *** no audio PID found in PMT.")

    if extinf_map:
        extinf = extinf_map.get(name)
        if extinf is not None:
            print(f"  manifest EXTINF       : {extinf:.3f}s")
            if video_pid and len(vpts) >= 2:
                print(f"  (compare against measured video PTS span above; "
                      f"large mismatches point at the playlist-duration "
                      f"issue)")
        else:
            print(f"  (no EXTINF entry found for {name} in the manifest "
                  f"you passed)")


def main():
    ap = argparse.ArgumentParser(description=__doc__,
                                  formatter_class=argparse.RawDescriptionHelpFormatter)
    ap.add_argument("--key", required=True, help="path to enc.key (16 raw bytes, or 32-char hex text)")
    ap.add_argument("--iv", help="hex IV, e.g. 2aa96af110f99fd7cfccddee74f52e86 (optional -- sequence-derived IV is tried automatically as a fallback/alternative)")
    ap.add_argument("--keyinfo", help="path to .keyinfo file (3rd line used as IV) if --iv not given")
    ap.add_argument("--segments", nargs="+", required=True, help="encrypted .ts segment paths, in order")
    ap.add_argument("--manifest", help="optional .m3u8 to cross-check EXTINF durations and show the actual #EXT-X-KEY line")
    args = ap.parse_args()

    key_hex = read_key_hex(args.key)
    iv_hex = read_iv(args)
    extinf_map = parse_manifest_extinf(args.manifest) if args.manifest else {}
    key_line = find_key_line_in_manifest(args.manifest)

    print("ts_diag.py report")
    print(f"segments analyzed: {len(args.segments)}")
    print(f"key file         : {args.key}  (not printed)")
    print(f"iv (fixed, given): {iv_hex if iv_hex else '(none given -- relying on sequence-derived IV only)'}")
    print("note: each segment is also tried with the HLS-spec sequence-")
    print("      number-derived IV (used automatically when the manifest's")
    print("      #EXT-X-KEY has no explicit IV= attribute) -- whichever")
    print("      candidate actually produces a valid TS sync byte is used,")
    print("      and reported per-segment below.")
    if args.manifest:
        print(f"manifest         : {args.manifest}  ({len(extinf_map)} EXTINF entries parsed)")
        if key_line:
            print(f"manifest's #EXT-X-KEY line:")
            print(f"  {key_line}")
            if "IV=" not in key_line.upper():
                print("  *** No IV= attribute in this line -- per HLS spec, the IV")
                print("      MUST be derived from each segment's media sequence")
                print("      number. If a fixed --iv was also given above, it will")
                print("      NOT be correct for this channel; the sequence-derived")
                print("      candidate below is the one that should actually work.")
        else:
            print("  (no #EXT-X-KEY line found in the manifest)")

    for seg in args.segments:
        analyze_one(seg, key_hex, iv_hex, extinf_map)

    print(f"\n{'='*70}")
    print("END OF REPORT -- share this text output back for review.")
    print(f"{'='*70}")


if __name__ == "__main__":
    main()
