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);
}
}