use alloc::collections::VecDeque;
use alloc::vec::Vec;
use broadcast_common::{Demand, Parse, Stage, Timestamp};
use crate::aac_asc::AudioSpecificConfig;
use crate::avc_config::{AVCConfigurationBox, AVCDecoderConfigurationRecord};
use crate::flv::{
AUDIO_SAMPLE_SIZE_BITS, CODEC_ID_AVC, FLV_HEADER_LEN, FLV_SIGNATURE, FLV_TIMESCALE,
FRAME_TYPE_KEYFRAME, FlvError, MAX_FLV_HEADER_LEN, PREV_TAG_SIZE_LEN, TAG_HEADER_LEN,
aac_packet_type, asc_rate_hz, avc_packet_type, build_aac_esds, read_si24, tag_type,
};
use crate::ir::DemuxEvent;
use crate::pipeline::{CodecConfig, Sample, TrackSpec};
#[derive(Debug)]
struct PendingSample {
sample: Sample,
dts: u32,
}
#[derive(Debug, Default)]
struct TrackState {
track_id: Option<u32>,
pending: Option<PendingSample>,
last_duration: u32,
}
impl TrackState {
fn advance(&mut self, sample: Sample, dts: u32, events: &mut VecDeque<DemuxEvent>) {
let track_id = self
.track_id
.expect("TrackState::advance called before the track's config resolved");
if let Some(prev) = self.pending.take() {
let duration = dts.saturating_sub(prev.dts);
self.last_duration = duration;
let mut emitted = prev.sample;
emitted.duration = Some(duration);
events.push_back(DemuxEvent::Sample {
track_id,
sample: emitted,
});
}
self.pending = Some(PendingSample { sample, dts });
}
fn flush(&mut self, events: &mut VecDeque<DemuxEvent>) {
let Some(track_id) = self.track_id else {
return;
};
if let Some(prev) = self.pending.take() {
let mut emitted = prev.sample;
emitted.duration = Some(self.last_duration);
events.push_back(DemuxEvent::Sample {
track_id,
sample: emitted,
});
}
}
}
#[derive(Debug)]
pub struct StreamingFlvDemux {
pending: Vec<u8>,
header_seen: bool,
video: TrackState,
audio: TrackState,
next_track_id: u32,
events: VecDeque<DemuxEvent>,
}
impl Default for StreamingFlvDemux {
fn default() -> Self {
Self::new()
}
}
impl StreamingFlvDemux {
pub fn new() -> Self {
Self {
pending: Vec::new(),
header_seen: false,
video: TrackState::default(),
audio: TrackState::default(),
next_track_id: 1,
events: VecDeque::new(),
}
}
pub fn feed(&mut self, input: &[u8]) -> Result<(), FlvError> {
self.pending.extend_from_slice(input);
loop {
if !self.header_seen {
if self.pending.len() < FLV_HEADER_LEN + PREV_TAG_SIZE_LEN {
break; }
if self.pending[0..3] != FLV_SIGNATURE {
return Err(FlvError::BadSignature([
self.pending[0],
self.pending[1],
self.pending[2],
]));
}
let data_offset = u32::from_be_bytes([
self.pending[5],
self.pending[6],
self.pending[7],
self.pending[8],
]);
if data_offset as usize > MAX_FLV_HEADER_LEN {
return Err(FlvError::HeaderTooLarge {
declared: data_offset,
max: MAX_FLV_HEADER_LEN,
});
}
let skip = (data_offset as usize).max(FLV_HEADER_LEN) + PREV_TAG_SIZE_LEN;
if self.pending.len() < skip {
break; }
self.pending.drain(0..skip);
self.header_seen = true;
continue;
}
if self.pending.len() < TAG_HEADER_LEN {
break; }
let tag_type_byte = self.pending[0];
let data_size =
u32::from_be_bytes([0, self.pending[1], self.pending[2], self.pending[3]]) as usize;
let ts_lo = u32::from_be_bytes([0, self.pending[4], self.pending[5], self.pending[6]]);
let ts_ext = self.pending[7] as u32;
let timestamp = (ts_ext << 24) | ts_lo;
let body_start = TAG_HEADER_LEN;
let body_end = body_start + data_size;
let total = body_end + PREV_TAG_SIZE_LEN;
if self.pending.len() < total {
break; }
let body: Vec<u8> = self.pending[body_start..body_end].to_vec();
Self::process_tag(
&mut self.video,
&mut self.audio,
&mut self.next_track_id,
tag_type_byte,
timestamp,
&body,
&mut self.events,
)?;
self.pending.drain(0..total);
}
Ok(())
}
pub fn poll_event(&mut self) -> Option<DemuxEvent> {
self.events.pop_front()
}
pub fn finish(&mut self) {
self.video.flush(&mut self.events);
self.audio.flush(&mut self.events);
}
fn bytes_wanted(&self) -> usize {
let have = self.pending.len();
if !self.header_seen {
let minimum = FLV_HEADER_LEN + PREV_TAG_SIZE_LEN;
if have < minimum {
return minimum - have;
}
let data_offset = u32::from_be_bytes([
self.pending[5],
self.pending[6],
self.pending[7],
self.pending[8],
]) as usize;
let skip = data_offset.max(FLV_HEADER_LEN) + PREV_TAG_SIZE_LEN;
return skip.saturating_sub(have).max(1);
}
if have < TAG_HEADER_LEN {
return TAG_HEADER_LEN - have;
}
let data_size =
u32::from_be_bytes([0, self.pending[1], self.pending[2], self.pending[3]]) as usize;
let total = TAG_HEADER_LEN + data_size + PREV_TAG_SIZE_LEN;
total.saturating_sub(have).max(1)
}
fn process_tag(
video: &mut TrackState,
audio: &mut TrackState,
next_track_id: &mut u32,
tag_type_byte: u8,
timestamp: u32,
body: &[u8],
events: &mut VecDeque<DemuxEvent>,
) -> Result<(), FlvError> {
match tag_type_byte {
tag_type::VIDEO => {
Self::process_video_tag(video, next_track_id, timestamp, body, events)
}
tag_type::AUDIO => {
Self::process_audio_tag(audio, next_track_id, timestamp, body, events)
}
tag_type::SCRIPT => Ok(()), _ => Ok(()), }
}
fn process_video_tag(
video: &mut TrackState,
next_track_id: &mut u32,
timestamp: u32,
body: &[u8],
events: &mut VecDeque<DemuxEvent>,
) -> Result<(), FlvError> {
if body.len() < 2 {
return Ok(());
}
let frame_type = body[0] >> 4;
let codec_id = body[0] & 0x0F;
if codec_id != CODEC_ID_AVC {
return Ok(()); }
let avc_packet_type_byte = body[1];
if body.len() < 5 {
return Ok(());
}
let composition_time = read_si24(&body[2..5]);
let data = &body[5..];
match avc_packet_type_byte {
avc_packet_type::SEQUENCE_HEADER if video.track_id.is_none() && !data.is_empty() => {
let record = AVCDecoderConfigurationRecord::parse(data)?;
let config = AVCConfigurationBox::new(record);
let (width, height) = config
.config
.sps
.first()
.and_then(|sps| crate::sps::decode_avc_sps(&sps.0).ok())
.map(|i| (i.width as u16, i.height as u16))
.unwrap_or((0, 0));
let track_id = *next_track_id;
*next_track_id += 1;
video.track_id = Some(track_id);
let spec = TrackSpec::new(
track_id,
FLV_TIMESCALE,
CodecConfig::Avc {
config,
width,
height,
},
);
events.push_back(DemuxEvent::TrackAdded(spec));
}
avc_packet_type::NALU if video.track_id.is_some() => {
let dts_abs = timestamp as i64;
let pts_abs = dts_abs + composition_time as i64;
let sample = Sample {
data: data.to_vec().into(),
dts: Some(dts_abs),
pts: Some(pts_abs),
duration: None, flags: crate::ir::SampleFlags::new(frame_type == FRAME_TYPE_KEYFRAME),
provenance: None,
};
video.advance(sample, timestamp, events);
}
avc_packet_type::END_OF_SEQUENCE => {}
_ => {}
}
Ok(())
}
fn process_audio_tag(
audio: &mut TrackState,
next_track_id: &mut u32,
timestamp: u32,
body: &[u8],
events: &mut VecDeque<DemuxEvent>,
) -> Result<(), FlvError> {
if body.is_empty() {
return Ok(());
}
let sound_format = body[0] >> 4;
if sound_format != crate::flv::SOUND_FORMAT_AAC {
return Ok(()); }
if body.len() < 2 {
return Ok(());
}
let aac_pkt_type = body[1];
let data = &body[2..];
match aac_pkt_type {
aac_packet_type::SEQUENCE_HEADER if audio.track_id.is_none() && !data.is_empty() => {
let asc = AudioSpecificConfig::parse(data)?;
let channels = asc.channel_configuration.raw() as u16;
let rate = asc_rate_hz(&asc);
let esds = build_aac_esds(data.to_vec());
let track_id = *next_track_id;
*next_track_id += 1;
audio.track_id = Some(track_id);
let spec = TrackSpec::new(
track_id,
FLV_TIMESCALE,
CodecConfig::Aac {
esds,
channel_count: channels,
sample_rate: rate,
sample_size: AUDIO_SAMPLE_SIZE_BITS,
},
);
events.push_back(DemuxEvent::TrackAdded(spec));
}
aac_packet_type::RAW if audio.track_id.is_some() => {
let dts_abs = timestamp as i64;
let sample = Sample {
data: data.to_vec().into(),
dts: Some(dts_abs),
pts: Some(dts_abs),
duration: None, flags: crate::ir::SampleFlags::SYNC,
provenance: None,
};
audio.advance(sample, timestamp, events);
}
_ => {}
}
Ok(())
}
}
impl Stage for StreamingFlvDemux {
type In<'a> = &'a [u8];
type Out = DemuxEvent;
type Error = FlvError;
fn feed(&mut self, input: &[u8], _now: Timestamp) -> Result<(), FlvError> {
StreamingFlvDemux::feed(self, input)
}
fn poll(&mut self) -> Option<Self::Out> {
self.poll_event()
}
fn finish(&mut self) -> Result<(), FlvError> {
StreamingFlvDemux::finish(self);
Ok(())
}
fn next_deadline(&self) -> Option<Timestamp> {
None
}
fn on_deadline(&mut self, _now: Timestamp) {}
fn demand(&self) -> Demand {
Demand::new(self.bytes_wanted())
}
}
#[cfg(test)]
mod tests {
use super::*;
use alloc::vec;
use bytes::Bytes;
fn minimal_avcc_bytes() -> Vec<u8> {
let sps_nal: [u8; 6] = [0x67, 0x42, 0x00, 0x1F, 0x00, 0x00];
let mut out = vec![
0x01, 0x42, 0x00, 0x1F, 0xFF, 0xE1, ];
out.extend_from_slice(&(sps_nal.len() as u16).to_be_bytes());
out.extend_from_slice(&sps_nal);
out.push(0x00); out
}
fn write_tag(out: &mut Vec<u8>, tag_type: u8, timestamp: u32, body: &[u8]) {
let start = out.len();
out.push(tag_type);
out.extend_from_slice(&(body.len() as u32).to_be_bytes()[1..]); out.push((timestamp >> 16) as u8);
out.push((timestamp >> 8) as u8);
out.push(timestamp as u8);
out.push((timestamp >> 24) as u8); out.extend_from_slice(&[0, 0, 0]); out.extend_from_slice(body);
let tag_size = (out.len() - start) as u32;
out.extend_from_slice(&tag_size.to_be_bytes()); }
fn flv_header() -> Vec<u8> {
header_with_data_offset(FLV_HEADER_LEN as u32)
}
fn header_with_data_offset(data_offset: u32) -> Vec<u8> {
let mut out = Vec::new();
out.extend_from_slice(&FLV_SIGNATURE);
out.push(1); out.push(0x01); out.extend_from_slice(&data_offset.to_be_bytes());
out.extend_from_slice(&0u32.to_be_bytes()); out
}
fn feed_and_drain(
demux: &mut StreamingFlvDemux,
input: &[u8],
) -> Result<Vec<DemuxEvent>, FlvError> {
demux.feed(input)?;
let mut events = Vec::new();
while let Some(ev) = demux.poll_event() {
events.push(ev);
}
Ok(events)
}
fn finish_and_drain(demux: &mut StreamingFlvDemux) -> Vec<DemuxEvent> {
demux.finish();
let mut events = Vec::new();
while let Some(ev) = demux.poll_event() {
events.push(ev);
}
events
}
fn synthetic_avc_flv(n: u32, step: u32) -> Vec<u8> {
let mut out = flv_header();
let mut seq_body = vec![(1u8 << 4) | CODEC_ID_AVC, avc_packet_type::SEQUENCE_HEADER];
seq_body.extend_from_slice(&[0, 0, 0]); seq_body.extend_from_slice(&minimal_avcc_bytes());
write_tag(&mut out, tag_type::VIDEO, 0, &seq_body);
for i in 0..n {
let ts = i * step;
let mut body = vec![(1u8 << 4) | CODEC_ID_AVC, avc_packet_type::NALU];
body.extend_from_slice(&[0, 0, 0]); body.push(i as u8); write_tag(&mut out, tag_type::VIDEO, ts, &body);
}
out
}
fn video_samples(events: &[DemuxEvent]) -> Vec<(u32, Bytes)> {
events
.iter()
.filter_map(|e| match e {
DemuxEvent::Sample { sample, .. } => Some((
sample
.duration
.expect("FLV video samples always carry a duration"),
sample.data.clone(),
)),
_ => None,
})
.collect()
}
#[test]
fn whole_buffer_one_feed_call_matches_one_shot_shape() {
let flv = synthetic_avc_flv(4, 100);
let mut demux = StreamingFlvDemux::new();
let mut events = feed_and_drain(&mut demux, &flv).unwrap();
events.extend(finish_and_drain(&mut demux));
let added = events
.iter()
.filter(|e| matches!(e, DemuxEvent::TrackAdded(_)))
.count();
assert_eq!(added, 1, "exactly one TrackAdded (video only)");
let samples = video_samples(&events);
assert_eq!(samples.len(), 4, "4 NALU tags -> 4 samples");
assert_eq!(
samples.iter().map(|(d, _)| *d).collect::<Vec<_>>(),
vec![100, 100, 100, 100]
);
assert_eq!(
samples.iter().map(|(_, b)| b[0]).collect::<Vec<_>>(),
vec![0, 1, 2, 3]
);
}
#[test]
fn chunked_feed_matches_whole_buffer_feed() {
let flv = synthetic_avc_flv(6, 33);
let mut whole = StreamingFlvDemux::new();
let mut whole_events = feed_and_drain(&mut whole, &flv).unwrap();
whole_events.extend(finish_and_drain(&mut whole));
let whole_samples = video_samples(&whole_events);
let mut chunked = StreamingFlvDemux::new();
let mut chunked_events = Vec::new();
for chunk in flv.chunks(7) {
chunked_events.extend(feed_and_drain(&mut chunked, chunk).unwrap());
}
chunked_events.extend(finish_and_drain(&mut chunked));
let chunked_samples = video_samples(&chunked_events);
assert_eq!(
whole_samples, chunked_samples,
"7-byte chunking must reproduce the whole-buffer result exactly"
);
}
#[test]
fn byte_at_a_time_feed_matches_whole_buffer_feed() {
let flv = synthetic_avc_flv(5, 40);
let mut whole = StreamingFlvDemux::new();
let mut whole_events = feed_and_drain(&mut whole, &flv).unwrap();
whole_events.extend(finish_and_drain(&mut whole));
let whole_samples = video_samples(&whole_events);
let mut byte_demux = StreamingFlvDemux::new();
let mut byte_events = Vec::new();
for b in &flv {
byte_events.extend(feed_and_drain(&mut byte_demux, core::slice::from_ref(b)).unwrap());
}
byte_events.extend(finish_and_drain(&mut byte_demux));
let byte_samples = video_samples(&byte_events);
assert_eq!(
whole_samples, byte_samples,
"byte-at-a-time feed must reproduce the whole-buffer result exactly"
);
}
#[test]
fn pending_buffer_stays_bounded_at_tag_boundaries() {
let flv = synthetic_avc_flv(200, 10);
let mut demux = StreamingFlvDemux::new();
let last_tag_len = TAG_HEADER_LEN + 2 + 3 + 1 + PREV_TAG_SIZE_LEN; let split = flv.len() - last_tag_len;
demux.feed(&flv[..split]).unwrap();
assert_eq!(
demux.pending.len(),
0,
"pending must be empty at a clean tag boundary, not accumulate the whole stream"
);
demux.feed(&flv[split..split + 1]).unwrap();
assert_eq!(
demux.pending.len(),
1,
"pending must hold only the in-progress partial tag"
);
assert!(
demux.pending.len() < last_tag_len,
"pending must never grow to hold a whole stream's worth of tags"
);
}
#[test]
fn bad_signature_is_an_error_not_a_panic() {
let mut bad = flv_header();
bad[0] = b'X'; let mut demux = StreamingFlvDemux::new();
let err = demux.feed(&bad).unwrap_err();
assert!(matches!(err, FlvError::BadSignature(_)));
}
#[test]
fn truncated_header_waits_without_erroring_or_panicking() {
let flv = synthetic_avc_flv(2, 10);
let mut demux = StreamingFlvDemux::new();
demux.feed(&flv[..5]).unwrap();
assert!(demux.poll_event().is_none());
}
#[test]
fn corrupt_sequence_header_is_a_codec_error_not_a_panic() {
let mut out = flv_header();
let mut seq_body = vec![(1u8 << 4) | CODEC_ID_AVC, avc_packet_type::SEQUENCE_HEADER];
seq_body.extend_from_slice(&[0, 0, 0]); seq_body.push(0xFF); write_tag(&mut out, tag_type::VIDEO, 0, &seq_body);
let mut demux = StreamingFlvDemux::new();
let err = demux.feed(&out).unwrap_err();
assert!(matches!(err, FlvError::Codec(_)));
}
fn zero_sps_avcc_bytes() -> Vec<u8> {
vec![
0x01, 0x42, 0x00, 0x1F, 0xFF, 0xE0, 0x00, ]
}
#[test]
fn zero_sps_avcc_is_an_error_not_a_panic() {
let mut out = flv_header();
let mut seq_body = vec![(1u8 << 4) | CODEC_ID_AVC, avc_packet_type::SEQUENCE_HEADER];
seq_body.extend_from_slice(&[0, 0, 0]); seq_body.extend_from_slice(&zero_sps_avcc_bytes());
write_tag(&mut out, tag_type::VIDEO, 0, &seq_body);
let mut demux = StreamingFlvDemux::new();
let err = demux.feed(&out).unwrap_err();
assert!(
matches!(err, FlvError::Codec(_)),
"expected FlvError::Codec (0-SPS avcC rejected), got {err:?}"
);
}
#[test]
fn absurd_data_offset_is_rejected_immediately_buffer_stays_small() {
let mut header = header_with_data_offset(0xFFFF_FFFF);
header.extend_from_slice(&[0, 0, 0, 0]); let fed_len = header.len();
let mut demux = StreamingFlvDemux::new();
let err = demux.feed(&header).unwrap_err();
assert!(
matches!(
err,
FlvError::HeaderTooLarge { declared, max }
if declared == 0xFFFF_FFFF && max == MAX_FLV_HEADER_LEN
),
"expected HeaderTooLarge, got {err:?}"
);
assert_eq!(
demux.pending.len(),
fed_len,
"pending must hold only what was actually fed ({fed_len} bytes) — the demux \
must not attempt to wait for the malicious DataOffset's implied ~4 GiB header"
);
}
#[test]
fn unknown_tag_type_is_skipped_leniently() {
let mut out = flv_header();
write_tag(&mut out, 0xAA, 0, &[1, 2, 3]);
let mut demux = StreamingFlvDemux::new();
demux.feed(&out).unwrap();
assert!(
demux.poll_event().is_none(),
"unknown tag types must be skipped, not error"
);
}
#[test]
fn header_is_only_consumed_once() {
let flv = synthetic_avc_flv(2, 25);
let mut demux = StreamingFlvDemux::new();
let mut events = Vec::new();
for b in &flv[..FLV_HEADER_LEN + PREV_TAG_SIZE_LEN] {
events.extend(feed_and_drain(&mut demux, core::slice::from_ref(b)).unwrap());
}
events.extend(
feed_and_drain(&mut demux, &flv[FLV_HEADER_LEN + PREV_TAG_SIZE_LEN..]).unwrap(),
);
events.extend(finish_and_drain(&mut demux));
let samples = video_samples(&events);
assert_eq!(samples.len(), 2);
assert_eq!(samples[0].1[0], 0);
assert_eq!(samples[1].1[0], 1);
}
}