use alloc::collections::BTreeMap;
use alloc::vec::Vec;
use core::marker::PhantomData;
use broadcast_common::Unpackage;
use mpeg_ps::program_stream::parse_all_packs;
use crate::ac3::Ac3SyncframeInfo;
use crate::annexb::iter_annexb_nals;
use crate::avc_config::{AVCConfigurationBox, AVCDecoderConfigurationRecord};
use crate::error::{Error, Result};
use crate::media::{Media, Track};
use crate::nalu_types::{AvcPps, AvcSps};
use crate::pipeline::{CodecConfig, Sample, TrackSpec};
const STREAM_ID_VIDEO_LO: u8 = 0xE0;
const STREAM_ID_VIDEO_HI: u8 = 0xEF;
const STREAM_ID_PRIVATE_1: u8 = 0xBD;
const PRIVATE1_AC3_HEADER_LEN: usize = 4;
const NAL_LENGTH_SIZE_MINUS_ONE: u8 = 3;
const H264_NAL_AUD: u8 = 9;
const H264_NAL_SPS: u8 = 7;
const H264_NAL_PPS: u8 = 8;
const H264_NAL_IDR: u8 = 5;
const H264_NAL_TYPE_MASK: u8 = 0x1F;
const VIDEO_TIMESCALE: u32 = 90_000;
const AUDIO_SAMPLE_SIZE_BITS: u16 = 16;
const TS_WRAP: i128 = 1 << 33;
const TS_WRAP_HALF: i128 = TS_WRAP / 2;
const DEFAULT_FRAME_DURATION: i128 = 3600;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum Codec {
H264,
Ac3,
}
impl Codec {
fn from_stream_id(stream_id: u8) -> Option<Self> {
match stream_id {
STREAM_ID_VIDEO_LO..=STREAM_ID_VIDEO_HI => Some(Codec::H264),
STREAM_ID_PRIVATE_1 => Some(Codec::Ac3),
_ => None,
}
}
}
struct ElementaryStream {
codec: Codec,
es_bytes: Vec<u8>,
stamps: Vec<Stamp>,
}
#[derive(Debug, Clone, Copy)]
struct Stamp {
offset: usize,
pts: Option<u64>,
dts: Option<u64>,
}
struct AccessUnit {
data: Vec<u8>,
pts: Option<u64>,
dts: Option<u64>,
}
#[derive(Debug, Default, Clone)]
pub struct PsDemux<'a> {
_marker: PhantomData<&'a [u8]>,
}
impl<'a> PsDemux<'a> {
pub fn new() -> Self {
Self {
_marker: PhantomData,
}
}
pub fn demux(&mut self, input: &'a [u8]) -> Result<Media> {
let (packs, _trailing) = parse_all_packs(input).map_err(Error::Ps)?;
let mut order: Vec<u8> = Vec::new();
let mut streams: BTreeMap<u8, ElementaryStream> = BTreeMap::new();
for pack in &packs {
for pes in &pack.pes_packets {
let sid = pes.stream_id.0;
let Some(codec) = Codec::from_stream_id(sid) else {
continue;
};
let payload: &[u8] = match codec {
Codec::Ac3 => {
if pes.payload.len() <= PRIVATE1_AC3_HEADER_LEN {
continue;
}
&pes.payload[PRIVATE1_AC3_HEADER_LEN..]
}
Codec::H264 => pes.payload,
};
if payload.is_empty() {
continue;
}
let (pts, dts) = pes
.header
.as_ref()
.map(|h| (h.pts.map(|p| p.0), h.dts.map(|d| d.0)))
.unwrap_or((None, None));
let es = streams.entry(sid).or_insert_with(|| {
order.push(sid);
ElementaryStream {
codec,
es_bytes: Vec::new(),
stamps: Vec::new(),
}
});
let offset = es.es_bytes.len();
if pts.is_some() || dts.is_some() {
es.stamps.push(Stamp { offset, pts, dts });
}
es.es_bytes.extend_from_slice(payload);
}
}
let mut tracks: Vec<Track> = Vec::new();
let mut track_id: u32 = 1;
for sid in &order {
let es = &streams[sid];
let built = match es.codec {
Codec::H264 => build_h264_track(es, track_id),
Codec::Ac3 => build_ac3_track(es, track_id),
};
if let Some(track) = built {
tracks.push(track);
track_id += 1;
}
}
Ok(Media::new(tracks, VIDEO_TIMESCALE))
}
}
impl<'a> Unpackage for PsDemux<'a> {
type Input = &'a [u8];
type Media = Media;
type Error = Error;
fn unpackage(&mut self, input: &'a [u8]) -> Result<Media> {
self.demux(input)
}
}
fn start_code_positions(data: &[u8]) -> Vec<usize> {
let mut positions = Vec::new();
let n = data.len();
let mut p = 0usize;
while p + 3 <= n {
if data[p] == 0 && data[p + 1] == 0 && data[p + 2] == 1 {
positions.push(p);
p += 3;
} else {
p += 1;
}
}
positions
}
fn split_access_units(data: &[u8]) -> Vec<(usize, usize)> {
let codes = start_code_positions(data);
let mut au_starts: Vec<usize> = Vec::new();
for &pos in &codes {
if pos + 3 < data.len() && (data[pos + 3] & H264_NAL_TYPE_MASK) == H264_NAL_AUD {
let start = if pos > 0 && data[pos - 1] == 0 {
pos - 1
} else {
pos
};
au_starts.push(start);
}
}
if au_starts.is_empty() {
return if data.is_empty() {
Vec::new()
} else {
alloc::vec![(0, data.len())]
};
}
let mut ranges: Vec<(usize, usize)> = Vec::with_capacity(au_starts.len());
let n = au_starts.len();
for i in 0..n {
let start = if i == 0 { 0 } else { au_starts[i] };
let end = if i + 1 < n {
au_starts[i + 1]
} else {
data.len()
};
ranges.push((start, end));
}
ranges
}
fn unwrap_ts(prev_unwrapped: i128, prev_raw: u64, raw: u64) -> i128 {
let mut delta = raw as i128 - prev_raw as i128;
if delta > TS_WRAP_HALF {
delta -= TS_WRAP;
} else if delta < -TS_WRAP_HALF {
delta += TS_WRAP;
}
prev_unwrapped + delta
}
fn assign_stamps(ranges: &[(usize, usize)], stamps: &[Stamp]) -> Vec<(Option<u64>, Option<u64>)> {
let mut out = alloc::vec![(None, None); ranges.len()];
let mut si = 0usize;
let (mut prev_pts_raw, mut prev_pts_uw): (Option<u64>, i128) = (None, 0);
let (mut prev_dts_raw, mut prev_dts_uw): (Option<u64>, i128) = (None, 0);
for (ai, &(start, _end)) in ranges.iter().enumerate() {
while si < stamps.len() && stamps[si].offset <= start {
let s = stamps[si];
let pts_uw = s.pts.map(|p| match prev_pts_raw {
Some(pr) => {
let uw = unwrap_ts(prev_pts_uw, pr, p);
prev_pts_uw = uw;
prev_pts_raw = Some(p);
uw
}
None => {
prev_pts_uw = p as i128;
prev_pts_raw = Some(p);
p as i128
}
});
let dts_uw = s.dts.map(|d| match prev_dts_raw {
Some(pr) => {
let uw = unwrap_ts(prev_dts_uw, pr, d);
prev_dts_uw = uw;
prev_dts_raw = Some(d);
uw
}
None => {
prev_dts_uw = d as i128;
prev_dts_raw = Some(d);
d as i128
}
});
if out[ai].0.is_none() && out[ai].1.is_none() {
out[ai] = (pts_uw.map(|v| v as u64), dts_uw.map(|v| v as u64));
}
si += 1;
}
}
out
}
fn build_h264_track(es: &ElementaryStream, track_id: u32) -> Option<Track> {
let ranges = split_access_units(&es.es_bytes);
if ranges.is_empty() {
return None;
}
let stamped = assign_stamps(&ranges, &es.stamps);
let mut sps: Option<Vec<u8>> = None;
let mut pps: Option<Vec<u8>> = None;
let mut units: Vec<AccessUnit> = Vec::with_capacity(ranges.len());
for (i, &(start, end)) in ranges.iter().enumerate() {
let au = &es.es_bytes[start..end];
for nal in iter_annexb_nals(au) {
match nal[0] & H264_NAL_TYPE_MASK {
H264_NAL_SPS if sps.is_none() => sps = Some(nal.to_vec()),
H264_NAL_PPS if pps.is_none() => pps = Some(nal.to_vec()),
_ => {}
}
}
units.push(AccessUnit {
data: au.to_vec(),
pts: stamped[i].0,
dts: stamped[i].1,
});
}
let sps = sps?;
let pps = pps?;
if sps.len() < 4 {
return None;
}
let record = AVCDecoderConfigurationRecord {
configuration_version: 1,
profile_indication: sps[1],
profile_compatibility: sps[2],
level_indication: sps[3],
length_size_minus_one: NAL_LENGTH_SIZE_MINUS_ONE,
sps: alloc::vec![AvcSps(sps)],
pps: alloc::vec![AvcPps(pps)],
chroma_format: None,
bit_depth_luma_minus8: None,
bit_depth_chroma_minus8: None,
sps_ext: alloc::vec![],
};
let config = AVCConfigurationBox::new(record);
let dts = interpolate_dts(&units);
let pts = interpolate_pts(&units, &dts);
let mut order: Vec<usize> = (0..units.len()).collect();
order.sort_by_key(|&i| dts[i]);
let samples: Vec<Sample> = order
.iter()
.enumerate()
.map(|(pos, &i)| {
let dur = frame_duration(&order, &dts, pos);
let is_idr = au_is_idr(&units[i].data);
let composition_offset = (pts[i] - dts[i]) as i32;
Sample::from_annexb(&units[i].data, dur, is_idr, composition_offset)
})
.collect();
Some(Track::new(
TrackSpec::new(
track_id,
VIDEO_TIMESCALE,
CodecConfig::Avc {
config,
width: 0,
height: 0,
},
),
samples,
))
}
fn au_is_idr(au: &[u8]) -> bool {
iter_annexb_nals(au).any(|nal| (nal[0] & H264_NAL_TYPE_MASK) == H264_NAL_IDR)
}
fn stamped_frame_duration(units: &[AccessUnit]) -> i128 {
let stamped: Vec<(usize, i128)> = units
.iter()
.enumerate()
.filter_map(|(i, u)| u.dts.map(|d| (i, d as i128)))
.collect();
if stamped.len() < 2 {
return DEFAULT_FRAME_DURATION;
}
let (i0, d0) = stamped[0];
let (i1, d1) = *stamped.last().unwrap();
let span_idx = (i1 - i0) as i128;
let span_dts = d1 - d0;
if span_idx > 0 && span_dts > 0 {
(span_dts / span_idx).max(1)
} else {
DEFAULT_FRAME_DURATION
}
}
fn interpolate_dts(units: &[AccessUnit]) -> Vec<i128> {
let dur = stamped_frame_duration(units);
let n = units.len();
let mut dts = alloc::vec![0i128; n];
let anchor = units
.iter()
.enumerate()
.find_map(|(i, u)| u.dts.map(|d| (i, d as i128)));
let (anchor_idx, anchor_dts) = anchor.unwrap_or((0, 0));
for (i, slot) in dts.iter_mut().enumerate() {
*slot = match units[i].dts {
Some(d) => d as i128,
None => anchor_dts + (i as i128 - anchor_idx as i128) * dur,
};
}
dts
}
fn interpolate_pts(units: &[AccessUnit], dts: &[i128]) -> Vec<i128> {
units
.iter()
.enumerate()
.map(|(i, u)| u.pts.map(|p| p as i128).unwrap_or(dts[i]))
.collect()
}
fn frame_duration(order: &[usize], dts: &[i128], pos: usize) -> u32 {
let n = order.len();
let dur = if pos + 1 < n {
(dts[order[pos + 1]] - dts[order[pos]]).max(0)
} else if pos > 0 {
(dts[order[pos]] - dts[order[pos - 1]]).max(0)
} else {
DEFAULT_FRAME_DURATION
};
dur as u32
}
fn split_ac3_frames(data: &[u8]) -> Vec<(usize, usize)> {
let mut syncs: Vec<usize> = Vec::new();
let mut i = 0usize;
while i + 1 < data.len() {
if data[i] == 0x0B && data[i + 1] == 0x77 {
syncs.push(i);
i += 2;
} else {
i += 1;
}
}
let mut ranges = Vec::with_capacity(syncs.len());
for k in 0..syncs.len() {
let start = syncs[k];
let end = if k + 1 < syncs.len() {
syncs[k + 1]
} else {
data.len()
};
ranges.push((start, end));
}
ranges
}
fn build_ac3_track(es: &ElementaryStream, track_id: u32) -> Option<Track> {
let info = Ac3SyncframeInfo::from_es(&es.es_bytes).ok()?;
let sample_rate = info.sample_rate;
let channel_count = info.channel_count() as u16;
let config = info.into_dac3();
let frames = split_ac3_frames(&es.es_bytes);
if frames.is_empty() {
return None;
}
let samples: Vec<Sample> = frames
.iter()
.map(|&(s, e)| Sample::from_raw(es.es_bytes[s..e].to_vec(), 0))
.collect();
Some(Track::new(
TrackSpec::new(
track_id,
sample_rate,
CodecConfig::Ac3 {
config,
channel_count,
sample_rate,
sample_size: AUDIO_SAMPLE_SIZE_BITS,
},
),
samples,
))
}