Skip to main content

wmv_decoder/wma/
vlc.rs

1//! VLC (Huffman) table builder and decoder.
2//!
3
4use crate::error::{DecoderError, Result};
5use crate::wma::bitstream::GetBitContext;
6
7pub const VLC_INIT_USE_STATIC: i32 = 1;
8pub const VLC_INIT_STATIC_OVERLONG: i32 = 2 | VLC_INIT_USE_STATIC;
9pub const VLC_INIT_INPUT_LE: i32 = 4;
10pub const VLC_INIT_OUTPUT_LE: i32 = 8;
11
12pub type VlcBaseType = i16;
13
14#[derive(Clone, Copy, Default)]
15pub struct VlcElem {
16    pub sym: VlcBaseType,
17    pub len: VlcBaseType,
18}
19
20#[derive(Default)]
21pub struct Vlc {
22    pub bits: i32,
23    pub table: Vec<VlcElem>,
24    pub table_size: i32,
25    pub table_allocated: i32,
26}
27
28#[derive(Clone, Copy)]
29struct VlcCode {
30    bits: u8,
31    symbol: VlcBaseType,
32    /// Codeword with the first bit-to-be-read in the MSB.
33    code: u32,
34}
35
36fn bitswap_32(x: u32) -> u32 {
37    x.reverse_bits()
38}
39
40fn alloc_table(vlc: &mut Vlc, size: i32, use_static: bool) -> Result<i32> {
41    let index = vlc.table_size;
42    vlc.table_size += size;
43    if vlc.table_size > vlc.table_allocated {
44        if use_static {
45            return Err(DecoderError::InvalidData(
46                "static VLC table too small".into(),
47            ));
48        }
49        vlc.table_allocated += 1 << vlc.bits;
50        let new_len = vlc.table_allocated as usize;
51        if new_len > vlc.table.len() {
52            vlc.table.resize(new_len, VlcElem { sym: 0, len: 0 });
53        }
54    }
55    Ok(index)
56}
57
58fn build_table(
59    vlc: &mut Vlc,
60    table_nb_bits: i32,
61    nb_codes: usize,
62    codes: &mut [VlcCode],
63    flags: i32,
64) -> Result<i32> {
65    if table_nb_bits > 30 {
66        return Err(DecoderError::InvalidData("table_nb_bits > 30".into()));
67    }
68    let table_size = 1 << table_nb_bits;
69    let table_index = alloc_table(vlc, table_size, (flags & VLC_INIT_USE_STATIC) != 0)?;
70
71    let base = table_index as usize;
72
73    for i in 0..nb_codes {
74        let mut n = codes[i].bits as i32;
75        let mut code = codes[i].code;
76        let symbol = codes[i].symbol;
77
78        if n <= table_nb_bits {
79            let mut j = (code >> (32 - table_nb_bits)) as i32;
80            let nb = 1 << (table_nb_bits - n);
81            let mut inc = 1;
82            if (flags & VLC_INIT_OUTPUT_LE) != 0 {
83                j = (bitswap_32(code) >> (32 - table_nb_bits)) as i32;
84                inc = 1 << n;
85            }
86            for _k in 0..nb {
87                let idx = base + j as usize;
88                let bits = vlc.table[idx].len;
89                let oldsym = vlc.table[idx].sym;
90                if (bits != 0 || oldsym != 0) && (bits != n as i16 || oldsym != symbol) {
91                    return Err(DecoderError::InvalidData("incorrect VLC codes".into()));
92                }
93                vlc.table[idx].len = n as i16;
94                vlc.table[idx].sym = symbol;
95                j += inc;
96            }
97        } else {
98            // Subtable.
99            n -= table_nb_bits;
100            let code_prefix = code >> (32 - table_nb_bits);
101            let mut subtable_bits = n;
102            codes[i].bits = n as u8;
103            codes[i].code = code << table_nb_bits;
104
105            let mut k = i + 1;
106            while k < nb_codes {
107                let nn = codes[k].bits as i32 - table_nb_bits;
108                if nn <= 0 {
109                    break;
110                }
111                let cc = codes[k].code;
112                if (cc >> (32 - table_nb_bits)) != code_prefix {
113                    break;
114                }
115                codes[k].bits = nn as u8;
116                codes[k].code = cc << table_nb_bits;
117                if nn > subtable_bits {
118                    subtable_bits = nn;
119                }
120                k += 1;
121            }
122            if subtable_bits > table_nb_bits {
123                subtable_bits = table_nb_bits;
124            }
125
126            let j = if (flags & VLC_INIT_OUTPUT_LE) != 0 {
127                (bitswap_32(code_prefix) >> (32 - table_nb_bits)) as i32
128            } else {
129                code_prefix as i32
130            };
131
132            let idx = base + j as usize;
133            vlc.table[idx].len = -(subtable_bits as i16);
134
135            let sub_index = build_table(vlc, subtable_bits, k - i, &mut codes[i..k], flags)?;
136
137            // Rebase after possible resize.
138            let base2 = table_index as usize;
139            let idx2 = base2 + j as usize;
140            vlc.table[idx2].sym = sub_index as i16;
141
142            // Skip processed range.
143            // Equivalent to `i = k - 1` in C loop.
144            // We cannot easily modify `i` in Rust for-loop, so handle via while in caller.
145        }
146    }
147
148    // Mark empty entries.
149    let base3 = table_index as usize;
150    for i in 0..table_size {
151        let idx = base3 + i as usize;
152        if vlc.table[idx].len == 0 {
153            vlc.table[idx].sym = -1;
154        }
155    }
156
157    Ok(table_index)
158}
159
160fn vlc_common_init(vlc: &mut Vlc, nb_bits: i32, flags: i32) {
161    vlc.bits = nb_bits;
162    vlc.table_size = 0;
163    if (flags & VLC_INIT_USE_STATIC) == 0 {
164        vlc.table.clear();
165        vlc.table_allocated = 0;
166    }
167}
168
169fn vlc_common_end(vlc: &mut Vlc, nb_bits: i32, codes: &mut [VlcCode], flags: i32) -> Result<()> {
170    // upstream's build_table expects codes grouped; for sparse init it sorts.
171    // We use a while loop in order to emulate the C for-loop that updates `i`.
172    let nb_codes = codes.len();
173
174    // Build table.
175    // Our build_table implementation above uses recursion but does not update outer loop index.
176    // To preserve upstream semantics, we rebuild using a local recursive builder that uses slices.
177
178    // Re-implement build_table logic with slice recursion, closer to C.
179    fn build(vlc: &mut Vlc, table_nb_bits: i32, codes: &mut [VlcCode], flags: i32) -> Result<i32> {
180        if table_nb_bits > 30 {
181            return Err(DecoderError::InvalidData("table_nb_bits > 30".into()));
182        }
183        let table_size = 1 << table_nb_bits;
184        let table_index = alloc_table(vlc, table_size, (flags & VLC_INIT_USE_STATIC) != 0)?;
185        let mut i: usize = 0;
186        while i < codes.len() {
187            let mut n = codes[i].bits as i32;
188            let mut code = codes[i].code;
189            let symbol = codes[i].symbol;
190
191            let base = table_index as usize;
192
193            if n <= table_nb_bits {
194                let mut j = (code >> (32 - table_nb_bits)) as i32;
195                let nb = 1 << (table_nb_bits - n);
196                let mut inc = 1;
197                if (flags & VLC_INIT_OUTPUT_LE) != 0 {
198                    j = (bitswap_32(code) >> (32 - table_nb_bits)) as i32;
199                    inc = 1 << n;
200                }
201                for _ in 0..nb {
202                    let idx = base + j as usize;
203                    let bits = vlc.table[idx].len;
204                    let oldsym = vlc.table[idx].sym;
205                    if (bits != 0 || oldsym != 0) && (bits != n as i16 || oldsym != symbol) {
206                        return Err(DecoderError::InvalidData("incorrect VLC codes".into()));
207                    }
208                    vlc.table[idx].len = n as i16;
209                    vlc.table[idx].sym = symbol;
210                    j += inc;
211                }
212                i += 1;
213            } else {
214                // Subtable.
215                n -= table_nb_bits;
216                let code_prefix = code >> (32 - table_nb_bits);
217                let mut subtable_bits = n;
218
219                codes[i].bits = n as u8;
220                codes[i].code = code << table_nb_bits;
221
222                let mut k = i + 1;
223                while k < codes.len() {
224                    let nn = codes[k].bits as i32 - table_nb_bits;
225                    if nn <= 0 {
226                        break;
227                    }
228                    let cc = codes[k].code;
229                    if (cc >> (32 - table_nb_bits)) != code_prefix {
230                        break;
231                    }
232                    codes[k].bits = nn as u8;
233                    codes[k].code = cc << table_nb_bits;
234                    if nn > subtable_bits {
235                        subtable_bits = nn;
236                    }
237                    k += 1;
238                }
239
240                if subtable_bits > table_nb_bits {
241                    subtable_bits = table_nb_bits;
242                }
243
244                let j = if (flags & VLC_INIT_OUTPUT_LE) != 0 {
245                    (bitswap_32(code_prefix) >> (32 - table_nb_bits)) as i32
246                } else {
247                    code_prefix as i32
248                };
249
250                {
251                    let idx = base + j as usize;
252                    vlc.table[idx].len = -(subtable_bits as i16);
253                }
254
255                let sub_index = build(vlc, subtable_bits, &mut codes[i..k], flags)?;
256
257                // Reload base after possible resize.
258                let base2 = table_index as usize;
259                let idx2 = base2 + j as usize;
260                vlc.table[idx2].sym = sub_index as i16;
261
262                i = k;
263            }
264        }
265
266        // Mark empty.
267        let base = table_index as usize;
268        for t in 0..table_size {
269            let idx = base + t as usize;
270            if vlc.table[idx].len == 0 {
271                vlc.table[idx].sym = -1;
272            }
273        }
274
275        Ok(table_index)
276    }
277
278    build(vlc, nb_bits, codes, flags)?;
279
280    if (flags & VLC_INIT_USE_STATIC) != 0 {
281        // Nothing.
282        let _ = nb_codes;
283    }
284
285    Ok(())
286}
287
288fn get_data_u32(table: &[u8], wrap: i32, i: usize, size: i32) -> u32 {
289    let off = i * wrap as usize;
290    match size {
291        1 => table[off] as u32,
292        2 => u16::from_ne_bytes([table[off], table[off + 1]]) as u32,
293        4 => u32::from_ne_bytes([table[off], table[off + 1], table[off + 2], table[off + 3]]),
294        _ => 0,
295    }
296}
297
298fn get_data_u16(table: &[u8], wrap: i32, i: usize, size: i32) -> u16 {
299    let off = i * wrap as usize;
300    match size {
301        1 => table[off] as u16,
302        2 => u16::from_ne_bytes([table[off], table[off + 1]]),
303        _ => 0,
304    }
305}
306
307/// Equivalent to upstream `ff_vlc_init_sparse()`.
308#[allow(clippy::too_many_arguments)]
309pub fn ff_vlc_init_sparse(
310    vlc: &mut Vlc,
311    nb_bits: i32,
312    nb_codes: usize,
313    bits: &[u8],
314    bits_wrap: i32,
315    bits_size: i32,
316    codes: &[u8],
317    codes_wrap: i32,
318    codes_size: i32,
319    symbols: Option<&[u8]>,
320    symbols_wrap: i32,
321    symbols_size: i32,
322    flags: i32,
323) -> Result<()> {
324    vlc_common_init(vlc, nb_bits, flags);
325
326    let mut buf: Vec<VlcCode> = Vec::with_capacity(nb_codes);
327
328    // Copy entries with len > nb_bits first.
329    for pass in 0..2 {
330        for i in 0..nb_codes {
331            let len = get_data_u32(bits, bits_wrap, i, bits_size) as u32;
332            let cond = if pass == 0 {
333                len > nb_bits as u32
334            } else {
335                len != 0 && len <= nb_bits as u32
336            };
337            if !cond {
338                continue;
339            }
340            if len > (3 * nb_bits) as u32 || len > 32 {
341                return Err(DecoderError::InvalidData(format!("Too long VLC ({len})")));
342            }
343            let mut code = get_data_u32(codes, codes_wrap, i, codes_size);
344            if code as u64 >= (1u64 << len) {
345                return Err(DecoderError::InvalidData(format!(
346                    "Invalid code {code:x} for {i}"
347                )));
348            }
349            if (flags & VLC_INIT_INPUT_LE) != 0 {
350                code = bitswap_32(code);
351            } else {
352                code <<= 32 - len;
353            }
354            let sym: i16 = if let Some(symtab) = symbols {
355                get_data_u16(symtab, symbols_wrap, i, symbols_size) as i16
356            } else {
357                i as i16
358            };
359            buf.push(VlcCode {
360                bits: len as u8,
361                symbol: sym,
362                code,
363            });
364        }
365        if pass == 0 {
366            buf.sort_by_key(|c| c.code >> 1);
367        }
368    }
369
370    vlc_common_end(vlc, nb_bits, &mut buf, flags)
371}
372
373/// Equivalent to upstream `ff_vlc_init_from_lengths()`.
374#[allow(clippy::too_many_arguments)]
375pub fn ff_vlc_init_from_lengths(
376    vlc: &mut Vlc,
377    nb_bits: i32,
378    nb_codes: usize,
379    lens: &[i8],
380    lens_wrap: i32,
381    symbols: Option<&[u8]>,
382    symbols_wrap: i32,
383    symbols_size: i32,
384    offset: i32,
385    flags: i32,
386) -> Result<()> {
387    vlc_common_init(vlc, nb_bits, flags);
388
389    let mut buf: Vec<VlcCode> = Vec::with_capacity(nb_codes);
390    let mut code: u64 = 0;
391    let len_max: i32 = 32.min(3 * nb_bits);
392
393    for i in 0..nb_codes {
394        let len = lens[(i * lens_wrap as usize)] as i32;
395        if len > 0 {
396            let sym_u = if let Some(symtab) = symbols {
397                get_data_u16(symtab, symbols_wrap, i, symbols_size) as u32
398            } else {
399                i as u32
400            };
401            let sym = (sym_u as i32 + offset) as i16;
402            buf.push(VlcCode {
403                bits: len as u8,
404                symbol: sym,
405                code: code as u32,
406            });
407        } else if len < 0 {
408            // Incomplete tree marker.
409        } else {
410            continue;
411        }
412
413        let mut abs_len = len;
414        if abs_len < 0 {
415            abs_len = -abs_len;
416        }
417        if abs_len > len_max || (code & ((1u64 << (32 - abs_len)) - 1)) != 0 {
418            return Err(DecoderError::InvalidData(format!(
419                "Invalid VLC (length {abs_len})"
420            )));
421        }
422        code += 1u64 << (32 - abs_len);
423        if code > (u32::MAX as u64) + 1 {
424            return Err(DecoderError::InvalidData("Overdetermined VLC tree".into()));
425        }
426    }
427
428    vlc_common_end(vlc, nb_bits, &mut buf, flags)
429}
430
431/// Equivalent to `get_vlc2()`.
432#[inline]
433pub fn get_vlc2(
434    gb: &mut GetBitContext<'_>,
435    table: &[VlcElem],
436    bits: i32,
437    max_depth: i32,
438) -> Result<i32> {
439    let mut code: i32;
440    let mut index = gb.show_bits(bits as usize)? as usize;
441    let mut n = table[index].len as i32;
442    code = table[index].sym as i32;
443
444    if max_depth > 1 && n < 0 {
445        gb.skip_bits(bits as usize)?;
446        let mut nb_bits = -n;
447        index = (gb.show_bits(nb_bits as usize)? as usize) + code as usize;
448        n = table[index].len as i32;
449        code = table[index].sym as i32;
450        if max_depth > 2 && n < 0 {
451            gb.skip_bits(nb_bits as usize)?;
452            nb_bits = -n;
453            index = (gb.show_bits(nb_bits as usize)? as usize) + code as usize;
454            n = table[index].len as i32;
455            code = table[index].sym as i32;
456        }
457    }
458
459    gb.skip_bits(n as usize)?;
460    Ok(code)
461}