use std::collections::HashMap;
use std::io::Write;
use oxideav_core::{
Error as CoreError, Muxer, Packet, Result as CoreResult, StreamInfo, WriteSeek,
};
use crate::{TS_PACKET_LEN, TS_SYNC_BYTE};
const PAT_PID: u16 = 0x0000;
const PMT_PID: u16 = 0x0100;
const FIRST_ES_PID: u16 = 0x1011;
const PROGRAM_NUMBER: u16 = 1;
pub fn open(output: Box<dyn WriteSeek>, streams: &[StreamInfo]) -> CoreResult<Box<dyn Muxer>> {
MpegTsMuxer::new(output, streams).map(|m| Box::new(m) as Box<dyn Muxer>)
}
pub struct MpegTsMuxer {
output: Box<dyn WriteSeek>,
tracks: Vec<TrackState>,
idx_to_track: HashMap<u32, usize>,
pcr_pid: u16,
pat_cc: u8,
pmt_cc: u8,
header_written: bool,
}
impl std::fmt::Debug for MpegTsMuxer {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("MpegTsMuxer")
.field("tracks", &self.tracks.len())
.field("pcr_pid", &self.pcr_pid)
.field("header_written", &self.header_written)
.finish()
}
}
#[derive(Debug, Clone)]
struct TrackState {
pid: u16,
stream_type: u8,
stream_id: u8,
cc: u8,
is_video: bool,
}
impl MpegTsMuxer {
fn new(output: Box<dyn WriteSeek>, streams: &[StreamInfo]) -> CoreResult<Self> {
if streams.is_empty() {
return Err(CoreError::invalid(
"mpegts muxer: at least one stream required",
));
}
let mut tracks = Vec::with_capacity(streams.len());
let mut idx_to_track = HashMap::with_capacity(streams.len());
let mut next_pid = FIRST_ES_PID;
for (i, s) in streams.iter().enumerate() {
let (stream_type, stream_id, is_video) =
stream_type_for_codec(s.params.codec_id.as_str()).ok_or_else(|| {
CoreError::unsupported(format!(
"mpegts muxer: no MPEG-TS stream_type for codec_id '{}'",
s.params.codec_id.as_str()
))
})?;
tracks.push(TrackState {
pid: next_pid,
stream_type,
stream_id,
cc: 0,
is_video,
});
idx_to_track.insert(s.index, i);
next_pid = next_pid.wrapping_add(1);
let _ = i;
}
let pcr_pid = tracks
.iter()
.find(|t| t.is_video)
.map(|t| t.pid)
.unwrap_or(tracks[0].pid);
Ok(Self {
output,
tracks,
idx_to_track,
pcr_pid,
pat_cc: 0,
pmt_cc: 0,
header_written: false,
})
}
fn write_psi_packet(&mut self, pid: u16, cc: &mut u8, section: &[u8]) -> CoreResult<()> {
if section.len() > TS_PACKET_LEN - 5 {
return Err(CoreError::invalid(format!(
"mpegts muxer: PSI section is {} bytes — multi-packet PSI is not yet supported",
section.len()
)));
}
let mut pkt = [0xFFu8; TS_PACKET_LEN];
pkt[0] = TS_SYNC_BYTE;
pkt[1] = 0b0100_0000 | ((pid >> 8) as u8 & 0b0001_1111);
pkt[2] = pid as u8;
pkt[3] = 0b0001_0000 | (*cc & 0x0F);
pkt[4] = 0;
let copy = section.len().min(TS_PACKET_LEN - 5);
pkt[5..5 + copy].copy_from_slice(§ion[..copy]);
self.output.write_all(&pkt).map_err(CoreError::Io)?;
*cc = (*cc + 1) & 0x0F;
Ok(())
}
fn write_pat(&mut self) -> CoreResult<()> {
let mut body = vec![
0x00, 0xB0,
0x0D, 0x00,
0x01, 0xC1, 0x00,
0x00, (PROGRAM_NUMBER >> 8) as u8,
PROGRAM_NUMBER as u8,
(0xE0 | ((PMT_PID >> 8) as u8 & 0x1F)),
PMT_PID as u8,
];
let crc = mpeg2_crc32(&body);
body.extend_from_slice(&crc.to_be_bytes());
let mut cc = self.pat_cc;
self.write_psi_packet(PAT_PID, &mut cc, &body)?;
self.pat_cc = cc;
Ok(())
}
fn write_pmt(&mut self) -> CoreResult<()> {
let mut body: Vec<u8> = Vec::new();
body.push(0x02); body.push(0xB0);
body.push(0x00);
body.extend_from_slice(&[
(PROGRAM_NUMBER >> 8) as u8,
PROGRAM_NUMBER as u8,
0xC1, 0x00, 0x00, 0xE0 | ((self.pcr_pid >> 8) as u8 & 0x1F),
self.pcr_pid as u8,
0xF0, 0x00, ]);
for t in &self.tracks {
body.push(t.stream_type);
body.push(0xE0 | ((t.pid >> 8) as u8 & 0x1F));
body.push(t.pid as u8);
body.push(0xF0); body.push(0x00); }
let section_length = (body.len() - 3 + 4) as u16;
body[1] = 0xB0 | ((section_length >> 8) as u8 & 0x0F);
body[2] = section_length as u8;
let crc = mpeg2_crc32(&body);
body.extend_from_slice(&crc.to_be_bytes());
let mut cc = self.pmt_cc;
self.write_psi_packet(PMT_PID, &mut cc, &body)?;
self.pmt_cc = cc;
Ok(())
}
fn write_pes_packet(&mut self, track_idx: usize, packet: &Packet) -> CoreResult<()> {
let track = self.tracks[track_idx].clone();
let pes = build_pes(track.stream_id, packet);
let pcr = if track.is_video {
packet.pts.map(|p| {
let pcr_base = (p.saturating_sub(900)).max(0) as u64;
(pcr_base, 0u16)
})
} else {
None
};
self.write_pes_bytes_as_ts(track.pid, track_idx, &pes, pcr)
}
fn write_pes_bytes_as_ts(
&mut self,
pid: u16,
track_idx: usize,
pes: &[u8],
first_packet_pcr: Option<(u64, u16)>,
) -> CoreResult<()> {
let mut cursor = 0usize;
let mut first = true;
while cursor < pes.len() {
let mut pkt = [0xFFu8; TS_PACKET_LEN];
pkt[0] = TS_SYNC_BYTE;
let pusi = if first { 0b0100_0000 } else { 0 };
pkt[1] = pusi | ((pid >> 8) as u8 & 0b0001_1111);
pkt[2] = pid as u8;
let cc = self.tracks[track_idx].cc;
let remaining = pes.len() - cursor;
let payload_room_no_af = TS_PACKET_LEN - 4;
let needs_pcr = first && first_packet_pcr.is_some();
let needs_stuffing = remaining < payload_room_no_af && !needs_pcr;
let af_present = needs_pcr || needs_stuffing;
let afc = if af_present { 0b11 } else { 0b01 };
pkt[3] = (afc << 4) | (cc & 0x0F);
self.tracks[track_idx].cc = (cc + 1) & 0x0F;
let mut payload_start = 4;
if af_present {
let mut af = Vec::<u8>::with_capacity(TS_PACKET_LEN - 4);
let mut flags = 0u8;
if let Some((pcr_base, pcr_ext)) = first_packet_pcr {
if first {
flags |= 0b0001_0000; let high32 = ((pcr_base >> 1) & 0xFFFF_FFFF) as u32;
let low_bit_of_base = (pcr_base & 1) as u8;
let mut pcr_bytes = [0u8; 6];
pcr_bytes[0..4].copy_from_slice(&high32.to_be_bytes());
pcr_bytes[4] =
(low_bit_of_base << 7) | 0b0111_1110 | ((pcr_ext >> 8) as u8 & 0x01);
pcr_bytes[5] = pcr_ext as u8;
af.push(flags);
af.extend_from_slice(&pcr_bytes);
} else {
af.push(flags);
}
} else {
af.push(flags);
}
let header_so_far = 4 + 1 + af.len();
let mut payload_room = TS_PACKET_LEN - header_so_far;
if payload_room > remaining {
let stuff_bytes = payload_room - remaining;
af.resize(af.len() + stuff_bytes, 0xFF);
payload_room = remaining;
}
let af_total_len = af.len() as u8;
pkt[4] = af_total_len;
pkt[5..5 + af.len()].copy_from_slice(&af);
payload_start = 4 + 1 + af.len();
let take = payload_room.min(remaining);
pkt[payload_start..payload_start + take]
.copy_from_slice(&pes[cursor..cursor + take]);
cursor += take;
} else {
let take = payload_room_no_af.min(remaining);
pkt[4..4 + take].copy_from_slice(&pes[cursor..cursor + take]);
cursor += take;
let _ = payload_start;
}
self.output.write_all(&pkt).map_err(CoreError::Io)?;
first = false;
}
Ok(())
}
}
impl Muxer for MpegTsMuxer {
fn format_name(&self) -> &str {
"mpegts"
}
fn write_header(&mut self) -> CoreResult<()> {
self.write_pat()?;
self.write_pmt()?;
self.header_written = true;
Ok(())
}
fn write_packet(&mut self, packet: &Packet) -> CoreResult<()> {
if !self.header_written {
return Err(CoreError::invalid(
"mpegts muxer: write_packet called before write_header",
));
}
let track_idx = *self.idx_to_track.get(&packet.stream_index).ok_or_else(|| {
CoreError::invalid(format!(
"mpegts muxer: packet for unknown stream_index {}",
packet.stream_index
))
})?;
self.write_pes_packet(track_idx, packet)
}
fn write_trailer(&mut self) -> CoreResult<()> {
self.write_pat()?;
self.write_pmt()?;
self.output.flush().map_err(CoreError::Io)?;
Ok(())
}
}
fn build_pes(stream_id: u8, packet: &Packet) -> Vec<u8> {
let has_pts = packet.pts.is_some();
let has_dts = packet.dts.is_some() && packet.dts != packet.pts;
let pts_dts_flags: u8 = match (has_pts, has_dts) {
(true, true) => 0b11,
(true, false) => 0b10,
_ => 0b00,
};
let opt_hdr_len: u8 = match pts_dts_flags {
0b11 => 10, 0b10 => 5,
_ => 0,
};
let mut out = Vec::with_capacity(9 + opt_hdr_len as usize + packet.data.len());
out.extend_from_slice(&[0x00, 0x00, 0x01]); out.push(stream_id);
out.push(0x00);
out.push(0x00);
out.push(0b1000_0000);
out.push(pts_dts_flags << 6);
out.push(opt_hdr_len);
if let Some(p) = packet.pts {
let prefix = if has_dts { 0b0011 } else { 0b0010 };
out.extend_from_slice(&encode_pts_dts(prefix, p as u64));
}
if let (true, Some(d)) = (has_dts, packet.dts) {
out.extend_from_slice(&encode_pts_dts(0b0001, d as u64));
}
out.extend_from_slice(&packet.data);
out
}
fn encode_pts_dts(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 << 1) | 1) & 0xFF) as u8),
((t14_0 >> 7) & 0xFF) as u8,
(((t14_0 << 1) | 1) & 0xFF) as u8,
]
}
fn stream_type_for_codec(codec_id: &str) -> Option<(u8, u8, bool)> {
Some(match codec_id {
"mpeg2video" => (0x02, 0xE0, true),
"h264" => (0x1B, 0xE0, true),
"hevc" => (0x24, 0xE0, true),
"vc1" => (0xEA, 0xE0, true),
"pcm_s16be" => (0x80, 0xBD, false),
"ac3" => (0x81, 0xBD, false),
"dts" => (0x82, 0xBD, false),
"truehd" => (0x83, 0xBD, false),
"eac3" => (0x84, 0xBD, false),
"hdmv_pgs_subtitle" => (0x90, 0xBD, false),
"hdmv_textst_subtitle" => (0x92, 0xBD, false),
_ => return None,
})
}
fn mpeg2_crc32(data: &[u8]) -> u32 {
let mut c: u32 = 0xFFFF_FFFF;
for &b in data {
c ^= (b as u32) << 24;
for _ in 0..8 {
c = if c & 0x8000_0000 != 0 {
(c << 1) ^ 0x04C1_1DB7
} else {
c << 1
};
}
}
c
}
#[cfg(test)]
mod tests {
use super::*;
use oxideav_core::{CodecId, CodecParameters, ReadSeek, TimeBase};
use std::io::Cursor;
use std::sync::{Arc, Mutex};
fn stream_info(idx: u32, codec_id: &str, is_video: bool) -> StreamInfo {
let params = if is_video {
CodecParameters::video(CodecId::new(codec_id))
} else {
CodecParameters::audio(CodecId::new(codec_id))
};
StreamInfo {
index: idx,
time_base: TimeBase::new(1, 90_000),
duration: None,
start_time: None,
params,
}
}
#[derive(Clone)]
struct SharedSink(Arc<Mutex<Cursor<Vec<u8>>>>);
impl SharedSink {
fn new() -> Self {
Self(Arc::new(Mutex::new(Cursor::new(Vec::new()))))
}
fn into_bytes(self) -> Vec<u8> {
let inner = Arc::try_unwrap(self.0)
.map_err(|_| "shared sink still has live references")
.unwrap()
.into_inner()
.unwrap();
inner.into_inner()
}
}
impl std::io::Write for SharedSink {
fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
self.0.lock().unwrap().write(buf)
}
fn flush(&mut self) -> std::io::Result<()> {
self.0.lock().unwrap().flush()
}
}
impl std::io::Seek for SharedSink {
fn seek(&mut self, pos: std::io::SeekFrom) -> std::io::Result<u64> {
self.0.lock().unwrap().seek(pos)
}
}
#[test]
fn mux_then_demux_round_trip_avc_plus_ac3() {
let streams = vec![stream_info(0, "h264", true), stream_info(1, "ac3", false)];
let sink = SharedSink::new();
{
let output: Box<dyn oxideav_core::WriteSeek> = Box::new(sink.clone());
let mut mx = MpegTsMuxer::new(output, &streams).expect("open");
mx.write_header().expect("hdr");
let video_pkt = Packet::new(0, TimeBase::new(1, 90_000), vec![0xDE, 0xAD, 0xBE, 0xEF])
.with_pts(12345);
let audio_pkt =
Packet::new(1, TimeBase::new(1, 90_000), vec![0xCA, 0xFE]).with_pts(22222);
mx.write_packet(&video_pkt).unwrap();
mx.write_packet(&audio_pkt).unwrap();
mx.write_trailer().unwrap();
}
let bytes = sink.into_bytes();
assert!(bytes.len() % TS_PACKET_LEN == 0);
let input: Box<dyn ReadSeek> = Box::new(Cursor::new(bytes));
let resolver = oxideav_core::NullCodecResolver;
let mut dmx = crate::demuxer::open(input, &resolver).expect("dmx open");
assert_eq!(dmx.streams().len(), 2);
let codecs: Vec<&str> = dmx
.streams()
.iter()
.map(|s| s.params.codec_id.as_str())
.collect();
assert!(codecs.contains(&"h264"));
assert!(codecs.contains(&"ac3"));
let p1 = dmx.next_packet().expect("p1");
let p2 = dmx.next_packet().expect("p2");
let mut pts_set = vec![p1.pts.unwrap(), p2.pts.unwrap()];
pts_set.sort();
assert_eq!(pts_set, vec![12345, 22222]);
}
#[test]
fn open_rejects_empty_stream_list() {
let output: Box<dyn oxideav_core::WriteSeek> = Box::new(Cursor::new(Vec::<u8>::new()));
let r = MpegTsMuxer::new(output, &[]);
assert!(r.is_err());
}
#[test]
fn open_rejects_unknown_codec_id() {
let s = stream_info(0, "this-is-not-a-known-codec", true);
let output: Box<dyn oxideav_core::WriteSeek> = Box::new(Cursor::new(Vec::<u8>::new()));
let r = MpegTsMuxer::new(output, &[s]);
assert!(r.is_err());
}
}