diff --git a/src/mux/ebml.rs b/src/mux/ebml.rs index 5e7798f..ca2661a 100644 --- a/src/mux/ebml.rs +++ b/src/mux/ebml.rs @@ -321,6 +321,12 @@ pub const EBML_DOC_TYPE_READ_VERSION: u32 = 0x4285; // Segment pub const SEGMENT: u32 = 0x1853_8067; +// SeekHead +pub const SEEK_HEAD: u32 = 0x114D_9B74; +pub const SEEK: u32 = 0x4DBB; +pub const SEEK_ID: u32 = 0x53AB; +pub const SEEK_POSITION: u32 = 0x53AC; + // Segment Info pub const INFO: u32 = 0x1549_A966; pub const TIMESTAMP_SCALE: u32 = 0x2A_D7B1; diff --git a/src/mux/m2ts.rs b/src/mux/m2ts.rs index 48fa3a8..bea1925 100644 --- a/src/mux/m2ts.rs +++ b/src/mux/m2ts.rs @@ -239,7 +239,9 @@ impl crate::pes::Stream for M2tsStream { fn write(&mut self, frame: &crate::pes::PesFrame) -> io::Result<()> { match &mut self.mode { - Mode::Write { muxer } => muxer.write_frame(frame.track, frame.pts, &frame.data), + Mode::Write { muxer } => { + muxer.write_frame(frame.track, frame.pts, frame.keyframe, &frame.data) + } Mode::Read { .. } => Err(crate::error::Error::StreamReadOnly.into()), } } @@ -283,3 +285,118 @@ impl crate::pes::Stream for M2tsStream { true } } + +#[cfg(test)] +mod tests { + use super::*; + use crate::disc::{ + Codec, ColorSpace, ContentFormat, DiscTitle, FrameRate, HdrFormat, Resolution, + Stream as DiscStream, VideoStream, + }; + use crate::pes::{PesFrame, Stream as PesStreamTrait}; + + const VIDEO_PID: u16 = 0x1011; + + fn make_title() -> DiscTitle { + DiscTitle { + playlist: String::new(), + playlist_id: 0, + duration_secs: 0.0, + size_bytes: 0, + clips: Vec::new(), + streams: vec![DiscStream::Video(VideoStream { + pid: VIDEO_PID, + codec: Codec::Hevc, + resolution: Resolution::R1080p, + frame_rate: FrameRate::F24, + hdr: HdrFormat::Sdr, + color_space: ColorSpace::Bt709, + secondary: false, + label: String::new(), + })], + chapters: Vec::new(), + extents: Vec::new(), + content_format: ContentFormat::BdTs, + codec_privates: vec![Some({ + // Minimal hvcC with one VPS-like array entry. + let marker: &[u8] = &[0x40, 0x01, 0x0C, 0x01]; + let mut hvcc = vec![0u8; 22]; + hvcc.push(1); // numArrays + hvcc.push(32); + hvcc.extend_from_slice(&1u16.to_be_bytes()); // numNalus + hvcc.extend_from_slice(&(marker.len() as u16).to_be_bytes()); + hvcc.extend_from_slice(marker); + hvcc + })], + } + } + + fn fake_idr_pes_data() -> Vec { + // 4-byte length prefix + NAL: type 19 (IDR_W_RADL). + let mut nal = vec![(19u8 << 1) & 0x7E, 0x01]; + for i in 0..200 { + nal.push((i & 0xFF) as u8); + } + let mut out = Vec::with_capacity(4 + nal.len()); + out.extend_from_slice(&(nal.len() as u32).to_be_bytes()); + out.extend_from_slice(&nal); + out + } + + /// Writer wrapper that shares an Arc>> so the test can + /// inspect the bytes after the muxer drops. + struct SharedSink(std::sync::Arc>>); + impl Write for SharedSink { + fn write(&mut self, b: &[u8]) -> io::Result { + self.0.lock().unwrap().extend_from_slice(b); + Ok(b.len()) + } + fn flush(&mut self) -> io::Result<()> { + Ok(()) + } + } + + #[test] + fn m2ts_stream_forwards_keyframe_to_rai() { + let title = make_title(); + let shared = std::sync::Arc::new(std::sync::Mutex::new(Vec::::new())); + let sink = SharedSink(shared.clone()); + let mut stream = M2tsStream::create(sink, &title).unwrap(); + let frame = PesFrame { + track: 0, + pts: 0, + keyframe: true, + data: fake_idr_pes_data(), + }; + stream.write(&frame).unwrap(); + stream.finish().unwrap(); + drop(stream); + + let buf = shared.lock().unwrap().clone(); + + // Skip FMKV metadata header via meta::read_header. + let mut cursor = std::io::Cursor::new(&buf); + let _meta = super::meta::read_header(&mut cursor) + .unwrap() + .expect("FMKV header present"); + let header_end = cursor.position() as usize; + let ts_bytes = &buf[header_end..]; + + // Find first PUSI packet on VIDEO_PID; verify RAI in AF flags. + let pkt = ts_bytes + .chunks(192) + .find(|p| { + let h = &p[4..]; + let pid = (((h[1] & 0x1F) as u16) << 8) | h[2] as u16; + pid == VIDEO_PID && (h[1] & 0x40) != 0 + }) + .expect("video PUSI packet present"); + let h = &pkt[4..]; + let afc = (h[3] >> 4) & 0x03; + assert!(afc & 0b10 != 0, "AF must be present"); + let af_len = h[4] as usize; + assert!(af_len >= 1, "AF length must include flags byte"); + let flags = h[5]; + assert_eq!(flags & 0x40, 0x40, "RAI bit set"); + } +} diff --git a/src/mux/m2ts_mux/mod.rs b/src/mux/m2ts_mux/mod.rs index 2623bee..4db2157 100644 --- a/src/mux/m2ts_mux/mod.rs +++ b/src/mux/m2ts_mux/mod.rs @@ -161,16 +161,21 @@ impl M2tsMux { } /// Write one video PES frame. `data` is either length-prefixed - /// NALUs (MKV-style) or already Annex B; both are accepted. - pub fn write_video(&mut self, pts_ns: i64, data: &[u8]) -> io::Result<()> { + /// NALUs (MKV-style) or already Annex B; both are accepted. `keyframe` + /// drives the random_access_indicator bit on the first packet of this + /// PES (and gates codec_private NAL prepending — those only attach to + /// the first keyframe). + pub fn write_video(&mut self, pts_ns: i64, keyframe: bool, data: &[u8]) -> io::Result<()> { let pts_90k = self.base_relative_pts(pts_ns); // PCR comes "before" the PTS it timestamps; clamp at 0 for the // first frame so we don't underflow. let pcr = pts_90k.saturating_sub(PCR_LEAD_90KHZ); - // Annex-B-ify the frame and prepend VPS/SPS/PPS once. + // Annex-B-ify the frame and prepend VPS/SPS/PPS once, on the + // FIRST keyframe (not first frame — non-key frames before the + // first keyframe can't carry params usefully). let mut es = Vec::with_capacity(data.len() + 64); - if !self.params_written { + if keyframe && !self.params_written { if let Some(cp) = &self.video_codec_private { let payload = hvcc_payload(cp); if !payload.is_empty() { @@ -184,7 +189,7 @@ impl M2tsMux { es.extend_from_slice(&annex_b); let pes = build_video_pes(pts_90k, &es); - self.write_pes(PID_VIDEO, &pes, Some(pcr)) + self.write_pes(PID_VIDEO, &pes, Some(pcr), keyframe) } /// Write one audio PES frame. Returns `Ok(())` and silently drops @@ -197,7 +202,7 @@ impl M2tsMux { } let pts_90k = self.base_relative_pts(pts_ns); let pes = build_audio_pes(pts_90k, data); - self.write_pes(PID_AUDIO, &pes, None) + self.write_pes(PID_AUDIO, &pes, None, false) } /// Drain the underlying writer. No TS-level trailer is mandatory — @@ -234,7 +239,13 @@ impl M2tsMux { /// The fit-the-tail logic on the last packet of the PES uses /// stuffing rather than a separate small packet, which is the /// standard MPEG-TS convention. - fn write_pes(&mut self, pid: u16, pes: &[u8], pcr: Option) -> io::Result<()> { + fn write_pes( + &mut self, + pid: u16, + pes: &[u8], + pcr: Option, + is_keyframe_video: bool, + ) -> io::Result<()> { // PSI cadence is enforced per TS packet — interleave a fresh // PAT+PMT into the packet stream every PSI_INTERVAL_PACKETS so // long single-PES emissions (e.g. one 60 KB video frame) don't @@ -250,11 +261,21 @@ impl M2tsMux { && (self.packets_written == 0 || self.video_packets_since_pcr >= PCR_INTERVAL_PACKETS); - let af_body: Vec = if attach_pcr { + // RAI rides only the FIRST packet of a keyframe video PES. + let attach_rai = first && is_keyframe_video && pid == PID_VIDEO; + + let mut af_body: Vec = if attach_pcr { build_pcr_adaptation(pcr.unwrap_or(0)) } else { Vec::new() }; + if attach_rai { + if af_body.is_empty() { + af_body.push(0x40); // flags: RAI only + } else { + af_body[0] |= 0x40; // OR RAI into existing PCR flags + } + } let remaining = pes.len() - offset; // Capacity for payload given AF body and 1-byte AF length. @@ -510,7 +531,7 @@ fn build_pcr_adaptation(pcr_90k: u64) -> Vec { // elementary_stream_priority(1) | PCR_flag(1) | OPCR_flag(1) | // splicing_point_flag(1) | transport_private_data_flag(1) | // adaptation_field_extension_flag(1) | PCR(48b). - let mut af = vec![0x50]; // PCR_flag=1, random_access_indicator=1 + let mut af = vec![0x10]; // PCR_flag=1; RAI is OR'd in by the caller when applicable let pcr_base = pcr_90k & 0x1_FFFF_FFFF; // 33-bit let pcr_ext: u16 = 0; // 9-bit, we keep it zero (no sub-tick precision) // Encode PCR: 33b base | 6b reserved | 9b extension = 48b @@ -598,7 +619,7 @@ mod tests { let mut frame = Vec::new(); frame.extend_from_slice(&4u32.to_be_bytes()); frame.extend_from_slice(&[0x40, 0x01, 0x0C, 0x01]); - mux.write_video(0, &frame).unwrap(); + mux.write_video(0, true, &frame).unwrap(); mux.finish().unwrap(); drop(mux); @@ -619,7 +640,7 @@ mod tests { let mut frame = Vec::new(); frame.extend_from_slice(&3u32.to_be_bytes()); frame.extend_from_slice(&[0x40, 0x01, 0x0C]); - mux.write_video(0, &frame).unwrap(); + mux.write_video(0, true, &frame).unwrap(); mux.write_audio(20_000_000, &[0x0B, 0x77, 0x12, 0x34]) .unwrap(); mux.finish().unwrap(); @@ -641,7 +662,7 @@ mod tests { let mut frame = Vec::new(); frame.extend_from_slice(&(big.len() as u32).to_be_bytes()); frame.extend_from_slice(&big); - mux.write_video(0, &frame).unwrap(); + mux.write_video(0, true, &frame).unwrap(); mux.finish().unwrap(); drop(mux); @@ -663,7 +684,7 @@ mod tests { let mut frame = Vec::new(); frame.extend_from_slice(&3u32.to_be_bytes()); frame.extend_from_slice(&[0xAA, 0xBB, 0xCC]); - mux.write_video(pts, &frame).unwrap(); + mux.write_video(pts, pts == 0, &frame).unwrap(); } mux.finish().unwrap(); drop(mux); @@ -678,4 +699,130 @@ mod tests { assert_eq!(w[1], (w[0] + 1) & 0x0F); } } + + /// Return the adaptation field body (length byte stripped) for one + /// 188-byte TS packet, or None when the packet has no AF. + fn af_body(packet: &[u8]) -> Option> { + let afc = (packet[3] >> 4) & 0x03; + if afc & 0b10 == 0 { + return None; + } + let af_len = packet[4] as usize; + if af_len == 0 { + return Some(Vec::new()); + } + Some(packet[5..5 + af_len].to_vec()) + } + + #[test] + fn rai_set_on_keyframe_pes_packet() { + let mut sink: Vec = Vec::new(); + let mut mux = M2tsMux::new(&mut sink); + let mut frame = Vec::new(); + frame.extend_from_slice(&4u32.to_be_bytes()); + frame.extend_from_slice(&[0x40, 0x01, 0x0C, 0x01]); + mux.write_video(0, true, &frame).unwrap(); + mux.finish().unwrap(); + drop(mux); + + // Find the first PUSI packet on PID_VIDEO. + let pkt = sink + .chunks(188) + .find(|p| u16::from_be_bytes([p[1] & 0x1F, p[2]]) == PID_VIDEO && (p[1] & 0x40) != 0) + .expect("video PUSI packet exists"); + let af = af_body(pkt).expect("AF present on first packet of keyframe video PES"); + assert!(!af.is_empty(), "AF flags byte present"); + assert_eq!(af[0] & 0x40, 0x40, "RAI bit set"); + } + + #[test] + fn pcr_packet_without_keyframe_has_rai_clear() { + let mut sink: Vec = Vec::new(); + let mut mux = M2tsMux::new(&mut sink); + // First frame is the keyframe (gates codec_private; also gets PCR). + let mut frame0 = Vec::new(); + frame0.extend_from_slice(&4u32.to_be_bytes()); + frame0.extend_from_slice(&[0x40, 0x01, 0x0C, 0x01]); + mux.write_video(0, true, &frame0).unwrap(); + // Push enough non-key video frames to cross PCR_INTERVAL_PACKETS + // video packets so a later PCR-bearing packet exists. + // Each frame is ~50 KB → ~270 packets, well over 40. + let big: Vec = (0..(50 * 1024)).map(|i| (i & 0xff) as u8).collect(); + for i in 1..3 { + let mut frame = Vec::new(); + frame.extend_from_slice(&(big.len() as u32).to_be_bytes()); + frame.extend_from_slice(&big); + mux.write_video((i as i64) * 40_000_000, false, &frame) + .unwrap(); + } + mux.finish().unwrap(); + drop(mux); + + // The first PUSI video packet carries PCR + RAI (keyframe). + // A later video PUSI packet with AF + PCR but NOT keyframe must + // have RAI clear. + let video_pusi: Vec<&[u8]> = sink + .chunks(188) + .filter(|p| u16::from_be_bytes([p[1] & 0x1F, p[2]]) == PID_VIDEO && (p[1] & 0x40) != 0) + .collect(); + assert!( + video_pusi.len() >= 2, + "expected ≥2 video PES starts, got {}", + video_pusi.len() + ); + // Find a later one with AF that carries PCR (flags & 0x10 set). + let later_pcr = video_pusi + .iter() + .skip(1) + .find_map(|p| { + let af = af_body(p)?; + if !af.is_empty() && (af[0] & 0x10) != 0 { + Some(af) + } else { + None + } + }) + .expect("later PCR-bearing PUSI exists"); + assert_eq!( + later_pcr[0] & 0x40, + 0, + "RAI must be clear on non-keyframe PCR packet" + ); + } + + #[test] + fn keyframe_video_with_pcr_combines_flags() { + // PCR attaches only when video_packets_since_pcr >= + // PCR_INTERVAL_PACKETS (40). The very first video packet emits a + // PAT+PMT first, so packets_written != 0 and attach_pcr is false on + // frame 0. We push: keyframe (no PCR) → many non-key (drives the + // PCR counter past the interval) → second keyframe (PCR + RAI). + let mut sink: Vec = Vec::new(); + let mut mux = M2tsMux::new(&mut sink); + let mut small = Vec::new(); + small.extend_from_slice(&4u32.to_be_bytes()); + small.extend_from_slice(&[0x40, 0x01, 0x0C, 0x01]); + mux.write_video(0, true, &small).unwrap(); + // ~50 KB ≈ 270 packets — well over PCR_INTERVAL_PACKETS. + let big: Vec = (0..(50 * 1024)).map(|i| (i & 0xff) as u8).collect(); + let mut big_frame = Vec::new(); + big_frame.extend_from_slice(&(big.len() as u32).to_be_bytes()); + big_frame.extend_from_slice(&big); + mux.write_video(40_000_000, false, &big_frame).unwrap(); + // Now a second keyframe — must combine RAI (keyframe) and PCR + // (counter exceeded). + mux.write_video(80_000_000, true, &small).unwrap(); + mux.finish().unwrap(); + drop(mux); + + // Collect video PUSI packets and find the third (second keyframe). + let video_pusi: Vec<&[u8]> = sink + .chunks(188) + .filter(|p| u16::from_be_bytes([p[1] & 0x1F, p[2]]) == PID_VIDEO && (p[1] & 0x40) != 0) + .collect(); + assert!(video_pusi.len() >= 3, "three video PES starts expected"); + let af = af_body(video_pusi[2]).expect("AF present"); + assert!(!af.is_empty(), "AF flags byte present"); + assert_eq!(af[0], 0x50, "flags == RAI | PCR"); + } } diff --git a/src/mux/mkv.rs b/src/mux/mkv.rs index ddcc24d..24873a7 100644 --- a/src/mux/mkv.rs +++ b/src/mux/mkv.rs @@ -159,6 +159,12 @@ struct CuePoint { cluster_pos: u64, // relative to Segment start } +/// SeekHead entry that needs its 8-byte SeekPosition back-patched after Cues are written. +struct SeekPositionFixup { + target_id: u32, + value_offset: u64, // absolute file offset of the 8-byte SeekPosition value +} + /// MKV muxer. Call write_frame() for each frame, then finish() at the end. pub struct MkvMuxer { writer: W, @@ -170,6 +176,10 @@ pub struct MkvMuxer { base_pts_ms: Option, cues: Vec, frame_count: u64, + seek_fixups: Vec, + info_offset: u64, + tracks_offset: u64, + chapters_offset: Option, } /// New cluster every 5 seconds. @@ -200,7 +210,34 @@ impl MkvMuxer { ebml::write_unknown_size(&mut writer)?; let segment_start = writer.stream_position()?; + // SeekHead with fixed-width SeekPosition placeholders. Order: Info, Tracks, [Chapters], Cues. + let mut seek_fixups: Vec = Vec::new(); + let seekhead_pos = ebml::start_master(&mut writer, ebml::SEEK_HEAD)?; + let mut targets: Vec = vec![ebml::INFO, ebml::TRACKS]; + if !chapters.is_empty() { + targets.push(ebml::CHAPTERS); + } + targets.push(ebml::CUES); + let seek_id_be = (ebml::SEEK as u16).to_be_bytes(); + let seek_inner_id_be = (ebml::SEEK_ID as u16).to_be_bytes(); + let seek_pos_id_be = (ebml::SEEK_POSITION as u16).to_be_bytes(); + for target_id in &targets { + writer.write_all(&[seek_id_be[0], seek_id_be[1], 0x92])?; + writer.write_all(&[seek_inner_id_be[0], seek_inner_id_be[1], 0x84])?; + writer.write_all(&target_id.to_be_bytes())?; + writer.write_all(&[seek_pos_id_be[0], seek_pos_id_be[1], 0x88])?; + let value_offset = writer.stream_position()?; + writer.write_all(&[0u8; 8])?; + seek_fixups.push(SeekPositionFixup { + target_id: *target_id, + value_offset, + }); + } + ebml::end_master(&mut writer, seekhead_pos)?; + // Info + let info_start = writer.stream_position()?; + let info_offset = info_start - segment_start; let info_pos = ebml::start_master(&mut writer, ebml::INFO)?; ebml::write_uint(&mut writer, ebml::TIMESTAMP_SCALE, 1_000_000)?; // 1ms precision if duration_secs > 0.0 { @@ -215,6 +252,8 @@ impl MkvMuxer { ebml::end_master(&mut writer, info_pos)?; // Tracks + let tracks_start = writer.stream_position()?; + let tracks_offset = tracks_start - segment_start; let tracks_pos = ebml::start_master(&mut writer, ebml::TRACKS)?; for (i, track) in tracks.iter().enumerate() { let entry_pos = ebml::start_master(&mut writer, ebml::TRACK_ENTRY)?; @@ -302,7 +341,10 @@ impl MkvMuxer { ebml::end_master(&mut writer, tracks_pos)?; // Chapters + let mut chapters_offset: Option = None; if !chapters.is_empty() { + let chapters_start = writer.stream_position()?; + chapters_offset = Some(chapters_start - segment_start); let chapters_pos = ebml::start_master(&mut writer, ebml::CHAPTERS)?; let edition_pos = ebml::start_master(&mut writer, ebml::EDITION_ENTRY)?; for (i, ch) in chapters.iter().enumerate() { @@ -330,6 +372,10 @@ impl MkvMuxer { base_pts_ms: None, cues: Vec::new(), frame_count: 0, + seek_fixups, + info_offset, + tracks_offset, + chapters_offset, }) } @@ -345,22 +391,22 @@ impl MkvMuxer { let base = *self.base_pts_ms.get_or_insert(raw_ms); let pts_ms = raw_ms - base; - // Start new cluster if needed - if !self.cluster_open || (pts_ms - self.cluster_ts_ms) >= CLUSTER_DURATION_MS { - if self.cluster_open { - // Close current cluster (it's a master with unknown size — we use known size) - // Actually, for streaming we keep clusters open-ended. Just start a new one. + // Cluster boundaries must coincide with a video keyframe so every + // Cues entry resolves to a seekable IDR at the cluster start. + let is_video_key = keyframe && track_idx == 0; + let needs_new_cluster = !self.cluster_open + || (is_video_key && (pts_ms - self.cluster_ts_ms) >= CLUSTER_DURATION_MS); + + if needs_new_cluster { + if !is_video_key { + return Ok(()); } self.start_cluster(pts_ms)?; - - // Add cue point at cluster start for keyframes (video track 0) - if keyframe && track_idx == 0 { - self.cues.push(CuePoint { - timestamp_ms: pts_ms, - track: track_idx + 1, - cluster_pos: self.cluster_pos - self.segment_start, - }); - } + self.cues.push(CuePoint { + timestamp_ms: pts_ms, + track: track_idx + 1, + cluster_pos: self.cluster_pos - self.segment_start, + }); } // Write SimpleBlock @@ -377,6 +423,8 @@ impl MkvMuxer { self.end_cluster()?; // Write Cues + let cues_start = self.writer.stream_position()?; + let cues_offset = cues_start - self.segment_start; if !self.cues.is_empty() { let cues_pos = ebml::start_master(&mut self.writer, ebml::CUES)?; for cue in &self.cues { @@ -395,6 +443,21 @@ impl MkvMuxer { ebml::end_master(&mut self.writer, cues_pos)?; } + // Back-patch SeekHead SeekPosition values now that all element offsets are known. + for fixup in &self.seek_fixups { + let offset = match fixup.target_id { + ebml::INFO => self.info_offset, + ebml::TRACKS => self.tracks_offset, + ebml::CHAPTERS => self.chapters_offset.unwrap_or(0), + ebml::CUES => cues_offset, + _ => 0, + }; + self.writer + .seek(std::io::SeekFrom::Start(fixup.value_offset))?; + self.writer.write_all(&offset.to_be_bytes())?; + } + self.writer.seek(std::io::SeekFrom::End(0))?; + self.writer.flush()?; Ok(()) } @@ -812,4 +875,422 @@ mod tests { "FlagForced element should not be present for non-forced subtitle" ); } + + // ============================================================ + // Seekability tests: SeekHead, keyframe-aligned clusters, Cues + // ============================================================ + + use std::sync::{Arc, Mutex}; + + /// Writer that lets the test inspect the buffer after `finish()` consumes the muxer. + struct SharedWriter(Arc>>>); + impl Write for SharedWriter { + fn write(&mut self, buf: &[u8]) -> io::Result { + self.0.lock().unwrap().write(buf) + } + fn flush(&mut self) -> io::Result<()> { + self.0.lock().unwrap().flush() + } + } + impl Seek for SharedWriter { + fn seek(&mut self, pos: io::SeekFrom) -> io::Result { + self.0.lock().unwrap().seek(pos) + } + } + + /// Build interleaved frames at 24 fps video (IDR every gop_secs) + 48 kHz audio (1024 samples per frame). + fn frames_for(duration_secs: f64, gop_secs: f64) -> Vec<(usize, i64, bool, Vec)> { + let video_interval_ns: i64 = 1_000_000_000 / 24; + let audio_interval_ns: i64 = (1024i64 * 1_000_000_000) / 48_000; + let gop_frames = (gop_secs * 24.0).round() as i64; + + let mut out: Vec<(usize, i64, bool, Vec)> = Vec::new(); + let total_ns = (duration_secs * 1_000_000_000.0) as i64; + + let mut vi: i64 = 0; + loop { + let pts = vi * video_interval_ns; + if pts >= total_ns { + break; + } + let keyframe = vi % gop_frames == 0; + out.push((0, pts, keyframe, vec![0xAB; 64])); + vi += 1; + } + + let mut ai: i64 = 0; + loop { + let pts = ai * audio_interval_ns; + if pts >= total_ns { + break; + } + out.push((1, pts, true, vec![0xCD; 32])); + ai += 1; + } + + out.sort_by_key(|f| f.1); + out + } + + /// Mux frames through a SharedWriter and return the final buffer. + fn mux_to_bytes( + tracks: &[MkvTrack], + chapters: &[Chapter], + frames: &[(usize, i64, bool, Vec)], + ) -> (Vec, u64) { + let shared = Arc::new(Mutex::new(Cursor::new(Vec::new()))); + let writer = SharedWriter(shared.clone()); + let mut muxer = MkvMuxer::new(writer, tracks, None, 0.0, chapters).unwrap(); + for (t, pts, kf, data) in frames { + muxer.write_frame(*t, *pts, *kf, data).unwrap(); + } + let frame_count = muxer.frame_count; + muxer.finish().unwrap(); + let data = shared.lock().unwrap().clone().into_inner(); + (data, frame_count) + } + + /// Find the Segment header in the buffer and return (segment_id_pos, segment_start_pos). + /// segment_start = position immediately after Segment's id + size bytes. + fn locate_segment(data: &[u8]) -> (usize, usize) { + let segment_id_pos = find_id(data, ebml::SEGMENT).expect("segment id not found"); + // Segment is written via write_id + write_unknown_size: 4 byte id + 8 byte size + (segment_id_pos, segment_id_pos + 4 + 8) + } + + /// Walk Segment's top-level children. Returns Vec<(id, data_start_offset, data_size)> + /// where data_start_offset is absolute file offset and data_size is the element body size. + fn segment_children(data: &[u8]) -> Vec<(u32, usize, u64)> { + let (_, seg_start) = locate_segment(data); + let mut out = Vec::new(); + let mut cursor = Cursor::new(&data[seg_start..]); + while (cursor.position() as usize) < data.len() - seg_start { + let pos_before = cursor.position(); + let (id, size, hdr_len) = match ebml::read_element_header(&mut cursor) { + Ok(v) => v, + Err(_) => break, + }; + let data_abs = seg_start + pos_before as usize + hdr_len; + out.push((id, data_abs, size)); + // Skip the body to advance to the next element. + cursor + .seek(io::SeekFrom::Current(size as i64)) + .expect("seek past element body"); + } + out + } + + /// Find every Cluster: returns Vec<(cluster_data_start_abs, cluster_data_size, cluster_timestamp_ms)>. + fn find_clusters(data: &[u8]) -> Vec<(usize, u64, u64)> { + let mut out = Vec::new(); + for (id, body_start, body_size) in segment_children(data) { + if id == ebml::CLUSTER { + let mut cursor = Cursor::new(&data[body_start..body_start + body_size as usize]); + let (tid, tsize, _) = ebml::read_element_header(&mut cursor).unwrap(); + assert_eq!( + tid, + ebml::CLUSTER_TIMESTAMP, + "cluster must start with timestamp" + ); + let ts = ebml::read_uint_val(&mut cursor, tsize as usize).unwrap(); + out.push((body_start, body_size, ts)); + } + } + out + } + + /// Parse the first SimpleBlock that appears in a cluster body slice. + /// Returns (track_num, flags_byte). track_num decoded from VINT. + fn first_simple_block(cluster_body: &[u8]) -> (u64, u8) { + let mut cursor = Cursor::new(cluster_body); + loop { + let (id, size, _) = ebml::read_element_header(&mut cursor).unwrap(); + if id == ebml::SIMPLE_BLOCK { + let body_start = cursor.position() as usize; + // Decode track VINT. + let b0 = cluster_body[body_start]; + let (track_num, vint_len) = if b0 & 0x80 != 0 { + ((b0 & 0x7F) as u64, 1usize) + } else if b0 & 0x40 != 0 { + let b1 = cluster_body[body_start + 1]; + ((((b0 & 0x3F) as u64) << 8) | b1 as u64, 2) + } else { + panic!("unsupported track vint width"); + }; + let flags = cluster_body[body_start + vint_len + 2]; + return (track_num, flags); + } + // Skip non-SimpleBlock child. + cursor.seek(io::SeekFrom::Current(size as i64)).unwrap(); + } + } + + /// Parse the Cues element body into Vec<(cue_time, cue_track, cue_cluster_position)>. + fn parse_cues(data: &[u8]) -> Vec<(u64, u64, u64)> { + let mut out = Vec::new(); + let (cues_id, cues_body_start, cues_body_size) = segment_children(data) + .into_iter() + .find(|(id, _, _)| *id == ebml::CUES) + .expect("cues element not found"); + assert_eq!(cues_id, ebml::CUES); + let cues_body = &data[cues_body_start..cues_body_start + cues_body_size as usize]; + let mut cursor = Cursor::new(cues_body); + while (cursor.position() as usize) < cues_body.len() { + let (id, size, _) = ebml::read_element_header(&mut cursor).unwrap(); + assert_eq!(id, ebml::CUE_POINT); + let cp_end = cursor.position() + size; + let mut cue_time = 0u64; + let mut cue_track = 0u64; + let mut cue_pos = 0u64; + while cursor.position() < cp_end { + let (sid, ssize, _) = ebml::read_element_header(&mut cursor).unwrap(); + match sid { + ebml::CUE_TIME => { + cue_time = ebml::read_uint_val(&mut cursor, ssize as usize).unwrap(); + } + ebml::CUE_TRACK_POSITIONS => { + let ctp_end = cursor.position() + ssize; + while cursor.position() < ctp_end { + let (iid, isize_, _) = ebml::read_element_header(&mut cursor).unwrap(); + match iid { + ebml::CUE_TRACK => { + cue_track = + ebml::read_uint_val(&mut cursor, isize_ as usize).unwrap(); + } + ebml::CUE_CLUSTER_POSITION => { + cue_pos = + ebml::read_uint_val(&mut cursor, isize_ as usize).unwrap(); + } + _ => { + cursor.seek(io::SeekFrom::Current(isize_ as i64)).unwrap(); + } + } + } + } + _ => { + cursor.seek(io::SeekFrom::Current(ssize as i64)).unwrap(); + } + } + } + out.push((cue_time, cue_track, cue_pos)); + } + out + } + + /// Parse the SeekHead body into Vec<(seek_id, seek_position)>. + fn parse_seekhead(data: &[u8]) -> Vec<(u32, u64)> { + let mut out = Vec::new(); + let (sh_id, sh_body_start, sh_body_size) = segment_children(data) + .into_iter() + .find(|(id, _, _)| *id == ebml::SEEK_HEAD) + .expect("seekhead not found"); + assert_eq!(sh_id, ebml::SEEK_HEAD); + let sh_body = &data[sh_body_start..sh_body_start + sh_body_size as usize]; + let mut cursor = Cursor::new(sh_body); + while (cursor.position() as usize) < sh_body.len() { + let (id, size, _) = ebml::read_element_header(&mut cursor).unwrap(); + assert_eq!(id, ebml::SEEK); + let seek_end = cursor.position() + size; + let mut seek_id_val: u32 = 0; + let mut seek_pos_val: u64 = 0; + while cursor.position() < seek_end { + let (sid, ssize, _) = ebml::read_element_header(&mut cursor).unwrap(); + match sid { + ebml::SEEK_ID => { + let raw = ebml::read_uint_val(&mut cursor, ssize as usize).unwrap(); + seek_id_val = raw as u32; + } + ebml::SEEK_POSITION => { + seek_pos_val = ebml::read_uint_val(&mut cursor, ssize as usize).unwrap(); + } + _ => { + cursor.seek(io::SeekFrom::Current(ssize as i64)).unwrap(); + } + } + } + out.push((seek_id_val, seek_pos_val)); + } + out + } + + #[test] + fn cluster_starts_only_on_video_keyframe() { + let tracks = [make_video_track(), make_audio_track()]; + let frames = frames_for(30.0, 1.0); + let (data, _) = mux_to_bytes(&tracks, &[], &frames); + let clusters = find_clusters(&data); + assert!(!clusters.is_empty(), "expected at least one cluster"); + for (body_start, body_size, _ts) in clusters { + let body = &data[body_start..body_start + body_size as usize]; + // Skip past the CLUSTER_TIMESTAMP element first. + let mut cursor = Cursor::new(body); + let (tid, tsize, _) = ebml::read_element_header(&mut cursor).unwrap(); + assert_eq!(tid, ebml::CLUSTER_TIMESTAMP); + cursor.seek(io::SeekFrom::Current(tsize as i64)).unwrap(); + let after_ts = cursor.position() as usize; + let (track_num, flags) = first_simple_block(&body[after_ts..]); + assert_eq!( + track_num, 1, + "first block in cluster must be track 1 (video)" + ); + assert_eq!( + flags & 0x80, + 0x80, + "first block in cluster must have keyframe flag set, got 0x{:02X}", + flags + ); + } + } + + #[test] + fn cue_count_equals_cluster_count() { + let tracks = [make_video_track(), make_audio_track()]; + let frames = frames_for(30.0, 1.0); + let (data, _) = mux_to_bytes(&tracks, &[], &frames); + let clusters = find_clusters(&data); + let cues = parse_cues(&data); + assert_eq!( + clusters.len(), + cues.len(), + "cluster count {} != cue count {}", + clusters.len(), + cues.len() + ); + // For 30s @ 5s min cluster duration with 1s GOP, expect 6 clusters / 6 cues. + assert_eq!( + clusters.len(), + 6, + "expected 6 clusters for 30s @ 5s cluster duration" + ); + } + + #[test] + fn cue_positions_resolve_to_clusters() { + let tracks = [make_video_track(), make_audio_track()]; + let frames = frames_for(30.0, 1.0); + let (data, _) = mux_to_bytes(&tracks, &[], &frames); + let (_, seg_start) = locate_segment(&data); + let cues = parse_cues(&data); + assert!(!cues.is_empty()); + for (_time, _track, pos) in cues { + let abs = seg_start + pos as usize; + let mut cursor = Cursor::new(&data[abs..]); + let (id, _size, _hdr_len) = ebml::read_element_header(&mut cursor).unwrap(); + assert_eq!( + id, + ebml::CLUSTER, + "cue position 0x{:X} did not resolve to a cluster", + pos + ); + } + } + + #[test] + fn cue_times_match_cluster_timestamps() { + let tracks = [make_video_track(), make_audio_track()]; + let frames = frames_for(30.0, 1.0); + let (data, _) = mux_to_bytes(&tracks, &[], &frames); + let (_, seg_start) = locate_segment(&data); + let cues = parse_cues(&data); + for (time, _track, pos) in cues { + let abs = seg_start + pos as usize; + let mut cursor = Cursor::new(&data[abs..]); + let (id, size, _hdr_len) = ebml::read_element_header(&mut cursor).unwrap(); + assert_eq!(id, ebml::CLUSTER); + let body_start = abs + (cursor.position() as usize); + let body = &data[body_start..body_start + size as usize]; + let mut bc = Cursor::new(body); + let (tid, tsize, _) = ebml::read_element_header(&mut bc).unwrap(); + assert_eq!(tid, ebml::CLUSTER_TIMESTAMP); + let cluster_ts = ebml::read_uint_val(&mut bc, tsize as usize).unwrap(); + assert_eq!( + cluster_ts, time, + "cluster timestamp {} != cue time {}", + cluster_ts, time + ); + } + } + + #[test] + fn seekhead_is_first_child_of_segment() { + let tracks = [make_video_track(), make_audio_track()]; + let (data, _) = mux_to_bytes(&tracks, &[], &frames_for(10.0, 1.0)); + let children = segment_children(&data); + assert!(!children.is_empty()); + assert_eq!( + children[0].0, + ebml::SEEK_HEAD, + "first child of segment must be SeekHead, got id 0x{:X}", + children[0].0 + ); + } + + #[test] + fn seekhead_points_to_real_elements() { + let tracks = [make_video_track(), make_audio_track()]; + let (data, _) = mux_to_bytes(&tracks, &[], &frames_for(10.0, 1.0)); + let (_, seg_start) = locate_segment(&data); + let entries = parse_seekhead(&data); + let required = [ebml::INFO, ebml::TRACKS, ebml::CUES]; + for &want_id in &required { + let entry = entries + .iter() + .find(|(id, _)| *id == want_id) + .unwrap_or_else(|| panic!("seekhead missing entry for id 0x{:X}", want_id)); + let abs = seg_start + entry.1 as usize; + let mut cursor = Cursor::new(&data[abs..]); + let (got_id, _, _) = ebml::read_element_header(&mut cursor).unwrap(); + assert_eq!( + got_id, want_id, + "seekhead entry for 0x{:X} resolves to wrong id 0x{:X}", + want_id, got_id + ); + } + } + + #[test] + fn seekhead_omits_chapters_when_empty() { + let tracks = [make_video_track()]; + let (data, _) = mux_to_bytes(&tracks, &[], &frames_for(5.0, 1.0)); + let entries = parse_seekhead(&data); + assert_eq!( + entries.len(), + 3, + "expected 3 seek entries (Info, Tracks, Cues), got {}", + entries.len() + ); + assert!( + entries.iter().all(|(id, _)| *id != ebml::CHAPTERS), + "seekhead should not contain Chapters entry when chapters are empty" + ); + } + + #[test] + fn pre_first_keyframe_frames_dropped() { + let tracks = [make_video_track()]; + let frames = vec![ + (0usize, 0i64, false, vec![0x11; 16]), + (0usize, 41_000_000i64, true, vec![0x22; 16]), + ]; + let (data, frame_count) = mux_to_bytes(&tracks, &[], &frames); + assert_eq!(frame_count, 1, "muxer.frame_count must equal 1"); + let clusters = find_clusters(&data); + assert_eq!(clusters.len(), 1, "expected exactly one cluster"); + let (body_start, body_size, _ts) = clusters[0]; + let body = &data[body_start..body_start + body_size as usize]; + let mut cursor = Cursor::new(body); + // Skip CLUSTER_TIMESTAMP. + let (tid, tsize, _) = ebml::read_element_header(&mut cursor).unwrap(); + assert_eq!(tid, ebml::CLUSTER_TIMESTAMP); + cursor.seek(io::SeekFrom::Current(tsize as i64)).unwrap(); + let mut sb_count = 0; + while (cursor.position() as usize) < body.len() { + let (id, sz, _) = ebml::read_element_header(&mut cursor).unwrap(); + if id == ebml::SIMPLE_BLOCK { + sb_count += 1; + } + cursor.seek(io::SeekFrom::Current(sz as i64)).unwrap(); + } + assert_eq!(sb_count, 1, "expected exactly one SimpleBlock in output"); + } } diff --git a/src/mux/tsmux.rs b/src/mux/tsmux.rs index a7170be..cbdb801 100644 --- a/src/mux/tsmux.rs +++ b/src/mux/tsmux.rs @@ -42,28 +42,39 @@ impl TsMuxer { /// Write a PES frame as BD-TS packets. /// Video frame data is expected as length-prefixed NALUs (MKV/PES format) /// and is converted to Annex B for transport stream. - pub fn write_frame(&mut self, track: usize, pts_ns: i64, data: &[u8]) -> io::Result<()> { + pub fn write_frame( + &mut self, + track: usize, + pts_ns: i64, + keyframe: bool, + data: &[u8], + ) -> io::Result<()> { if track >= self.pids.len() { return Ok(()); // unknown track, skip } - let base = *self.base_pts_ns.get_or_insert(pts_ns); - let pts_ns = pts_ns - base; - let pid = self.pids[track]; let is_video = (0x1011..=0x101F).contains(&pid); - // For video: convert length-prefixed NALUs to Annex B (start codes) - // On first keyframe, prepend parameter sets from codec_private + // Drop non-key video before any keyframe — decoder has no IDR or + // parameter sets to anchor on. + if is_video && !keyframe && !self.params_written[track] { + return Ok(()); + } + + let base = *self.base_pts_ns.get_or_insert(pts_ns); + let pts_ns = pts_ns - base; + + // For video: convert length-prefixed NALUs to Annex B (start codes). + // Prepend codec_private parameter sets on the FIRST keyframe only. let es_data = if is_video && !data.is_empty() { let mut annex_b = Vec::new(); - // Prepend codec_private parameter sets on first keyframe - if !self.params_written[track] { + if keyframe && !self.params_written[track] { if let Some(ref cp) = self.codec_privates[track] { if let Some(params) = hvcc_to_annex_b(cp) { annex_b.extend_from_slice(¶ms); - self.params_written[track] = true; } } + self.params_written[track] = true; } annex_b.extend_from_slice(&length_prefixed_to_annex_b(data)); annex_b @@ -85,8 +96,24 @@ impl TsMuxer { let mut first = true; while offset < pes_packet.len() { let remaining = pes_packet.len() - offset; - let payload_len = remaining.min(TS_PAYLOAD); - let need_stuffing = payload_len < TS_PAYLOAD; + + // Invariant: TP_extra(4) + TS_header(4) + AF(af_bytes) + payload(payload_len) = 192, + // i.e. af_bytes + payload_len = TS_PAYLOAD (184). + // RAI on first packet of a keyframe video PES requires AF with flags=0x40. + let want_rai = first && keyframe && is_video; + + // Pick payload_len and af_bytes per case. + let (af_bytes, payload_len): (usize, usize) = if want_rai { + // Minimum AF = 2 bytes (length=1, flags=0x40). Payload caps at 182. + let max_payload = TS_PAYLOAD - 2; + let p = remaining.min(max_payload); + (TS_PAYLOAD - p, p) + } else if remaining >= TS_PAYLOAD { + (0, TS_PAYLOAD) // no AF, full payload + } else { + // Stuffing-only AF, payload = remaining. + (TS_PAYLOAD - remaining, remaining) + }; // TP_extra_header (4 bytes — arrival time, set to 0) let tp_extra = [0u8; 4]; @@ -102,38 +129,45 @@ impl TsMuxer { ts_header[1] |= 0x40; // PUSI } ts_header[2] = pid as u8; - ts_header[3] = 0x10 | cc; // no adaptation, has payload + ts_header[3] = if af_bytes > 0 { + 0x30 | cc // AF + payload + } else { + 0x10 | cc // payload only + }; - if need_stuffing { - // Adaptation field for stuffing - let stuff_len = TS_PAYLOAD - payload_len; - ts_header[3] = 0x30 | cc; // adaptation + payload + self.writer.write_all(&tp_extra)?; + self.writer.write_all(&ts_header)?; - self.writer.write_all(&tp_extra)?; - self.writer.write_all(&ts_header)?; - - // Write adaptation field: length byte + flags byte + 0xFF padding - // stuff_len == 1: AF length = 0 (just the length byte, no flags) - // stuff_len >= 2: AF length = stuff_len-1, flags = 0, rest 0xFF + if af_bytes > 0 { static STUFF_FF: [u8; 184] = [0xFF; 184]; - if stuff_len == 1 { - self.writer.write_all(&[0u8])?; // adaptation_field_length = 0 + if want_rai { + // RAI AF: length byte + flags(0x40) + (af_bytes - 2) stuffing. + let af_len_field = (af_bytes - 1) as u8; + self.writer.write_all(&[af_len_field])?; + self.writer.write_all(&[0x40u8])?; + let stuff = af_bytes - 2; + if stuff > 0 { + self.writer.write_all(&STUFF_FF[..stuff])?; + } } else { - self.writer.write_all(&[(stuff_len - 1) as u8])?; // AF length - self.writer.write_all(&[0u8])?; // flags - if stuff_len > 2 { - self.writer.write_all(&STUFF_FF[..stuff_len - 2])?; + // Stuffing-only AF. + // af_bytes == 1: length=0, no flags. + // af_bytes >= 2: length = af_bytes-1, flags=0, rest 0xFF. + if af_bytes == 1 { + self.writer.write_all(&[0u8])?; + } else { + self.writer.write_all(&[(af_bytes - 1) as u8])?; + self.writer.write_all(&[0u8])?; + if af_bytes > 2 { + self.writer.write_all(&STUFF_FF[..af_bytes - 2])?; + } } } - self.writer - .write_all(&pes_packet[offset..offset + payload_len])?; - } else { - self.writer.write_all(&tp_extra)?; - self.writer.write_all(&ts_header)?; - self.writer - .write_all(&pes_packet[offset..offset + payload_len])?; } + self.writer + .write_all(&pes_packet[offset..offset + payload_len])?; + offset += payload_len; first = false; } @@ -257,3 +291,206 @@ fn length_prefixed_to_annex_b(data: &[u8]) -> Vec { } out } + +#[cfg(test)] +mod tests { + use super::*; + + const BD_PACKET_SIZE: usize = 192; + const VIDEO_PID: u16 = 0x1011; + + /// Parsed BD-TS packet (192 bytes total: 4 TP_extra + 4 TS header + 184 body). + struct TsPacket { + pid: u16, + pusi: bool, + #[allow(dead_code)] + cc: u8, + /// Adaptation field body (length byte stripped) when present. + af: Option>, + /// Payload bytes (after AF, if any). + payload: Vec, + } + + /// Walk 192-byte BD-TS packets. + fn parse_bd_ts(buf: &[u8]) -> Vec { + let mut out = Vec::new(); + for chunk in buf.chunks(BD_PACKET_SIZE) { + if chunk.len() != BD_PACKET_SIZE { + break; + } + // Skip TP_extra_header (4 bytes), parse TS header. + let h = &chunk[4..]; + assert_eq!(h[0], 0x47, "bad sync byte"); + let pusi = (h[1] & 0x40) != 0; + let pid = (((h[1] & 0x1F) as u16) << 8) | h[2] as u16; + let afc = (h[3] >> 4) & 0x03; + let cc = h[3] & 0x0F; + let body = &h[4..]; // 184 bytes + + let (af, payload) = match afc { + 0b01 => (None, body.to_vec()), + 0b11 => { + let af_len = body[0] as usize; + let af_body = body[1..1 + af_len].to_vec(); + let payload = body[1 + af_len..].to_vec(); + (Some(af_body), payload) + } + 0b10 => { + let af_len = body[0] as usize; + (Some(body[1..1 + af_len].to_vec()), Vec::new()) + } + _ => (None, Vec::new()), + }; + + out.push(TsPacket { + pid, + pusi, + cc, + af, + payload, + }); + } + out + } + + /// Build a fake HEVC NAL with a 4-byte length prefix. + /// nal_type=19/20 are IDR; 1 is non-key (TRAIL_N/R). + fn fake_hevc_nal(nal_type: u8, body_len: usize) -> Vec { + let mut nal = Vec::with_capacity(2 + body_len); + // 2-byte NAL header: forbidden_zero(1)=0 | nal_unit_type(6) | layer_id(6)=0 | tid_plus1(3)=1 + nal.push((nal_type & 0x3F) << 1); + nal.push(0x01); + for i in 0..body_len { + nal.push((i & 0xFF) as u8); + } + let mut framed = Vec::with_capacity(4 + nal.len()); + framed.extend_from_slice(&(nal.len() as u32).to_be_bytes()); + framed.extend_from_slice(&nal); + framed + } + + #[test] + fn keyframe_param_threads_through() { + let mut sink: Vec = Vec::new(); + { + let mut mux = TsMuxer::new(&mut sink, &[VIDEO_PID]); + let idr = fake_hevc_nal(19, 100); + mux.write_frame(0, 0, true, &idr).unwrap(); + let p = fake_hevc_nal(1, 80); + mux.write_frame(0, 41_000_000, false, &p).unwrap(); + mux.finish().unwrap(); + } + assert!(!sink.is_empty()); + let packets = parse_bd_ts(&sink); + assert!(packets.iter().any(|p| p.pid == VIDEO_PID && p.pusi)); + } + + #[test] + fn rai_set_on_first_packet_of_keyframe_pes() { + let mut sink: Vec = Vec::new(); + { + let mut mux = TsMuxer::new(&mut sink, &[VIDEO_PID]); + let idr = fake_hevc_nal(19, 200); + mux.write_frame(0, 0, true, &idr).unwrap(); + mux.finish().unwrap(); + } + let packets = parse_bd_ts(&sink); + let first_pusi = packets + .iter() + .find(|p| p.pid == VIDEO_PID && p.pusi) + .expect("video PUSI packet exists"); + let af = first_pusi.af.as_ref().expect("AF present on keyframe PES"); + assert!(!af.is_empty(), "AF body has flags byte"); + assert_eq!(af[0] & 0x40, 0x40, "RAI bit set"); + } + + #[test] + fn rai_clear_on_non_keyframe_pes() { + let mut sink: Vec = Vec::new(); + { + let mut mux = TsMuxer::new(&mut sink, &[VIDEO_PID]); + let idr = fake_hevc_nal(19, 100); + mux.write_frame(0, 0, true, &idr).unwrap(); + let p = fake_hevc_nal(1, 100); + mux.write_frame(0, 41_000_000, false, &p).unwrap(); + mux.finish().unwrap(); + } + let packets = parse_bd_ts(&sink); + // Second PUSI packet on the video PID belongs to the non-key frame. + let pusi_video: Vec<&TsPacket> = packets + .iter() + .filter(|p| p.pid == VIDEO_PID && p.pusi) + .collect(); + assert!(pusi_video.len() >= 2, "two PUSI packets expected"); + let second = pusi_video[1]; + match &second.af { + None => {} + Some(af) if af.is_empty() => {} // length=0 case + Some(af) => assert_eq!(af[0] & 0x40, 0, "RAI must be clear on non-key PES"), + } + } + + #[test] + fn codec_private_prepended_only_on_first_keyframe() { + // Build a minimal hvcC with one recognizable NAL. + let marker: &[u8] = &[0xDE, 0xAD, 0xBE, 0xEF, 0xCA, 0xFE]; + let mut hvcc = vec![0u8; 22]; + hvcc.push(1); // numArrays + hvcc.push(32); // VPS NAL type byte (high bits arbitrary) + hvcc.extend_from_slice(&1u16.to_be_bytes()); // numNalus + hvcc.extend_from_slice(&(marker.len() as u16).to_be_bytes()); + hvcc.extend_from_slice(marker); + + let mut sink: Vec = Vec::new(); + { + let mut mux = TsMuxer::new(&mut sink, &[VIDEO_PID]); + mux.set_codec_private(0, hvcc); + // Non-IDR before any IDR: should be dropped. + let p = fake_hevc_nal(1, 50); + mux.write_frame(0, 0, false, &p).unwrap(); + // IDR: should carry codec_private NALs prepended. + let idr = fake_hevc_nal(19, 50); + mux.write_frame(0, 41_000_000, true, &idr).unwrap(); + mux.finish().unwrap(); + } + let packets = parse_bd_ts(&sink); + // Concatenate all video PID payloads in emission order. + let video_bytes: Vec = packets + .iter() + .filter(|p| p.pid == VIDEO_PID) + .flat_map(|p| p.payload.clone()) + .collect(); + // marker bytes must appear in the stream (codec_private was prepended). + let pos_marker = video_bytes + .windows(marker.len()) + .position(|w| w == marker) + .expect("codec_private marker bytes present in TS payload"); + // Find IDR body byte (0x26 = (19<<1)). pos_idr must be AFTER marker. + let idr_header = (19u8 << 1) & 0x7E; + let pos_idr = video_bytes + .iter() + .position(|&b| b == idr_header) + .expect("IDR NAL header present in TS payload"); + assert!( + pos_marker < pos_idr, + "codec_private must precede IDR in TS payload" + ); + } + + #[test] + fn non_key_before_first_keyframe_dropped() { + let mut sink: Vec = Vec::new(); + { + let mut mux = TsMuxer::new(&mut sink, &[VIDEO_PID]); + let p = fake_hevc_nal(1, 80); + mux.write_frame(0, 0, false, &p).unwrap(); + mux.finish().unwrap(); + } + // Nothing should be emitted for that PID. + let packets = parse_bd_ts(&sink); + assert!( + !packets.iter().any(|p| p.pid == VIDEO_PID), + "non-key before first keyframe must be dropped" + ); + } +}