| Bring vidya in cfd3e36 nandi 19d ago | 1 | //! H.264-in-MP4 decode session (openh264 + mp4 demux). |
| 2 | |
| 3 | use std::io::Cursor; |
| 4 | use std::sync::Arc; |
| 5 | use std::time::Duration; |
| 6 | |
| 7 | use openh264::decoder::{Decoder, DecoderConfig, Flush}; |
| 8 | use openh264::formats::YUVSource; |
| 9 | use openh264::OpenH264API; |
| 10 | |
| 11 | use super::avcc::Mp4BitstreamConverter; |
| 12 | |
| 13 | /// Decoded poster / playback session for one MP4 byte buffer. |
| 14 | pub struct DecodeSession { |
| 15 | mp4: mp4::Mp4Reader<Cursor<Arc<[u8]>>>, |
| 16 | track_id: u32, |
| 17 | sample_count: u32, |
| 18 | timescale: u32, |
| 19 | converter: Mp4BitstreamConverter, |
| 20 | decoder: Decoder, |
| 21 | pub width: u32, |
| 22 | pub height: u32, |
| 23 | /// Presentation time (seconds) for sample index `i` (0-based). |
| 24 | sample_pts: Vec<f64>, |
| 25 | pub duration: Duration, |
| 26 | /// 1-based sample id last successfully fed to the decoder. |
| 27 | next_sample: u32, |
| 28 | buffer: Vec<u8>, |
| 29 | rgba: Vec<u8>, |
| 30 | /// Latest decoded frame (RGBA). |
| 31 | pub frame: Option<(u32, u32, Vec<u8>)>, |
| 32 | } |
| 33 | |
| 34 | impl DecodeSession { |
| 35 | pub fn open(bytes: Arc<[u8]>) -> Result<Self, String> { |
| 36 | if bytes.len() < 12 { |
| 37 | return Err("Video too small".into()); |
| 38 | } |
| 39 | // Reject obvious non-MP4 (WebM starts with EBML 0x1A45DFA3). |
| 40 | if bytes.starts_with(&[0x1A, 0x45, 0xDF, 0xA3]) { |
| 41 | return Err("WebM is not supported yet (H.264 MP4 only)".into()); |
| 42 | } |
| 43 | |
| 44 | let size = bytes.len() as u64; |
| 45 | let mut mp4 = mp4::Mp4Reader::read_header(Cursor::new(bytes), size) |
| 46 | .map_err(|e| format!("MP4 header: {e}"))?; |
| 47 | |
| 48 | let (track_id, width, height, sample_count, timescale, duration_secs) = { |
| 49 | let track = mp4 |
| 50 | .tracks() |
| 51 | .iter() |
| 52 | .find(|(_, t)| matches!(t.track_type(), Ok(mp4::TrackType::Video))) |
| 53 | .or_else(|| { |
| 54 | mp4.tracks() |
| 55 | .iter() |
| 56 | .find(|(_, t)| matches!(t.media_type(), Ok(mp4::MediaType::H264))) |
| 57 | }) |
| 58 | .map(|(_, t)| t) |
| 59 | .ok_or_else(|| "No video track in MP4".to_string())?; |
| 60 | |
| 61 | match track.media_type() { |
| 62 | Ok(mp4::MediaType::H264) => {} |
| 63 | Ok(other) => { |
| 64 | return Err(format!("Unsupported video codec in MP4: {other:?}")); |
| 65 | } |
| 66 | Err(e) => return Err(format!("Track media type: {e}")), |
| 67 | } |
| 68 | |
| 69 | let track_id = track.track_id(); |
| 70 | let width = u32::from(track.width()); |
| 71 | let height = u32::from(track.height()); |
| 72 | let sample_count = track.sample_count(); |
| 73 | let timescale = track.timescale().max(1); |
| 74 | let duration_secs = track.duration(); |
| 75 | ( |
| 76 | track_id, |
| 77 | width, |
| 78 | height, |
| 79 | sample_count, |
| 80 | timescale, |
| 81 | duration_secs, |
| 82 | ) |
| 83 | }; |
| 84 | |
| 85 | if sample_count == 0 || width == 0 || height == 0 { |
| 86 | return Err("Empty video track".into()); |
| 87 | } |
| 88 | |
| 89 | let converter = { |
| 90 | let track = mp4 |
| 91 | .tracks() |
| 92 | .get(&track_id) |
| 93 | .ok_or_else(|| "Missing track after probe".to_string())?; |
| 94 | Mp4BitstreamConverter::for_mp4_track(track)? |
| 95 | }; |
| 96 | |
| 97 | let decoder_options = DecoderConfig::new().flush_after_decode(Flush::NoFlush); |
| 98 | let decoder = Decoder::with_api_config(OpenH264API::from_source(), decoder_options) |
| 99 | .map_err(|e| format!("OpenH264 init: {e}"))?; |
| 100 | |
| 101 | let mut sample_pts = Vec::with_capacity(sample_count as usize); |
| 102 | for i in 1..=sample_count { |
| 103 | if let Ok(Some(sample)) = mp4.read_sample(track_id, i) { |
| 104 | sample_pts.push(sample.start_time as f64 / timescale as f64); |
| 105 | } else { |
| 106 | sample_pts.push(sample_pts.last().copied().unwrap_or(0.0)); |
| 107 | } |
| 108 | } |
| 109 | |
| 110 | let mut session = Self { |
| 111 | mp4, |
| 112 | track_id, |
| 113 | sample_count, |
| 114 | timescale, |
| 115 | converter, |
| 116 | decoder, |
| 117 | width, |
| 118 | height, |
| 119 | sample_pts, |
| 120 | duration: duration_secs, |
| 121 | next_sample: 1, |
| 122 | buffer: Vec::new(), |
| 123 | rgba: vec![0; (width as usize) * (height as usize) * 4], |
| 124 | frame: None, |
| 125 | }; |
| 126 | |
| 127 | // Decode until we have a poster frame. |
| 128 | session.decode_until_frame()?; |
| 129 | if session.frame.is_none() { |
| 130 | return Err("Could not decode a video frame".into()); |
| 131 | } |
| 132 | Ok(session) |
| 133 | } |
| 134 | |
| 135 | /// Advance decoding so a frame at or past `t` seconds is available. |
| 136 | pub fn seek_playhead(&mut self, t: f64) -> Result<(), String> { |
| 137 | let t = t.clamp(0.0, self.duration.as_secs_f64().max(0.0)); |
| 138 | // Find the first sample whose PTS is >= t (or the last sample). |
| 139 | let mut target = self.sample_count; |
| 140 | for (idx, pts) in self.sample_pts.iter().enumerate() { |
| 141 | if *pts + f64::EPSILON >= t { |
| 142 | target = (idx as u32) + 1; |
| 143 | break; |
| 144 | } |
| 145 | } |
| 146 | if target < self.next_sample { |
| 147 | // Need to restart decoder to go backwards / loop. |
| 148 | self.reset_decoder()?; |
| 149 | } |
| 150 | while self.next_sample <= target { |
| 151 | if !self.feed_next_sample()? { |
| 152 | break; |
| 153 | } |
| 154 | } |
| 155 | Ok(()) |
| 156 | } |
| 157 | |
| 158 | fn reset_decoder(&mut self) -> Result<(), String> { |
| 159 | let decoder_options = DecoderConfig::new().flush_after_decode(Flush::NoFlush); |
| 160 | self.decoder = Decoder::with_api_config(OpenH264API::from_source(), decoder_options) |
| 161 | .map_err(|e| format!("OpenH264 re-init: {e}"))?; |
| 162 | // Rebuild converter (SPS inject state). |
| 163 | let track = self |
| 164 | .mp4 |
| 165 | .tracks() |
| 166 | .get(&self.track_id) |
| 167 | .ok_or_else(|| "Missing track".to_string())?; |
| 168 | self.converter = Mp4BitstreamConverter::for_mp4_track(track)?; |
| 169 | self.next_sample = 1; |
| 170 | Ok(()) |
| 171 | } |
| 172 | |
| 173 | fn decode_until_frame(&mut self) -> Result<(), String> { |
| 174 | while self.frame.is_none() && self.next_sample <= self.sample_count { |
| 175 | self.feed_next_sample()?; |
| 176 | } |
| 177 | Ok(()) |
| 178 | } |
| 179 | |
| 180 | /// Feed one sample. Returns false when the stream is exhausted. |
| 181 | fn feed_next_sample(&mut self) -> Result<bool, String> { |
| 182 | if self.next_sample > self.sample_count { |
| 183 | return Ok(false); |
| 184 | } |
| 185 | let sample_id = self.next_sample; |
| 186 | self.next_sample += 1; |
| 187 | |
| 188 | let Some(sample) = self |
| 189 | .mp4 |
| 190 | .read_sample(self.track_id, sample_id) |
| 191 | .map_err(|e| format!("Read sample: {e}"))? |
| 192 | else { |
| 193 | return Ok(true); |
| 194 | }; |
| 195 | |
| 196 | self.converter |
| 197 | .convert_packet(&sample.bytes, &mut self.buffer); |
| 198 | if self.buffer.is_empty() { |
| 199 | return Ok(true); |
| 200 | } |
| 201 | |
| 202 | match self.decoder.decode(&self.buffer) { |
| 203 | Ok(Some(yuv)) => { |
| 204 | let (w, h) = yuv.dimensions(); |
| 205 | let need = w * h * 4; |
| 206 | if self.rgba.len() != need { |
| 207 | self.rgba.resize(need, 0); |
| 208 | self.width = w as u32; |
| 209 | self.height = h as u32; |
| 210 | } |
| 211 | yuv.write_rgba8(&mut self.rgba); |
| 212 | self.frame = Some((self.width, self.height, self.rgba.clone())); |
| 213 | } |
| 214 | Ok(None) => {} |
| 215 | Err(e) => { |
| 216 | // Soft-fail individual samples (B-frames / corruption). |
| 217 | let _ = e; |
| 218 | } |
| 219 | } |
| 220 | Ok(true) |
| 221 | } |
| 222 | |
| 223 | pub fn ended(&self, t: f64) -> bool { |
| 224 | t >= self.duration.as_secs_f64() && self.next_sample > self.sample_count |
| 225 | } |
| 226 | |
| 227 | #[allow(dead_code)] |
| 228 | pub fn timescale(&self) -> u32 { |
| 229 | self.timescale |
| 230 | } |
| 231 | } |
| 232 | |
| 233 | #[cfg(all(test, feature = "video"))] |
| 234 | mod tests { |
| 235 | use super::*; |
| 236 | use std::path::PathBuf; |
| 237 | |
| 238 | #[test] |
| 239 | fn decodes_sample_h264_mp4() { |
| 240 | let path = PathBuf::from("/tmp/sleek-test.mp4"); |
| 241 | if !path.exists() { |
| 242 | eprintln!("skip: missing {}", path.display()); |
| 243 | return; |
| 244 | } |
| 245 | let bytes: Arc<[u8]> = std::fs::read(&path).unwrap().into(); |
| 246 | let mut session = DecodeSession::open(bytes).expect("open"); |
| 247 | assert!(session.width > 0 && session.height > 0); |
| 248 | assert!(session.frame.is_some()); |
| 249 | session.seek_playhead(0.5).expect("seek"); |
| 250 | assert!(session.frame.is_some()); |
| 251 | } |
| 252 | } |