Skip to main content

wmv_decoder/
vc1.rs

1/// VC-1 / WMV9 Sequence & Picture Header Parser + Bitplane Decoder
2///
3use crate::bitreader::BitReader;
4use crate::error::{DecoderError, Result};
5
6// ─── Enums ───────────────────────────────────────────────────────────────────
7
8#[derive(Debug, Clone, Copy, PartialEq, Eq)]
9pub enum Profile {
10    Simple,
11    Main,
12    Advanced,
13}
14
15#[derive(Debug, Clone, Copy, PartialEq, Eq)]
16pub enum FrameType {
17    I,
18    P,
19    B,
20    BI,
21    Skipped,
22}
23
24impl std::fmt::Display for FrameType {
25    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
26        match self {
27            FrameType::I => write!(f, "I"),
28            FrameType::P => write!(f, "P"),
29            FrameType::B => write!(f, "B"),
30            FrameType::BI => write!(f, "BI"),
31            FrameType::Skipped => write!(f, "skip"),
32        }
33    }
34}
35
36#[derive(Debug, Clone, Copy, PartialEq, Eq)]
37pub enum QuantizerMode {
38    Implicit,
39    Explicit,
40    NonUniform,
41    Uniform,
42}
43
44// ─── Sequence Header ─────────────────────────────────────────────────────────
45
46#[derive(Debug, Clone)]
47pub struct SequenceHeader {
48    pub profile: Profile,
49    pub max_b_frames: u8,
50    pub frame_rate_num: u32,
51    pub frame_rate_den: u32,
52    pub loop_filter: bool,
53    pub multires: bool,
54    pub fastuvmc: bool,
55    pub extended_mv: bool,
56    pub dquant: u8,
57    pub vstransform: bool,
58    pub overlap: bool,
59    pub syncmarker: bool,
60    pub rangered: bool,
61    pub quantizer_mode: QuantizerMode,
62    pub finterpflag: bool,
63    pub transacfrm: u8,  // inter AC table index 0-3 (TRANSACFRM)
64    pub transacfrm2: u8, // intra AC table index 0-3 (TRANSACFRM2)
65    pub mvtab: u8,       // MV table index 0-3 (MVTAB)
66    pub cbptab: u8,      // CBP table index 0-3 (CBPTAB)
67    pub dctab: bool,     // DC table select (DCTAB)
68    pub width: u32,
69    pub height: u32,
70    pub display_width: u32,
71    pub display_height: u32,
72}
73
74impl SequenceHeader {
75    pub fn parse(data: &[u8]) -> Result<Self> {
76        if data.len() < 4 {
77            return Err(DecoderError::InvalidData(
78                "Sequence header too short".into(),
79            ));
80        }
81        let mut br = BitReader::new(data);
82
83        // WMV9: first 2 bits = profile
84        let profile_bits = br.read_bits(2).unwrap_or(0) as u8;
85        let profile = match profile_bits {
86            0 => Profile::Simple,
87            1 => Profile::Main,
88            3 => Profile::Advanced,
89            _ => Profile::Main,
90        };
91
92        br.read_bits(2); // reserved
93
94        let frmrtq_postproc = br.read_bits(3).unwrap_or(0);
95        let _bitrtq_postproc = br.read_bits(5).unwrap_or(0);
96        let loop_filter = br.read_bit().unwrap_or(false);
97        let _res_sm = br.read_bit().unwrap_or(false);
98        let multires = br.read_bit().unwrap_or(false);
99        let _res_fasttx = br.read_bit().unwrap_or(true);
100        let fastuvmc = br.read_bit().unwrap_or(false);
101        let extended_mv = br.read_bit().unwrap_or(false);
102        let dquant = br.read_bits(2).unwrap_or(0) as u8;
103        let vstransform = br.read_bit().unwrap_or(false);
104        let _res_transtab = br.read_bit().unwrap_or(false);
105        let overlap = br.read_bit().unwrap_or(false);
106        let _resync_marker = br.read_bit().unwrap_or(false);
107        let rangered = br.read_bit().unwrap_or(false);
108        let max_b_frames = br.read_bits(3).unwrap_or(0) as u8;
109        let quant_bits = br.read_bits(2).unwrap_or(0) as u8;
110        let finterpflag = br.read_bit().unwrap_or(false);
111        let syncmarker = br.read_bit().unwrap_or(false);
112        // Additional fields per SMPTE 421M §8.1.1 (Simple/Main)
113        let transacfrm = br.read_bits(2).unwrap_or(0) as u8;
114        let transacfrm2 = br.read_bits(2).unwrap_or(0) as u8;
115        let mvtab = br.read_bits(2).unwrap_or(0) as u8;
116        let cbptab = br.read_bits(2).unwrap_or(0) as u8;
117        let dctab = br.read_bit().unwrap_or(false);
118
119        let quantizer_mode = match quant_bits {
120            0 => QuantizerMode::Implicit,
121            1 => QuantizerMode::Explicit,
122            2 => QuantizerMode::NonUniform,
123            _ => QuantizerMode::Uniform,
124        };
125
126        let (frame_rate_num, frame_rate_den) = match frmrtq_postproc {
127            0 => (6, 1),
128            1 => (8, 1),
129            2 => (10, 1),
130            3 => (12, 1),
131            4 => (15, 1),
132            5 => (24000, 1001),
133            6 => (24, 1),
134            7 => (25, 1),
135            _ => (30, 1),
136        };
137
138        Ok(SequenceHeader {
139            profile,
140            max_b_frames,
141            frame_rate_num,
142            frame_rate_den,
143            loop_filter,
144            multires,
145            fastuvmc,
146            extended_mv,
147            dquant,
148            vstransform,
149            overlap,
150            syncmarker,
151            rangered,
152            quantizer_mode,
153            finterpflag,
154            transacfrm,
155            transacfrm2,
156            mvtab,
157            cbptab,
158            dctab,
159            width: 0,
160            height: 0,
161            display_width: 0,
162            display_height: 0,
163        })
164    }
165}
166
167// ─── Bitplane ────────────────────────────────────────────────────────────────
168// SMPTE 421M §8.7.  Used to signal skipped MBs and direct-mode flags.
169
170#[derive(Debug, Clone, Copy, PartialEq, Eq)]
171enum BitplaneMode {
172    Norm2,
173    Diff2,
174    Norm6,
175    Diff6,
176    RowSkip,
177    ColSkip,
178}
179
180pub struct Bitplane {
181    pub data: Vec<u8>, // one byte per macroblock (0 or 1)
182    pub is_raw: bool,
183}
184
185impl Bitplane {
186    pub fn decode(br: &mut BitReader<'_>, mb_w: usize, mb_h: usize) -> Option<Self> {
187        let n_mb = mb_w * mb_h;
188        let mut data = vec![0u8; n_mb];
189
190        // 3-bit mode code
191        let mode_bits = br.read_bits(3)?;
192        let mode = match mode_bits {
193            0 => BitplaneMode::Norm2,
194            1 => BitplaneMode::Norm6,
195            2 => BitplaneMode::Diff2,
196            3 => BitplaneMode::Diff6,
197            4 => BitplaneMode::RowSkip,
198            5 => BitplaneMode::ColSkip,
199            _ => {
200                // Raw: one bit per MB
201                for i in 0..n_mb {
202                    data[i] = br.read_bit()? as u8;
203                }
204                return Some(Bitplane { data, is_raw: true });
205            }
206        };
207
208        match mode {
209            BitplaneMode::Norm6 | BitplaneMode::Diff6 => {
210                // Tile-coded 6 MBs per codeword
211                let tile_size = 6usize;
212                let mut inv = br.read_bit()? as u8; // invert flag for Diff modes
213                if !matches!(mode, BitplaneMode::Diff2 | BitplaneMode::Diff6) {
214                    inv = 0;
215                }
216                let mut i = 0;
217                while i < n_mb {
218                    let tile = br.read_bits(tile_size as u8)? as usize;
219                    for b in 0..tile_size.min(n_mb - i) {
220                        data[i + b] = (((tile >> (tile_size - 1 - b)) & 1) as u8) ^ inv;
221                    }
222                    i += tile_size;
223                }
224            }
225            BitplaneMode::Norm2 | BitplaneMode::Diff2 => {
226                let inv = if matches!(mode, BitplaneMode::Diff2) {
227                    br.read_bit()? as u8
228                } else {
229                    0
230                };
231                let mut i = 0;
232                while i < n_mb {
233                    let pair = br.read_bits(2)? as u8;
234                    data[i] = ((pair >> 1) & 1) ^ inv;
235                    if i + 1 < n_mb {
236                        data[i + 1] = (pair & 1) ^ inv;
237                    }
238                    i += 2;
239                }
240            }
241            BitplaneMode::RowSkip => {
242                for row in 0..mb_h {
243                    if br.read_bit()? {
244                        continue;
245                    }
246                    for col in 0..mb_w {
247                        data[row * mb_w + col] = br.read_bit()? as u8;
248                    }
249                }
250            }
251            BitplaneMode::ColSkip => {
252                for col in 0..mb_w {
253                    if br.read_bit()? {
254                        continue;
255                    }
256                    for row in 0..mb_h {
257                        data[row * mb_w + col] = br.read_bit()? as u8;
258                    }
259                }
260            }
261        }
262
263        Some(Bitplane {
264            data,
265            is_raw: false,
266        })
267    }
268}
269
270// ─── Picture Header ───────────────────────────────────────────────────────────
271
272#[derive(Debug, Clone)]
273pub struct PictureHeader {
274    pub frame_type: FrameType,
275    pub pqindex: u8,
276    pub pquant: u8,
277    pub halfqp: bool,
278    pub pqual_mode: u8,
279    pub mvrange: u8,
280    pub rptfrm: u8,
281    pub pts_ms: u32,
282    pub rangeredfrm: bool,
283    /// Bit offset where the macroblock layer starts (from beginning of the frame payload).
284    ///
285    /// This includes the full picture header and any bitplanes decoded from it.
286    pub header_bits: usize,
287    /// Skipped-MB bitplane (None if not present or raw-mode)
288    pub skipmb_plane: Option<Vec<u8>>,
289    /// Direct-mode bitplane for B-frames
290    pub directmb_plane: Option<Vec<u8>>,
291    /// B-frame temporal fraction from SMPTE 421M §7.1.3.6 Table 40.
292    pub bfrac_num: i32,
293    pub bfrac_den: i32,
294}
295
296impl PictureHeader {
297    pub fn parse(
298        data: &[u8],
299        seq: &SequenceHeader,
300        pts_ms: u32,
301        mb_w: usize,
302        mb_h: usize,
303    ) -> Result<Self> {
304        let mut br = BitReader::new(data);
305
306        // ── frame type ──────────────────────────────────────────────────────
307        let frame_type = if seq.max_b_frames > 0 {
308            match br.read_bits(2).unwrap_or(0xFF) {
309                0b11 => FrameType::I,
310                0b10 => FrameType::P,
311                0b00 => FrameType::B,
312                0b01 => FrameType::BI,
313                _ => return Err(DecoderError::InvalidData("Unknown frame type".into())),
314            }
315        } else {
316            match br.read_bit().unwrap_or(false) {
317                false => FrameType::P,
318                true => FrameType::I,
319            }
320        };
321
322        // ── range reduction ─────────────────────────────────────────────────
323        let rangeredfrm = seq.rangered && br.read_bit().unwrap_or(false);
324
325        // ── quantizer ───────────────────────────────────────────────────────
326        let pqindex = br.read_bits(5).unwrap_or(1) as u8;
327        let (pquant, halfqp, pqual_mode) = Self::decode_quantizer(pqindex, seq);
328
329        // ── MV range ────────────────────────────────────────────────────────
330        let mvrange = if seq.extended_mv {
331            let mut r = 0u8;
332            while br.read_bit().unwrap_or(false) {
333                r += 1;
334                if r >= 3 {
335                    break;
336                }
337            }
338            r
339        } else {
340            0
341        };
342
343        // ── repeat frame count (I-frame) ─────────────────────────────────
344        let rptfrm = if frame_type == FrameType::I {
345            br.read_bits(2).unwrap_or(0) as u8
346        } else {
347            0
348        };
349
350        // ── bitplanes ───────────────────────────────────────────────────────
351        // P-frame: skipped-MB bitplane
352        let skipmb_plane = if frame_type == FrameType::P {
353            Bitplane::decode(&mut br, mb_w, mb_h).map(|bp| bp.data)
354        } else {
355            None
356        };
357
358        // B-frame: direct-mode bitplane + skipped-MB bitplane
359        let directmb_plane = if frame_type == FrameType::B {
360            Bitplane::decode(&mut br, mb_w, mb_h).map(|bp| bp.data)
361        } else {
362            None
363        };
364
365        let skipmb_plane = if frame_type == FrameType::B {
366            Bitplane::decode(&mut br, mb_w, mb_h).map(|bp| bp.data)
367        } else {
368            skipmb_plane
369        };
370
371        // ── BFRACTION (B-frames only, SMPTE 421M §7.1.3.6 Table 40) ──────────
372        const BFRAC: [(i32, i32); 8] = [
373            (1, 2),
374            (1, 3),
375            (2, 3),
376            (1, 4),
377            (3, 4),
378            (1, 5),
379            (2, 5),
380            (1, 2),
381        ];
382        let (bfrac_num, bfrac_den) = if frame_type == FrameType::B {
383            let idx = br.read_bits(3).unwrap_or(0) as usize;
384            BFRAC[idx.min(7)]
385        } else {
386            (1, 2)
387        };
388
389        let header_bits = br.bits_read();
390
391        Ok(PictureHeader {
392            frame_type,
393            pqindex,
394            pquant,
395            halfqp,
396            pqual_mode,
397            mvrange,
398            rptfrm,
399            pts_ms,
400            rangeredfrm,
401            header_bits,
402            skipmb_plane,
403            directmb_plane,
404            bfrac_num,
405            bfrac_den,
406        })
407    }
408
409    // ── Legacy parse (no bitplane, backward compat) ─────────────────────────
410    pub fn parse_simple(data: &[u8], seq: &SequenceHeader, pts_ms: u32) -> Result<Self> {
411        let mb_w = ((seq.width + 15) / 16).max(1) as usize;
412        let mb_h = ((seq.height + 15) / 16).max(1) as usize;
413        Self::parse(data, seq, pts_ms, mb_w, mb_h)
414    }
415
416    fn decode_quantizer(pqindex: u8, seq: &SequenceHeader) -> (u8, bool, u8) {
417        match seq.quantizer_mode {
418            QuantizerMode::Implicit => {
419                // SMPTE 421M Table 5
420                let pquant = if pqindex <= 8 {
421                    pqindex
422                } else {
423                    const MAP: [u8; 23] = [
424                        9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21, 22, 23, 24, 25, 27, 29,
425                        31, 33, 63, 0,
426                    ];
427                    MAP.get(pqindex as usize - 9).copied().unwrap_or(pqindex)
428                };
429                let halfqp = pqindex >= 9 && pquant == 0;
430                (pquant, halfqp, 0)
431            }
432            _ => (pqindex, false, 0),
433        }
434    }
435}