use broadcast_common::{Package, Parse, Serialize, Unpackage};
use bytes::Bytes;
use transmux::init_segment::{MovieBox, SampleEntryVariant, StblChild};
use transmux::{CmafMux, CodecConfig, FlvDemux, FlvMux};
const FLV: &[u8] = include_bytes!("../../fixtures/flv/av.flv");
const CSV: &str = include_str!("../../fixtures/flv/av.packets.csv");
const REF_MP4: &[u8] = include_bytes!("../../fixtures/ts/demux-oracle/h264_aac.ref.mp4");
#[derive(Debug, Clone, Copy)]
struct Pkt {
is_video: bool,
pts: i64,
dts: i64,
size: usize,
keyframe: bool,
}
fn oracle() -> Vec<Pkt> {
CSV.lines()
.filter(|l| !l.starts_with('#') && !l.trim().is_empty())
.map(|l| {
let c: Vec<&str> = l.split(',').collect();
Pkt {
is_video: c[0] == "video",
pts: c[2].parse().unwrap(),
dts: c[3].parse().unwrap(),
size: c[5].parse().unwrap(),
keyframe: c[6] == "1",
}
})
.collect()
}
fn find_top_box<'a>(data: &'a [u8], fourcc: &[u8; 4]) -> Option<&'a [u8]> {
let mut off = 0usize;
while off + 8 <= data.len() {
let size =
u32::from_be_bytes([data[off], data[off + 1], data[off + 2], data[off + 3]]) as usize;
let ty = &data[off + 4..off + 8];
let end = if size == 0 { data.len() } else { off + size };
if ty == fourcc {
return Some(&data[off..end]);
}
if size < 8 {
break;
}
off = end;
}
None
}
fn ref_mp4_avcc(mp4: &[u8]) -> Vec<u8> {
let moov = find_top_box(mp4, b"moov").expect("moov in ref.mp4");
let movie = MovieBox::parse(moov).expect("parse ref.mp4 moov");
for trak in &movie.tracks {
let Some(stbl) = trak
.mdia
.as_ref()
.and_then(|m| m.minf.as_ref())
.and_then(|m| m.stbl.as_ref())
else {
continue;
};
let Some(stsd) = stbl.children.iter().find_map(|c| match c {
StblChild::Stsd(s) => Some(s),
_ => None,
}) else {
continue;
};
if let Some(SampleEntryVariant::Avc1(avc1)) = stsd.entries.first() {
let mut body = vec![0u8; avc1.config.config.serialized_len()];
let n = avc1.config.config.serialize_into(&mut body).unwrap();
body.truncate(n);
return body;
}
}
panic!("no avc1 sample entry in ref.mp4");
}
#[test]
fn enumerate_two_tracks_avc_320x240_and_aac() {
let mut demux = FlvDemux::new();
let media = demux.unpackage(FLV).expect("demux av.flv");
assert_eq!(media.tracks.len(), 2, "must enumerate 2 tracks (AVC + AAC)");
match media.tracks[0].config() {
CodecConfig::Avc { width, height, .. } => {
assert_eq!((*width, *height), (320, 240), "AVC dims from SPS");
}
other => panic!("track 0 must be AVC, got {other:?}"),
}
assert!(
matches!(media.tracks[1].config(), CodecConfig::Aac { .. }),
"track 1 must be AAC"
);
}
#[test]
fn avcc_matches_ref_mp4_and_asc_decodes() {
let mut demux = FlvDemux::new();
let media = demux.unpackage(FLV).expect("demux av.flv");
let CodecConfig::Avc { config, .. } = media.tracks[0].config() else {
panic!("track 0 must be AVC");
};
let mut flv_avcc = alloc_body(config.config.serialized_len());
let n = config.config.serialize_into(&mut flv_avcc).unwrap();
flv_avcc.truncate(n);
let ref_avcc = ref_mp4_avcc(REF_MP4);
assert_eq!(
flv_avcc, ref_avcc,
"FLV-demuxed avcC must be byte-identical to the ref.mp4 avcC"
);
let CodecConfig::Aac {
esds,
channel_count,
sample_rate,
..
} = media.tracks[1].config()
else {
panic!("track 1 must be AAC");
};
let asc_bytes = esds
.es_descriptor
.decoder_config
.as_ref()
.unwrap()
.decoder_specific_info
.as_ref()
.unwrap()
.data
.clone();
let asc = transmux::AudioSpecificConfig::parse(&asc_bytes).expect("parse ASC");
assert_eq!(asc.channel_configuration.raw(), 1, "ASC channels = 1");
assert_eq!(*channel_count, 1, "config channel_count = 1");
assert_eq!(*sample_rate, 44100, "ASC sample rate = 44100 Hz");
}
#[test]
fn timestamps_and_keyframes_match_oracle() {
let mut demux = FlvDemux::new();
let media = demux.unpackage(FLV).expect("demux av.flv");
let ora = oracle();
let vid_ora: Vec<Pkt> = ora.iter().copied().filter(|p| p.is_video).collect();
let aud_ora: Vec<Pkt> = ora.iter().copied().filter(|p| !p.is_video).collect();
let vid = &media.tracks[0].samples;
let aud = &media.tracks[1].samples;
assert_eq!(vid.len(), 75, "video sample count");
assert_eq!(aud.len(), 131, "audio sample count");
assert_eq!(vid_ora.len(), 75);
assert_eq!(aud_ora.len(), 131);
check_timing(vid, &vid_ora, "video");
check_timing(aud, &aud_ora, "audio");
for (i, (s, p)) in vid.iter().zip(&vid_ora).enumerate() {
assert_eq!(
s.flags.is_sync, p.keyframe,
"video sample {i} keyframe flag"
);
}
let kf = vid.iter().filter(|s| s.flags.is_sync).count();
assert_eq!(kf, 3, "exactly 3 video keyframes");
}
fn check_timing(samples: &[transmux::Sample], ora: &[Pkt], kind: &str) {
let base_dts = ora[0].dts;
let mut dts = 0i64; for (i, (s, p)) in samples.iter().zip(ora).enumerate() {
assert_eq!(s.data.len(), p.size, "{kind} sample {i} payload size");
let exp_dts_rel = p.dts - base_dts;
assert_eq!(dts, exp_dts_rel, "{kind} sample {i} DTS (relative)");
let exp_pts_rel = p.pts - base_dts;
assert_eq!(
dts + s.composition_offset() as i64,
exp_pts_rel,
"{kind} sample {i} PTS (= DTS + composition offset)"
);
if i + 1 < samples.len() {
dts += s.duration.unwrap_or(0) as i64;
}
}
}
#[test]
fn flv_round_trip_preserves_samples_and_timing() {
let mut demux = FlvDemux::new();
let media = demux.unpackage(FLV).expect("demux av.flv");
let mut mux = FlvMux::new();
let flv2 = mux.package(&media).expect("mux to FLV");
let mut demux2 = FlvDemux::new();
let media2 = demux2.unpackage(&flv2).expect("re-demux FLV");
assert_eq!(media2.tracks.len(), 2, "round-trip track count");
for (a, b) in media.tracks.iter().zip(&media2.tracks) {
assert_eq!(
a.samples.len(),
b.samples.len(),
"track {} sample count preserved",
a.track_id()
);
for (i, (sa, sb)) in a.samples.iter().zip(&b.samples).enumerate() {
assert_eq!(sa.data, sb.data, "track {} sample {i} bytes", a.track_id());
assert_eq!(
sa.composition_offset(),
sb.composition_offset(),
"track {} sample {i} composition offset",
a.track_id()
);
assert_eq!(
sa.flags.is_sync,
sb.flags.is_sync,
"track {} sample {i} sync flag",
a.track_id()
);
}
}
let vid = &media.tracks[0].samples;
assert_eq!(vid.len(), 75);
for s in vid {
let nals = transmux::iter_length_prefixed_nals(&s.data).expect("length-prefixed NALs");
let total: usize = nals.iter().map(|n| n.len() + 4).sum();
assert_eq!(
total,
s.data.len(),
"video sample is well-formed length-prefixed"
);
}
assert_eq!(media.tracks[1].samples.len(), 131);
}
#[test]
fn cross_hub_flv_to_cmaf() {
let mut demux = FlvDemux::new();
let media = demux.unpackage(FLV).expect("demux av.flv");
let flv_nals: Vec<Bytes> = media.tracks[0]
.samples
.iter()
.map(|s| s.data.clone())
.collect();
let mut cmaf = CmafMux::new(1);
let seg = cmaf.package(&media).expect("CMAF package");
let moov = find_top_box(&seg, b"moov").expect("moov in CMAF");
let movie = MovieBox::parse(moov).expect("parse moov");
assert_eq!(movie.tracks.len(), 2, "CMAF moov has 2 tracks");
let mut saw_avc1 = false;
let mut saw_mp4a = false;
for trak in &movie.tracks {
let stbl = trak
.mdia
.as_ref()
.and_then(|m| m.minf.as_ref())
.and_then(|m| m.stbl.as_ref())
.expect("stbl");
let stsd = stbl
.children
.iter()
.find_map(|c| match c {
StblChild::Stsd(s) => Some(s),
_ => None,
})
.expect("stsd");
match stsd.entries.first().expect("entry") {
SampleEntryVariant::Avc1(avc1) => {
saw_avc1 = true;
assert!(!avc1.config.config.sps.is_empty(), "avcC has SPS");
}
SampleEntryVariant::Mp4a(_) => saw_mp4a = true,
_ => {}
}
}
assert!(saw_avc1, "CMAF must carry avc1/avcC");
assert!(saw_mp4a, "CMAF must carry mp4a");
let mdat = find_top_box(&seg, b"mdat").expect("mdat in CMAF");
let mdat_body = &mdat[8..];
assert!(
mdat_body.starts_with(&flv_nals[0]),
"first FLV video sample NAL bytes appear verbatim in the CMAF mdat"
);
}
#[test]
fn streaming_payloads_are_flv_shaped_not_ts_shaped() {
let mut demux = FlvDemux::new();
let media = demux.unpackage(FLV).expect("demux av.flv");
let headers = transmux::flv_sequence_header_payloads(&media).expect("build sequence headers");
assert_eq!(headers.len(), 2, "one video + one audio sequence header");
let vid_hdr = headers
.iter()
.find(|p| p.kind == transmux::FlvPayloadKind::Video)
.expect("video sequence header present");
let aud_hdr = headers
.iter()
.find(|p| p.kind == transmux::FlvPayloadKind::Audio)
.expect("audio sequence header present");
assert_eq!(
vid_hdr.body[0], 0x17,
"video seq header FrameType/CodecID byte"
);
assert_ne!(vid_hdr.body[0], 0x47, "must not be a TS sync byte");
assert_eq!(
vid_hdr.body[1], 0,
"video seq header AVCPacketType = sequence header"
);
let CodecConfig::Avc { config, .. } = media.tracks[0].config() else {
panic!("track 0 must be AVC");
};
let mut avcc = alloc_body(config.config.serialized_len());
let n = config.config.serialize_into(&mut avcc).unwrap();
avcc.truncate(n);
assert_eq!(
&vid_hdr.body[5..],
&avcc[..],
"avcC bytes verbatim in the sequence-header payload"
);
assert_eq!(
aud_hdr.body[0], 0xAE,
"audio seq header AudioTagHeader byte (mono AAC)"
);
assert_eq!(
aud_hdr.body[1], 0,
"audio seq header AACPacketType = sequence header"
);
let frames = transmux::flv_frame_payloads(&media).expect("build frame payloads");
assert_eq!(frames.len(), 75 + 131, "one payload per demuxed sample");
let video_frames: Vec<_> = frames
.iter()
.filter(|p| p.kind == transmux::FlvPayloadKind::Video)
.collect();
let audio_frames: Vec<_> = frames
.iter()
.filter(|p| p.kind == transmux::FlvPayloadKind::Audio)
.collect();
assert_eq!(video_frames.len(), 75, "video frame payload count");
assert_eq!(audio_frames.len(), 131, "audio frame payload count");
for p in &video_frames {
assert_ne!(
p.body[0], 0x47,
"video frame payload must not look like a TS packet"
);
assert_eq!(p.body[0] & 0x0F, 0x07, "video frame CodecID nibble = AVC");
assert!(
matches!(p.body[0] >> 4, 1 | 2),
"video frame FrameType nibble = keyframe or inter"
);
assert_eq!(p.body[1], 1, "video frame AVCPacketType = NALU");
}
for p in &audio_frames {
assert_ne!(
p.body[0], 0x47,
"audio frame payload must not look like a TS packet"
);
assert_eq!(p.body[0] >> 4, 0x0A, "audio frame SoundFormat nibble = AAC");
assert_eq!(p.body[1], 1, "audio frame AACPacketType = raw AU");
}
let first_sample_dts = media.tracks[0].samples[0].dts.expect("video dts");
assert_eq!(
video_frames[0].timestamp_ms as i64, first_sample_dts,
"first video payload ms timestamp == sample absolute dts"
);
}
fn alloc_body(len: usize) -> Vec<u8> {
vec![0u8; len]
}