use crate::error::Error;
use std::io::{self, Write};
const TS_PACKET_SIZE: usize = 188;
const MAX_AF_LEN: usize = TS_PACKET_SIZE - 4 - 1;
const SYNC_BYTE: u8 = 0x47;
const STUFF_BYTE: u8 = 0xFF;
pub(super) struct Packet {
buf: [u8; TS_PACKET_SIZE],
len: usize,
}
impl Packet {
pub(super) fn new() -> Self {
Self {
buf: [0u8; TS_PACKET_SIZE],
len: 0,
}
}
fn push(&mut self, b: u8) {
if self.len < TS_PACKET_SIZE {
self.buf[self.len] = b;
self.len += 1;
}
}
fn extend(&mut self, bytes: &[u8]) {
let n = bytes.len().min(TS_PACKET_SIZE - self.len);
self.buf[self.len..self.len + n].copy_from_slice(&bytes[..n]);
self.len += n;
}
pub(super) fn set_header(
&mut self,
pid: u16,
payload_unit_start: bool,
has_payload: bool,
has_adaptation: bool,
cc: u8,
) {
self.len = 0;
self.push(SYNC_BYTE);
let pus_bit = if payload_unit_start { 0x40 } else { 0 };
self.push(pus_bit | ((pid >> 8) as u8 & 0x1F));
self.push(pid as u8);
let afc = match (has_adaptation, has_payload) {
(false, false) => 0b00, (false, true) => 0b01, (true, false) => 0b10, (true, true) => 0b11, };
self.push((afc << 4) | (cc & 0x0F));
}
pub(super) fn append_adaptation(&mut self, body: &[u8], stuffing: usize) -> io::Result<()> {
let af_len = body.len() + stuffing;
if af_len > MAX_AF_LEN {
return Err(Error::M2tsPacketMalformed.into());
}
self.push(af_len as u8);
self.extend(body);
for _ in 0..stuffing {
self.push(STUFF_BYTE);
}
Ok(())
}
pub(super) fn append_payload(&mut self, payload: &[u8]) -> io::Result<()> {
if self.len + payload.len() > TS_PACKET_SIZE {
return Err(Error::M2tsPacketMalformed.into());
}
self.extend(payload);
Ok(())
}
pub(super) fn pad_to_188(&mut self) {
while self.len < TS_PACKET_SIZE {
self.push(STUFF_BYTE);
}
}
pub(super) fn bytes(&self) -> &[u8] {
&self.buf[..self.len]
}
pub(super) fn len(&self) -> usize {
self.len
}
}
pub(super) struct PacketWriter<W: Write> {
inner: W,
}
impl<W: Write> PacketWriter<W> {
pub(super) fn new(inner: W) -> Self {
Self { inner }
}
pub(super) fn write_packet(&mut self, packet: &Packet) -> io::Result<()> {
let bytes = packet.bytes();
if bytes.len() != TS_PACKET_SIZE {
return Err(Error::M2tsPacketMalformed.into());
}
self.inner.write_all(bytes)
}
pub(super) fn flush(&mut self) -> io::Result<()> {
self.inner.flush()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn pad_fills_to_188() {
let mut p = Packet::new();
p.set_header(0x100, true, true, false, 0);
p.append_payload(&[1, 2, 3]).unwrap();
p.pad_to_188();
assert_eq!(p.bytes().len(), 188);
assert_eq!(p.bytes()[0], SYNC_BYTE);
assert_eq!(p.bytes()[7], STUFF_BYTE);
}
#[test]
fn append_adaptation_rejects_overflow() {
let mut p = Packet::new();
p.set_header(0x100, true, true, true, 0);
let err = p.append_adaptation(&[0x00], MAX_AF_LEN).unwrap_err();
assert_eq!(err.kind(), io::ErrorKind::InvalidData);
}
#[test]
fn append_payload_rejects_overflow() {
let mut p = Packet::new();
p.set_header(0x100, true, true, false, 0);
let err = p.append_payload(&[0u8; 185]).unwrap_err();
assert_eq!(err.kind(), io::ErrorKind::InvalidData);
}
#[test]
fn write_packet_rejects_short_packet() {
let mut p = Packet::new();
p.set_header(0x100, true, true, false, 0);
p.append_payload(&[1, 2, 3]).unwrap(); let mut sink: Vec<u8> = Vec::new();
let mut w = PacketWriter::new(&mut sink);
let err = w.write_packet(&p).unwrap_err();
assert_eq!(err.kind(), io::ErrorKind::InvalidData);
assert!(sink.is_empty(), "short packet must not be written");
}
#[test]
fn header_pid_round_trips() {
let mut p = Packet::new();
p.set_header(0x1ABC, false, true, false, 0xA);
let pid = u16::from_be_bytes([p.bytes()[1] & 0x1F, p.bytes()[2]]);
assert_eq!(pid, 0x1ABC);
assert_eq!(p.bytes()[3] & 0x0F, 0xA);
}
#[test]
fn header_sync_byte_and_pusi_bit() {
let mut p = Packet::new();
p.set_header(0x0100, true, true, false, 0);
assert_eq!(p.bytes()[0], SYNC_BYTE);
assert_eq!(p.bytes()[1] & 0x40, 0x40, "PUSI set");
assert_eq!(p.bytes()[1] & 0x80, 0, "TEI clear");
assert_eq!(p.bytes()[1] & 0x20, 0, "transport_priority clear");
let mut p2 = Packet::new();
p2.set_header(0x0100, false, true, false, 0);
assert_eq!(p2.bytes()[1] & 0x40, 0, "PUSI clear when not a unit start");
}
#[test]
fn header_afc_bits_per_combination() {
let cases = [
(false, true, 0b01u8),
(true, false, 0b10),
(true, true, 0b11),
(false, false, 0b00),
];
for (af, pl, want) in cases {
let mut p = Packet::new();
p.set_header(0x0100, true, pl, af, 0);
assert_eq!((p.bytes()[3] >> 4) & 0x03, want, "AFC for af={af} pl={pl}");
}
}
#[test]
fn append_adaptation_length_byte_matches_written_bytes() {
let mut p = Packet::new();
p.set_header(0x0100, true, true, true, 0);
p.append_adaptation(&[0x10, 0xAA, 0xBB], 4).unwrap(); assert_eq!(p.bytes()[4], 3 + 4, "length byte = body+stuffing");
assert_eq!(&p.bytes()[5..8], &[0x10, 0xAA, 0xBB]);
assert_eq!(&p.bytes()[8..12], &[0xFF; 4]);
}
#[test]
fn append_adaptation_at_exact_max_succeeds() {
let mut p = Packet::new();
p.set_header(0x0100, true, true, true, 0);
assert!(p.append_adaptation(&[0x00], MAX_AF_LEN - 1).is_ok());
assert_eq!(p.bytes()[4] as usize, MAX_AF_LEN);
}
#[test]
fn append_payload_at_exact_boundary_fills_188() {
let mut p = Packet::new();
p.set_header(0x0100, true, true, false, 0);
assert!(p.append_payload(&[0xAB; 184]).is_ok());
assert_eq!(p.len(), 188);
}
#[test]
fn pad_to_188_is_idempotent_when_already_full() {
let mut p = Packet::new();
p.set_header(0x0100, true, true, false, 0);
p.append_payload(&[0xAB; 184]).unwrap();
assert_eq!(p.len(), 188);
p.pad_to_188();
assert_eq!(p.len(), 188, "no growth past 188");
}
#[test]
fn write_packet_rejects_long_packet() {
let mut p = Packet::new();
p.set_header(0x0100, true, true, false, 0);
p.append_payload(&[1, 2, 3, 4, 5]).unwrap(); let mut sink: Vec<u8> = Vec::new();
let mut w = PacketWriter::new(&mut sink);
assert!(w.write_packet(&p).is_err());
assert!(sink.is_empty());
}
#[test]
fn write_packet_accepts_exactly_188() {
let mut p = Packet::new();
p.set_header(0x0100, true, true, false, 0);
p.append_payload(&[0x5A; 184]).unwrap();
let mut sink: Vec<u8> = Vec::new();
{
let mut w = PacketWriter::new(&mut sink);
w.write_packet(&p).unwrap();
}
assert_eq!(sink.len(), 188);
assert_eq!(sink[0], SYNC_BYTE);
}
#[test]
fn pid_high_bits_masked_to_13_bits() {
let mut p = Packet::new();
p.set_header(0xE100, false, true, false, 0);
assert_eq!(
p.bytes()[1] & 0xE0,
0,
"top 3 bits of byte1 are flags, not PID"
);
let pid = u16::from_be_bytes([p.bytes()[1] & 0x1F, p.bytes()[2]]);
assert_eq!(pid, 0xE100 & 0x1FFF, "PID masked to 13 bits");
}
}