1use std::collections::HashMap;
4#[cfg(target_os = "uefi")]
5use std::collections::hash_map::DefaultHasher;
6#[cfg(target_os = "uefi")]
7use std::hash::BuildHasherDefault;
8use std::io::{Read, Seek, SeekFrom};
9
10use crate::asf::{AsfFile, AsfPayload, VideoStreamInfo};
11use crate::decoder::{MacroblockDecoder, YuvFrame};
12use crate::error::{DecoderError, Result};
13#[cfg(feature = "audio")]
14use crate::wma::{PcmFrameF32, WmaDecoder};
15use crate::wmv2::{Wmv2FrameHeader, Wmv2FrameType, Wmv2Params};
16
17#[derive(Clone)]
19pub struct DecodedFrame {
20 pub pts_ms: u32,
21 pub is_key_frame: bool,
22 pub frame: YuvFrame,
23}
24
25#[cfg(feature = "audio")]
27#[derive(Clone)]
28pub struct DecodedAudioFrame {
29 pub pts_ms: u32,
30 pub frame: PcmFrameF32,
31}
32
33pub struct Wmv2Decoder {
37 params: Wmv2Params,
38 mb_dec: MacroblockDecoder,
39 cur: YuvFrame,
40 locked_hdr_off: Option<usize>,
41}
42
43impl Wmv2Decoder {
44 pub fn new(width: u32, height: u32, extradata: &[u8]) -> Self {
48 let params = Wmv2Params::new(width, height);
49 let mut mb_dec = MacroblockDecoder::new(width, height);
50 mb_dec.wmv2_set_extradata(extradata);
51 let cur = YuvFrame::new(width, height);
52 Self {
53 params,
54 mb_dec,
55 cur,
56 locked_hdr_off: None,
57 }
58 }
59
60 pub fn width(&self) -> u32 {
61 self.params.width
62 }
63
64 pub fn height(&self) -> u32 {
65 self.params.height
66 }
67
68 pub fn current_frame(&self) -> &YuvFrame {
72 &self.cur
73 }
74
75 pub fn decode_frame(
79 &mut self,
80 payload: &[u8],
81 is_key_frame: bool,
82 ) -> Result<Option<&YuvFrame>> {
83 if payload.is_empty() {
84 return Ok(None);
85 }
86
87 let mut best_score: i64 = -1;
88 let mut best_off: usize = 0;
89 let mut best_hdr: Option<Wmv2FrameHeader> = None;
90
91 let mut offs: Vec<usize> = Vec::with_capacity(18);
93 if let Some(o) = self.locked_hdr_off {
94 offs.push(o);
95 }
96 for o in 0..=16 {
97 if Some(o) != self.locked_hdr_off {
98 offs.push(o);
99 }
100 }
101
102 for off in offs {
103 if off > payload.len() {
104 continue;
105 }
106 let cands = Wmv2FrameHeader::parse_candidates(
107 &payload[off..],
108 self.mb_dec.width_mb,
109 self.mb_dec.height_mb,
110 );
111 if cands.is_empty() {
112 continue;
113 }
114 for h in cands {
115 if is_key_frame && h.frame_type != Wmv2FrameType::I {
117 continue;
118 }
119
120 let mut sc: i64 = if h.frame_skipped {
122 1
123 } else if is_key_frame {
124 2
125 } else {
126 self.mb_dec.probe_wmv2_payload(&payload[off..], &h) as i64
127 };
128
129 if Some(off) == self.locked_hdr_off {
130 sc += 64;
131 }
132
133 if sc > best_score {
134 best_score = sc;
135 best_off = off;
136 best_hdr = Some(h);
137 }
138 }
139 }
140
141 let Some(hdr) = best_hdr else {
142 return Ok(None);
143 };
144
145 if self.locked_hdr_off.is_none() {
146 self.locked_hdr_off = Some(best_off);
147 }
148
149 let frame_data = &payload[best_off..];
150 self.mb_dec
151 .decode_wmv2_frame(frame_data, &hdr, &self.params, &mut self.cur)?;
152 Ok(Some(&self.cur))
153 }
154
155 pub fn decode_frame_owned(
157 &mut self,
158 payload: &[u8],
159 is_key_frame: bool,
160 ) -> Result<Option<YuvFrame>> {
161 let Some(f) = self.decode_frame(payload, is_key_frame)? else {
162 return Ok(None);
163 };
164 Ok(Some(f.clone()))
165 }
166}
167
168#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
173struct FrameKey {
174 stream_number: u8,
175 object_id: u32,
176}
177
178#[derive(Debug, Clone)]
179struct FrameAssembly {
180 total: usize,
181 pts_ms: u32,
182 is_key: bool,
183 data: Vec<u8>,
184 ranges: Vec<(usize, usize)>,
185}
186
187impl FrameAssembly {
188 fn new(total: usize, pts_ms: u32, is_key: bool) -> Self {
189 Self {
190 total,
191 pts_ms,
192 is_key,
193 data: vec![0u8; total],
194 ranges: Vec::new(),
195 }
196 }
197
198 fn insert(&mut self, offset: usize, frag: &[u8]) {
199 if self.total == 0 || offset >= self.total || frag.is_empty() {
200 return;
201 }
202 let end = (offset + frag.len()).min(self.total);
203 let n = end - offset;
204 self.data[offset..end].copy_from_slice(&frag[..n]);
205 self.add_range(offset, end);
206 }
207
208 fn add_range(&mut self, start: usize, end: usize) {
209 if start >= end {
210 return;
211 }
212 self.ranges.push((start, end));
213 self.ranges.sort_by_key(|r| r.0);
214
215 let mut merged: Vec<(usize, usize)> = Vec::with_capacity(self.ranges.len());
216 for (s, e) in self.ranges.drain(..) {
217 if let Some(last) = merged.last_mut() {
218 if s <= last.1 {
219 last.1 = last.1.max(e);
220 continue;
221 }
222 }
223 merged.push((s, e));
224 }
225 self.ranges = merged;
226 }
227
228 fn covered_len(&self) -> usize {
229 self.ranges.iter().map(|(s, e)| e - s).sum()
230 }
231
232 fn is_complete(&self) -> bool {
233 self.total > 0
234 && self.covered_len() >= self.total
235 && self.ranges.len() == 1
236 && self.ranges[0] == (0, self.total)
237 }
238}
239
240#[cfg(target_os = "uefi")]
241type InFlightMap = HashMap<FrameKey, FrameAssembly, BuildHasherDefault<DefaultHasher>>;
242#[cfg(not(target_os = "uefi"))]
243type InFlightMap = HashMap<FrameKey, FrameAssembly>;
244
245#[derive(Default)]
246struct FrameAssembler {
247 in_flight: InFlightMap,
248}
249
250impl FrameAssembler {
251 fn push(&mut self, payload: AsfPayload) -> Option<(u32, bool, Vec<u8>)> {
252 if payload.data.is_empty() {
253 return None;
254 }
255
256 let key = FrameKey {
257 stream_number: payload.stream_number,
258 object_id: payload.object_id,
259 };
260
261 if payload.obj_offset == 0 {
263 let osz = payload.obj_size as usize;
264 if osz == 0 || osz == payload.data.len() {
265 return Some((payload.pts_ms, payload.is_key_frame, payload.data));
266 }
267 }
268
269 if payload.obj_size == 0 {
271 return Some((payload.pts_ms, payload.is_key_frame, payload.data));
272 }
273
274 let total = payload.obj_size as usize;
275 let entry = self
276 .in_flight
277 .entry(key)
278 .or_insert_with(|| FrameAssembly::new(total, payload.pts_ms, payload.is_key_frame));
279
280 entry.is_key |= payload.is_key_frame;
282 entry.insert(payload.obj_offset as usize, &payload.data);
283
284 if entry.is_complete() {
285 let assembly = self.in_flight.remove(&key).unwrap();
286 return Some((assembly.pts_ms, assembly.is_key, assembly.data));
287 }
288 None
289 }
290}
291
292pub struct AsfWmv2Decoder<R: Read + Seek> {
297 reader: R,
298 asf: AsfFile,
299 video_info: VideoStreamInfo,
300 assembler: FrameAssembler,
301 decoder: Wmv2Decoder,
302}
303
304#[cfg(feature = "audio")]
309pub struct AsfWmaDecoder<R: Read + Seek> {
310 reader: R,
311 asf: AsfFile,
312 audio_stream_number: u8,
313 decoder: WmaDecoder,
314 assembler: FrameAssembler,
315 last_pts_ms: u32,
316 flushed_eof: bool,
317}
318
319#[cfg(feature = "audio")]
320impl<R: Read + Seek> AsfWmaDecoder<R> {
321 pub fn open(mut reader: R) -> Result<Self> {
326 let asf = AsfFile::open(&mut reader)?;
327 let mut chosen = None;
328 for a in asf.audio_streams.iter() {
329 if matches!(a.format_tag, 0x0160 | 0x0161) {
330 chosen = Some(a.clone());
331 break;
332 }
333 }
334 let Some(audio_info) = chosen else {
335 return Err(DecoderError::Unsupported(
336 "No supported WMA (0x0160/0x0161) audio stream found".into(),
337 ));
338 };
339
340 reader.seek(SeekFrom::Start(asf.data_offset))?;
341 let decoder = WmaDecoder::new(&audio_info)?;
342
343 Ok(Self {
344 reader,
345 asf,
346 audio_stream_number: audio_info.stream_number,
347 decoder,
348 assembler: FrameAssembler::default(),
349 last_pts_ms: 0,
350 flushed_eof: false,
351 })
352 }
353
354 pub fn sample_rate(&self) -> u32 {
355 self.decoder.sample_rate()
356 }
357
358 pub fn channels(&self) -> u16 {
359 self.decoder.channels()
360 }
361
362 pub fn next_frame(&mut self) -> Result<Option<DecodedAudioFrame>> {
366 loop {
367 let payloads = match self.asf.read_packet(&mut self.reader) {
368 Ok(p) => p,
369 Err(DecoderError::EndOfStream) => {
370 if self.flushed_eof {
371 return Ok(None);
372 }
373 self.flushed_eof = true;
374 if let Some(frame) = self.decoder.decode_packet(&[], self.last_pts_ms)? {
375 return Ok(Some(DecodedAudioFrame {
376 pts_ms: frame.pts_ms,
377 frame,
378 }));
379 }
380 return Ok(None);
381 }
382 Err(e) => return Err(e),
383 };
384
385 for payload in payloads {
386 if payload.stream_number != self.audio_stream_number {
387 continue;
388 }
389 let Some((pts_ms, _is_key, data)) = self.assembler.push(payload) else {
390 continue;
391 };
392 self.last_pts_ms = pts_ms;
393 if let Some(frame) = self.decoder.decode_packet(&data, pts_ms)? {
394 return Ok(Some(DecodedAudioFrame { pts_ms, frame }));
395 }
396 }
397 }
398 }
399}
400
401impl<R: Read + Seek> AsfWmv2Decoder<R> {
402 pub fn open(mut reader: R) -> Result<Self> {
406 let asf = AsfFile::open(&mut reader)?;
407 let mut video_info: Option<VideoStreamInfo> = None;
408 for v in asf.video_streams.iter() {
409 let four_cc = std::str::from_utf8(&v.codec_four_cc)
410 .unwrap_or("")
411 .to_uppercase();
412 if matches!(four_cc.as_str(), "WMV2" | "WMV1") {
413 video_info = Some(v.clone());
414 break;
415 }
416 }
417 let Some(video_info) = video_info else {
418 return Err(DecoderError::Unsupported(
419 "No WMV2/WMV1 video stream found".into(),
420 ));
421 };
422
423 reader.seek(SeekFrom::Start(asf.data_offset))?;
424
425 let decoder = Wmv2Decoder::new(video_info.width, video_info.height, &video_info.extra_data);
426
427 Ok(Self {
428 reader,
429 asf,
430 video_info,
431 assembler: FrameAssembler::default(),
432 decoder,
433 })
434 }
435
436 pub fn video_stream_info(&self) -> &VideoStreamInfo {
438 &self.video_info
439 }
440
441 pub fn next_frame(&mut self) -> Result<Option<DecodedFrame>> {
445 loop {
446 let payloads = match self.asf.read_packet(&mut self.reader) {
447 Ok(p) => p,
448 Err(DecoderError::EndOfStream) => return Ok(None),
449 Err(e) => return Err(e),
450 };
451
452 for payload in payloads {
453 if payload.stream_number != self.video_info.stream_number {
454 continue;
455 }
456 let Some((pts_ms, is_key, data)) = self.assembler.push(payload) else {
457 continue;
458 };
459
460 if let Some(frame) = self.decoder.decode_frame_owned(&data, is_key)? {
461 return Ok(Some(DecodedFrame {
462 pts_ms,
463 is_key_frame: is_key,
464 frame,
465 }));
466 }
467 }
468 }
469 }
470}