import sys
from collections import Counter

def parse_ts(data):
    pkts = []
    for i in range(0, len(data) - 187, 188):
        p = data[i:i+188]
        if p[0] != 0x47:
            continue
        pid  = ((p[1] & 0x1F) << 8) | p[2]
        pusi = (p[1] >> 6) & 1
        afc  = (p[3] >> 4) & 3
        cc   = p[3] & 0xF
        if   afc == 1: pay_off = 4
        elif afc == 3: pay_off = 5 + p[4]
        else:          pay_off = -1
        pkts.append({'pid':pid,'pusi':pusi,'afc':afc,'cc':cc,
                     'pay_off':pay_off,'raw':p})
    return pkts

def get_payload(pkt):
    o = pkt['pay_off']
    if o < 0 or o >= 188: return b''
    return pkt['raw'][o:]

def parse_pat(pay):
    if len(pay) < 8: return None
    ptr = pay[0]
    sec = pay[ptr+1:]
    if len(sec) < 8 or sec[0] != 0x00: return None
    sec_len = ((sec[1] & 0x0F) << 8) | sec[2]
    pmts = {}
    for i in range(8, 3 + sec_len - 4, 4):
        if i + 3 >= len(sec): break
        pn = (sec[i] << 8) | sec[i+1]
        pp = ((sec[i+2] & 0x1F) << 8) | sec[i+3]
        if pn: pmts[pn] = pp
    return pmts

def parse_pmt(pay):
    if len(pay) < 13: return None
    ptr = pay[0]
    sec = pay[ptr+1:]
    if len(sec) < 12 or sec[0] != 0x02: return None
    sec_len = ((sec[1] & 0x0F) << 8) | sec[2]
    pcr_pid = ((sec[8] & 0x1F) << 8) | sec[9]
    pil     = ((sec[10] & 0x0F) << 8) | sec[11]
    streams = {}
    es  = 12 + pil
    ee  = 3 + sec_len - 4
    while es + 4 < len(sec) and es + 4 < ee:
        st = sec[es]
        sp = ((sec[es+1] & 0x1F) << 8) | sec[es+2]
        el = ((sec[es+3] & 0x0F) << 8) | sec[es+4]
        streams[sp] = st
        es += 5 + el
    return {'pcr': pcr_pid, 'streams': streams}

def parse_pes_pts(pay):
    if len(pay) < 14: return None
    if pay[0] != 0 or pay[1] != 0 or pay[2] != 1: return None
    if (pay[7] & 0x80) == 0: return None
    p = pay[9:]
    pts = ((p[0] & 0x0E) << 29) | (p[1] << 22) | \
          ((p[2] & 0xFE) << 14) | (p[3] << 7)  | \
          ((p[4] & 0xFE) >> 1)
    return pts

if len(sys.argv) < 2:
    print("Usage: analyze_seg_aac.py <segment.ts>")
    sys.exit(1)

data = open(sys.argv[1], 'rb').read()
print("File : %s" % sys.argv[1])
print("Size : %d bytes,  %d TS packets" % (len(data), len(data)//188))

pkts = parse_ts(data)
print("Valid: %d packets with sync byte 0x47" % len(pkts))

# PAT
pat_pmt = {}
for p in pkts:
    if p['pid'] == 0 and p['pusi']:
        pay = get_payload(p)
        r = parse_pat(pay)
        if r:
            pat_pmt = r
            print("PAT  : PMT PIDs = %s" % {hex(k):hex(v) for k,v in r.items()})
            break

# PMT
pmt_info = None
for pmt_pid in pat_pmt.values():
    for p in pkts:
        if p['pid'] == pmt_pid and p['pusi']:
            pay = get_payload(p)
            r = parse_pmt(pay)
            if r:
                pmt_info = r
                print("PMT  : pcr=0x%04X" % r['pcr'])
                TYPE = {0x01:'MPEG1-V',0x02:'MPEG2-V',0x1B:'H264',0x24:'HEVC',
                        0x03:'MP1',0x04:'MP2',0x0F:'AAC-ADTS',0x11:'AAC-LATM',
                        0x06:'private',0x81:'AC3',0x82:'EAC3'}
                for pid,st in r['streams'].items():
                    print("       PID=0x%04X  type=0x%02X (%s)" % (
                          pid, st, TYPE.get(st, 'unknown')))
                break
    if pmt_info: break

if not pmt_info:
    print("ERROR: PMT not found in segment")
    sys.exit(1)

# Identify PIDs
vid_pid = None
aud_pid = None
aud_type = None
VID = (0x01,0x02,0x10,0x1B,0x24,0x27,0x42)
AUD = (0x03,0x04,0x06,0x0F,0x11,0x81,0x82)
for pid,st in pmt_info['streams'].items():
    if st in VID: vid_pid  = pid
    if st in AUD: aud_pid  = pid; aud_type = st

print("")
print("vid_pid = 0x%04X" % (vid_pid or 0))
print("aud_pid = 0x%04X  type=0x%02X" % (aud_pid or 0, aud_type or 0))

# PID packet counts
pid_counts = Counter(p['pid'] for p in pkts)
print("\nPID packet counts:")
for pid,cnt in pid_counts.most_common(12):
    tag = ""
    if pid == vid_pid:  tag = " <-- VIDEO"
    if pid == aud_pid:  tag = " <-- AUDIO"
    if pid == 0:        tag = " <-- PAT"
    for pmt_pid in pat_pmt.values():
        if pid == pmt_pid: tag = " <-- PMT"
    print("  0x%04X : %4d pkts%s" % (pid, cnt, tag))

if aud_pid is None:
    print("\nERROR: No audio PID declared in PMT")
    sys.exit(1)

aud_pkts = [p for p in pkts if p['pid'] == aud_pid]
print("\nAudio PID 0x%04X : %d TS packets" % (aud_pid, len(aud_pkts)))

if len(aud_pkts) == 0:
    print("ERROR: Zero audio packets written into segment!")
    sys.exit(1)

# Reassemble PES payloads
pes_list = []
cur = bytearray()
cur_pts = None
for p in aud_pkts:
    pay = get_payload(p)
    if p['pusi']:
        if cur:
            pes_list.append((cur_pts, bytes(cur)))
        cur = bytearray()
        cur_pts = parse_pes_pts(pay)
        if len(pay) > 9:
            hdr = 9 + pay[8]
            cur += pay[hdr:]
    else:
        cur += pay
if cur:
    pes_list.append((cur_pts, bytes(cur)))

print("PES units reassembled: %d" % len(pes_list))

adts_ok = 0
adts_bad = 0
for i, (pts, es) in enumerate(pes_list):
    ok = len(es) >= 2 and es[0] == 0xFF and (es[1] & 0xF0) == 0xF0
    if ok: adts_ok += 1
    else:  adts_bad += 1
    if i < 8:
        sync = "ADTS OK" if ok else ("BAD [%s]" % es[:6].hex())
        print("  PES[%d] pts=%-14s es_len=%4d  %s" % (
              i, str(pts), len(es), sync))

print("")
print("ADTS sync OK  : %d / %d" % (adts_ok, len(pes_list)))
if adts_bad:
    print("ADTS sync BAD : %d  <-- player cannot decode these" % adts_bad)

# Check video packets
if vid_pid:
    vid_pkts = [p for p in pkts if p['pid'] == vid_pid]
    print("\nVideo PID 0x%04X : %d TS packets" % (vid_pid, len(vid_pkts)))

print("\nDone.")

# ---- EXTRA: dump raw bytes of first audio PUSI TS packet ----
print("\n--- First audio PUSI packet raw dump ---")
for p in pkts:
    if p['pid'] == aud_pid and p['pusi']:
        raw = p['raw']
        print("TS header: %s" % raw[:4].hex())
        pay_o = p['pay_off']
        print("payload_offset=%d" % pay_o)
        pay = raw[pay_o:]
        print("PES start code: %s" % pay[:3].hex())
        print("stream_id: 0x%02X" % pay[3])
        print("PES_packet_length: %d" % ((pay[4]<<8)|pay[5]))
        print("flags: 0x%02X 0x%02X" % (pay[6], pay[7]))
        print("PES_header_data_length: %d" % pay[8])
        hdr = 9 + pay[8]
        print("ES starts at offset: %d" % hdr)
        print("First 8 ES bytes: %s" % pay[hdr:hdr+8].hex())
        adts_ok = pay[hdr]==0xFF and (pay[hdr+1]&0xF0)==0xF0
        print("ADTS sync: %s" % ("OK" if adts_ok else "BAD"))
        break