use crate::{TsError, TsPacket};
fn has_optional_pes_header(stream_id: u8) -> bool {
!matches!(
stream_id,
0xBC | 0xBE | 0xBF | 0xF0 | 0xF1 | 0xFF | 0xF2 | 0xF8 )
}
#[derive(Debug, Clone)]
pub struct PesPacket {
pub stream_id: u8,
pub pts_90k: Option<u64>,
pub dts_90k: Option<u64>,
pub payload: Vec<u8>,
}
impl PesPacket {
pub fn parse(bytes: &[u8]) -> Result<Self, TsError> {
if bytes.len() < 6 {
return Err(TsError::Truncated {
what: "PES header",
have: bytes.len(),
need: 6,
});
}
if bytes[0] != 0x00 || bytes[1] != 0x00 || bytes[2] != 0x01 {
return Err(TsError::BadPesStartCode([bytes[0], bytes[1], bytes[2]]));
}
let stream_id = bytes[3];
let _pes_packet_length = u16::from_be_bytes([bytes[4], bytes[5]]);
if !has_optional_pes_header(stream_id) {
return Ok(Self {
stream_id,
pts_90k: None,
dts_90k: None,
payload: bytes[6..].to_vec(),
});
}
if bytes.len() < 9 {
return Err(TsError::Truncated {
what: "PES optional header",
have: bytes.len(),
need: 9,
});
}
let pts_dts_flags = (bytes[7] >> 6) & 0b11;
let pes_header_data_length = bytes[8] as usize;
let header_end = 9 + pes_header_data_length;
if bytes.len() < header_end {
return Err(TsError::Truncated {
what: "PES optional header body",
have: bytes.len(),
need: header_end,
});
}
let mut pts_90k = None;
let mut dts_90k = None;
let opt = &bytes[9..header_end];
let mut cursor = 0usize;
match pts_dts_flags {
0b10 => {
if opt.len() < 5 {
return Err(TsError::Truncated {
what: "PES PTS",
have: opt.len(),
need: 5,
});
}
pts_90k = Some(decode_timestamp(&opt[..5])?);
cursor += 5;
}
0b11 => {
if opt.len() < 10 {
return Err(TsError::Truncated {
what: "PES PTS+DTS",
have: opt.len(),
need: 10,
});
}
pts_90k = Some(decode_timestamp(&opt[..5])?);
dts_90k = Some(decode_timestamp(&opt[5..10])?);
cursor += 10;
}
0b00 => { }
_ => return Err(TsError::Unsupported("PES PTS_DTS_flags = 0b01")),
}
let _ = cursor;
Ok(Self {
stream_id,
pts_90k,
dts_90k,
payload: bytes[header_end..].to_vec(),
})
}
}
fn decode_timestamp(b: &[u8]) -> Result<u64, TsError> {
if b.len() < 5 {
return Err(TsError::Truncated {
what: "PTS/DTS",
have: b.len(),
need: 5,
});
}
let t32_30 = ((b[0] >> 1) & 0b0000_0111) as u64;
let t29_22 = b[1] as u64;
let t21_15 = ((b[2] >> 1) & 0b0111_1111) as u64;
let t14_7 = b[3] as u64;
let t6_0 = ((b[4] >> 1) & 0b0111_1111) as u64;
let ts = (t32_30 << 30) | (t29_22 << 22) | (t21_15 << 15) | (t14_7 << 7) | t6_0;
Ok(ts)
}
#[derive(Debug, Default)]
pub struct PesReassembler {
buf: Vec<u8>,
started: bool,
}
impl PesReassembler {
pub fn new() -> Self {
Self::default()
}
pub fn feed(&mut self, ts: &TsPacket<'_>) -> Result<Option<PesPacket>, TsError> {
if ts.payload_unit_start {
let finished = if self.started {
Some(PesPacket::parse(&self.buf)?)
} else {
None
};
self.buf.clear();
self.buf.extend_from_slice(ts.payload);
self.started = true;
Ok(finished)
} else if self.started {
self.buf.extend_from_slice(ts.payload);
Ok(None)
} else {
Ok(None)
}
}
pub fn flush(&mut self) -> Result<Option<PesPacket>, TsError> {
if !self.started {
return Ok(None);
}
let buf = std::mem::take(&mut self.buf);
self.started = false;
Ok(Some(PesPacket::parse(&buf)?))
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::packet::{TS_PACKET_LEN, TS_SYNC_BYTE};
fn encode_timestamp(prefix: u8, ts: u64) -> [u8; 5] {
let t = ts & 0x1_FFFF_FFFF;
let t32_30 = ((t >> 30) & 0b0111) as u8;
let t29_15 = ((t >> 15) & 0x7FFF) as u16;
let t14_0 = (t & 0x7FFF) as u16;
[
(prefix << 4) | (t32_30 << 1) | 0b1,
((t29_15 >> 7) & 0xFF) as u8,
(((t29_15 & 0x7F) << 1) as u8) | 0b1,
((t14_0 >> 7) & 0xFF) as u8,
(((t14_0 & 0x7F) << 1) as u8) | 0b1,
]
}
fn build_pes(stream_id: u8, pts: u64, payload: &[u8]) -> Vec<u8> {
let pts_bytes = encode_timestamp(0b0010, pts);
let mut v = Vec::new();
v.extend_from_slice(&[0x00, 0x00, 0x01, stream_id]);
let pes_packet_length: u16 = (3 + 5 + payload.len()) as u16;
v.extend_from_slice(&pes_packet_length.to_be_bytes());
v.push(0b1000_0000);
v.push(0b1000_0000);
v.push(5);
v.extend_from_slice(&pts_bytes);
v.extend_from_slice(payload);
v
}
fn pes_into_ts(pid: u16, pes: &[u8], chunk_len: usize) -> Vec<u8> {
assert!(chunk_len > 0 && chunk_len <= 184);
let mut out = Vec::new();
let mut cursor = 0;
let mut first = true;
let mut cc: u8 = 0;
while cursor < pes.len() {
let pusi = if first { 0b0100_0000 } else { 0 };
let pid_hi = ((pid >> 8) & 0x1F) as u8;
let pid_lo = (pid & 0xFF) as u8;
let remaining = pes.len() - cursor;
let take = remaining.min(chunk_len).min(184);
let af_total = 184 - take;
let af_control: u8 = if af_total > 0 { 0b11 } else { 0b01 };
let b3 = (af_control << 4) | (cc & 0x0F);
let mut pkt = vec![TS_SYNC_BYTE, pusi | pid_hi, pid_lo, b3];
if af_total > 0 {
pkt.push((af_total - 1) as u8);
if af_total >= 2 {
pkt.push(0); pkt.extend(std::iter::repeat(0xFF).take(af_total - 2));
}
}
pkt.extend_from_slice(&pes[cursor..cursor + take]);
cursor += take;
assert_eq!(pkt.len(), TS_PACKET_LEN);
out.extend_from_slice(&pkt);
cc = (cc + 1) & 0x0F;
first = false;
}
out
}
#[test]
fn decode_timestamp_round_trip() {
for ts in [0u64, 1, 90_000, 0x1_FFFF_FFFF] {
let enc = encode_timestamp(0b0010, ts);
let dec = decode_timestamp(&enc).unwrap();
assert_eq!(dec, ts);
}
}
#[test]
fn parse_complete_pes_packet_pts_only() {
let pes = build_pes(0xE0, 90_000, b"hello world");
let parsed = PesPacket::parse(&pes).unwrap();
assert_eq!(parsed.stream_id, 0xE0);
assert_eq!(parsed.pts_90k, Some(90_000));
assert_eq!(parsed.dts_90k, None);
assert_eq!(parsed.payload, b"hello world");
}
#[test]
fn parse_complete_pes_packet_pts_dts() {
let pts_bytes = encode_timestamp(0b0011, 200_000);
let dts_bytes = encode_timestamp(0b0001, 180_000);
let payload = b"AVCdata";
let mut v = Vec::new();
v.extend_from_slice(&[0x00, 0x00, 0x01, 0xE0]);
let pes_packet_length: u16 = (3 + 10 + payload.len()) as u16;
v.extend_from_slice(&pes_packet_length.to_be_bytes());
v.push(0b1000_0000);
v.push(0b1100_0000); v.push(10);
v.extend_from_slice(&pts_bytes);
v.extend_from_slice(&dts_bytes);
v.extend_from_slice(payload);
let parsed = PesPacket::parse(&v).unwrap();
assert_eq!(parsed.pts_90k, Some(200_000));
assert_eq!(parsed.dts_90k, Some(180_000));
assert_eq!(parsed.payload, payload);
}
#[test]
fn reassemble_pes_split_across_three_ts_packets() {
let payload: Vec<u8> = (0..400u32).map(|i| (i & 0xFF) as u8).collect();
let pes = build_pes(0xE0, 12345, &payload);
let ts_buf = pes_into_ts(0x100, &pes, 150);
let next_pes = build_pes(0xE0, 67890, b"X");
let next_ts = pes_into_ts(0x100, &next_pes, 184);
let mut full = ts_buf;
full.extend_from_slice(&next_ts);
let mut r = PesReassembler::new();
let mut packets = Vec::new();
for pkt in crate::iter_packets(&full) {
let pkt = pkt.unwrap();
if let Some(done) = r.feed(&pkt).unwrap() {
packets.push(done);
}
}
if let Some(done) = r.flush().unwrap() {
packets.push(done);
}
assert_eq!(packets.len(), 2);
assert_eq!(packets[0].pts_90k, Some(12345));
assert_eq!(&packets[0].payload[..payload.len()], &payload[..]);
assert_eq!(packets[1].pts_90k, Some(67890));
}
#[test]
fn pusi_starts_new_packet_and_emits_previous() {
let pes1 = build_pes(0xC0, 1000, b"first");
let pes2 = build_pes(0xC0, 2000, b"second");
let ts1 = pes_into_ts(0x101, &pes1, 184);
let ts2 = pes_into_ts(0x101, &pes2, 184);
let mut r = PesReassembler::new();
let mut emitted = Vec::new();
for buf in [&ts1, &ts2] {
for pkt in crate::iter_packets(buf) {
let pkt = pkt.unwrap();
if let Some(done) = r.feed(&pkt).unwrap() {
emitted.push(done);
}
}
}
assert_eq!(emitted.len(), 1);
assert_eq!(emitted[0].pts_90k, Some(1000));
assert_eq!(&emitted[0].payload[..5], b"first");
let flushed = r.flush().unwrap().expect("buffered pes2");
assert_eq!(flushed.pts_90k, Some(2000));
assert_eq!(&flushed.payload[..6], b"second");
}
#[test]
fn padding_stream_has_no_optional_header() {
let mut v = Vec::new();
v.extend_from_slice(&[0x00, 0x00, 0x01, 0xBE]);
let payload = [0xFFu8; 12];
let len: u16 = payload.len() as u16;
v.extend_from_slice(&len.to_be_bytes());
v.extend_from_slice(&payload);
let p = PesPacket::parse(&v).unwrap();
assert_eq!(p.stream_id, 0xBE);
assert_eq!(p.pts_90k, None);
assert_eq!(p.dts_90k, None);
assert_eq!(p.payload, &payload);
}
#[test]
fn bad_pes_start_code_rejected() {
let mut v = vec![0x00, 0x00, 0x02, 0xE0, 0, 0];
v.extend_from_slice(&[0u8; 3]);
let err = PesPacket::parse(&v).unwrap_err();
match err {
TsError::BadPesStartCode(_) => {}
other => panic!("expected BadPesStartCode, got {other:?}"),
}
}
}