Skip to main content

wmv_decoder/
api.rs

1//! Public library API.
2
3use 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/// A decoded video frame with timing metadata.
18#[derive(Clone)]
19pub struct DecodedFrame {
20    pub pts_ms: u32,
21    pub is_key_frame: bool,
22    pub frame: YuvFrame,
23}
24
25/// A decoded audio frame with timing metadata.
26#[cfg(feature = "audio")]
27#[derive(Clone)]
28pub struct DecodedAudioFrame {
29    pub pts_ms: u32,
30    pub frame: PcmFrameF32,
31}
32
33/// WMV2 (Windows Media Video 8) decoder.
34///
35/// The picture header parsing and macroblock decode paths are aligned with upstream.
36pub struct Wmv2Decoder {
37    params: Wmv2Params,
38    mb_dec: MacroblockDecoder,
39    cur: YuvFrame,
40    locked_hdr_off: Option<usize>,
41}
42
43impl Wmv2Decoder {
44    /// Create a decoder for a fixed resolution.
45    ///
46    /// `extradata` is the 4-byte WMV2 ext header typically carried in ASF stream properties.
47    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    /// Borrow the internal YUV420p frame buffer.
69    ///
70    /// The returned reference stays valid until the next successful decode.
71    pub fn current_frame(&self) -> &YuvFrame {
72        &self.cur
73    }
74
75    /// Decode one assembled WMV2 frame payload.
76    ///
77    /// Returns `Ok(None)` if no plausible picture header can be found.
78    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        // Try the previously locked offset first, then fall back to a small scan.
92        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                // ASF keyframe marking should correspond to WMV2 I pictures.
116                if is_key_frame && h.frame_type != Wmv2FrameType::I {
117                    continue;
118                }
119
120                // upstream-aligned scoring strategy.
121                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    /// Decode and return an owned frame buffer (clone).
156    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// ─────────────────────────────────────────────────────────────────────────────
169// ASF media-object reassembly (frame reassembly)
170// ─────────────────────────────────────────────────────────────────────────────
171
172#[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        // Fast path: complete media object in one payload (or size unknown).
262        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 the total object size is unknown, we cannot reliably reassemble.
270        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        // Update meta (first PTS wins; keyframe if any fragment says so).
281        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
292/// ASF + WMV2 decoding pipeline.
293///
294/// This type owns the `Read+Seek` source, parses ASF headers, reassembles media objects
295/// and decodes WMV2 frames into `YuvFrame`.
296pub struct AsfWmv2Decoder<R: Read + Seek> {
297    reader: R,
298    asf: AsfFile,
299    video_info: VideoStreamInfo,
300    assembler: FrameAssembler,
301    decoder: Wmv2Decoder,
302}
303
304/// ASF + WMA (v1/v2) decoding pipeline.
305///
306/// This type owns the `Read+Seek` source, parses ASF headers, reassembles media objects
307/// and decodes WMA packets into PCM.
308#[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    /// Open an ASF/WMV stream and initialize the WMA decoder.
322    ///
323    /// The decoder selects the first audio stream with format tag 0x0160 (WMAv1)
324    /// or 0x0161 (WMAv2).
325    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    /// Decode the next audio frame.
363    ///
364    /// Returns `Ok(None)` on end-of-stream.
365    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    /// Open an ASF/WMV stream and initialize the WMV2 decoder.
403    ///
404    /// The decoder selects the first video stream whose FourCC is WMV2 or WMV1.
405    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    /// Return the selected video stream info.
437    pub fn video_stream_info(&self) -> &VideoStreamInfo {
438        &self.video_info
439    }
440
441    /// Decode the next video frame.
442    ///
443    /// Returns `Ok(None)` on end-of-stream.
444    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}