Skip to main content

wmv_decoder/wma/
decoder.rs

1use std::f32;
2
3use crate::asf::AudioStreamInfo;
4use crate::error::{DecoderError, Result};
5
6use super::bitstream::GetBitContext;
7use super::common::ff_wma_get_frame_len_bits;
8use super::mdct::MdctNaive;
9use super::tables;
10use super::vlc::{ff_vlc_init_from_lengths, ff_vlc_init_sparse, get_vlc2, Vlc, VlcElem};
11
12const BLOCK_MIN_BITS: i32 = 7;
13const BLOCK_MAX_BITS: i32 = 11;
14const BLOCK_MAX_SIZE: usize = 1 << BLOCK_MAX_BITS;
15const BLOCK_NB_SIZES: usize = (BLOCK_MAX_BITS - BLOCK_MIN_BITS + 1) as usize;
16
17const HIGH_BAND_MAX_SIZE: usize = 16;
18const NB_LSP_COEFS: usize = 10;
19
20const MAX_CODED_SUPERFRAME_SIZE: usize = 32768;
21const MAX_CHANNELS: usize = 2;
22
23const NOISE_TAB_SIZE: usize = 8192;
24const LSP_POW_BITS: usize = 7;
25
26const VLCBITS: i32 = 9;
27const VLCMAX: i32 = (22 + VLCBITS - 1) / VLCBITS;
28
29const EXPVLCBITS: i32 = 8;
30const EXPMAX: i32 = (19 + EXPVLCBITS - 1) / EXPVLCBITS;
31
32const HGAINVLCBITS: i32 = 9;
33const HGAINMAX: i32 = (13 + HGAINVLCBITS - 1) / HGAINVLCBITS;
34
35/// A decoded PCM chunk.
36#[derive(Debug, Clone)]
37pub struct PcmFrameF32 {
38    pub pts_ms: u32,
39    pub sample_rate: u32,
40    pub channels: u16,
41    /// Interleaved samples.
42    pub samples: Vec<f32>,
43}
44
45#[derive(Clone, Copy, Debug)]
46enum WmaVersion {
47    V1,
48    V2,
49}
50
51impl WmaVersion {
52    fn id(&self) -> i32 {
53        match self {
54            WmaVersion::V1 => 1,
55            WmaVersion::V2 => 2,
56        }
57    }
58}
59
60/// Direct translation of upstream `WMACodecContext` for WMAv1/2.
61pub struct WmaDecoder {
62    version: WmaVersion,
63
64    channels: usize,
65    sample_rate: u32,
66    bit_rate: u32,
67    block_align: u16,
68
69    // Flags derived from `flags2`.
70    use_exp_vlc: bool,
71    use_bit_reservoir: bool,
72    use_variable_block_len: bool,
73    use_noise_coding: bool,
74
75    byte_offset_bits: i32,
76
77    // VLC tables.
78    exp_vlc: Vlc,
79    hgain_vlc: Vlc,
80    coef_vlc: [Vlc; 2],
81    run_table: [Vec<u16>; 2],
82    level_table: [Vec<f32>; 2],
83
84    // Frame / block config.
85    frame_len_bits: i32,
86    frame_len: usize,
87    nb_block_sizes: usize,
88
89    reset_block_lengths: bool,
90    block_len_bits: i32,
91    next_block_len_bits: i32,
92    prev_block_len_bits: i32,
93    block_len: usize,
94    block_num: i32,
95    block_pos: usize,
96
97    ms_stereo: bool,
98    channel_coded: [bool; MAX_CHANNELS],
99
100    // Exponent bands.
101    exponent_sizes: [usize; BLOCK_NB_SIZES],
102    exponent_bands: [[u16; 25]; BLOCK_NB_SIZES],
103    high_band_start: [usize; BLOCK_NB_SIZES],
104    coefs_start: usize,
105    coefs_end: [usize; BLOCK_NB_SIZES],
106    exponent_high_sizes: [usize; BLOCK_NB_SIZES],
107    exponent_high_bands: [[u16; HIGH_BAND_MAX_SIZE]; BLOCK_NB_SIZES],
108
109    high_band_coded: [[bool; HIGH_BAND_MAX_SIZE]; MAX_CHANNELS],
110    high_band_values: [[i32; HIGH_BAND_MAX_SIZE]; MAX_CHANNELS],
111
112    // Exponents and coefficients.
113    exponents_bsize: [usize; MAX_CHANNELS],
114    exponents: [Vec<f32>; MAX_CHANNELS],
115    max_exponent: [f32; MAX_CHANNELS],
116    coefs1: [Vec<f32>; MAX_CHANNELS],
117    coefs: [Vec<f32>; MAX_CHANNELS],
118
119    // MDCT.
120    mdct: Vec<MdctNaive>,
121    windows: Vec<Vec<f32>>, // per block-size: half window of length block_len
122    output: Vec<f32>,       // 2*BLOCK_MAX_SIZE
123    frame_out: [Vec<f32>; MAX_CHANNELS],
124
125    // Bit reservoir.
126    last_superframe: Vec<u8>,
127    last_bitoffset: usize,
128    last_superframe_len: usize,
129    eof_done: bool,
130
131    // Noise.
132    noise_table: Vec<f32>,
133    noise_index: usize,
134    noise_mult: f32,
135
136    // LSP to curve.
137    lsp_cos_table: Vec<f32>,
138    lsp_pow_e_table: [f32; 256],
139    lsp_pow_m_table1: [f32; 1 << LSP_POW_BITS],
140    lsp_pow_m_table2: [f32; 1 << LSP_POW_BITS],
141
142    exponents_initialized: [bool; MAX_CHANNELS],
143}
144
145fn ilog2_u32(x: u32) -> i32 {
146    31 - (x.leading_zeros() as i32)
147}
148
149fn ff_exp10f(x: f32) -> f32 {
150    // ff_exp10f(x) = exp2f(M_LOG2_10 * x)
151    (std::f32::consts::LOG2_10 * x).exp2()
152}
153
154fn sine_window_init(n: usize) -> Vec<f32> {
155    // Translated from ff_sine_window_init.
156    let mut w = vec![0f32; n];
157    let den = 2.0f32 * n as f32;
158    for i in 0..n {
159        w[i] = ((i as f32 + 0.5) * (std::f32::consts::PI / den)).sin();
160    }
161    w
162}
163
164fn vector_fmul_reverse(dst: &mut [f32], src0: &[f32], win: &[f32]) {
165    let len = dst.len();
166    for i in 0..len {
167        dst[i] = src0[i] * win[len - 1 - i];
168    }
169}
170
171fn butterflies_float(v1: &mut [f32], v2: &mut [f32]) {
172    for i in 0..v1.len() {
173        let t = v1[i] - v2[i];
174        v1[i] += v2[i];
175        v2[i] = t;
176    }
177}
178
179fn pow_m1_4_tables(
180    x: f32,
181    lsp_pow_e_table: &[f32; 256],
182    lsp_pow_m_table1: &[f32; 1 << LSP_POW_BITS],
183    lsp_pow_m_table2: &[f32; 1 << LSP_POW_BITS],
184) -> f32 {
185    // Direct translation of `pow_m1_4` from upstream wmadec.c, but parameterized to avoid borrowing `self`.
186    let u = x.to_bits();
187    let e = (u >> 23) as usize;
188    let m = ((u >> (23 - LSP_POW_BITS)) & ((1 << LSP_POW_BITS) - 1) as u32) as usize;
189    let t_bits = ((u << LSP_POW_BITS) & ((1 << 23) - 1)) | (127 << 23);
190    let t = f32::from_bits(t_bits);
191    let a = lsp_pow_m_table1[m];
192    let b = lsp_pow_m_table2[m];
193    lsp_pow_e_table[e] * (a + b * t)
194}
195
196fn wma_lsp_to_curve_tables(
197    out: &mut [f32],
198    n: usize,
199    lsp: &[f32; NB_LSP_COEFS],
200    lsp_cos_table: &[f32],
201    lsp_pow_e_table: &[f32; 256],
202    lsp_pow_m_table1: &[f32; 1 << LSP_POW_BITS],
203    lsp_pow_m_table2: &[f32; 1 << LSP_POW_BITS],
204) -> f32 {
205    // Direct translation of `wma_lsp_to_curve` from upstream wmadec.c, parameterized to avoid borrowing `self`.
206    let mut val_max = 0.0f32;
207    for i in 0..n {
208        let mut p = 0.5f32;
209        let mut q = 0.5f32;
210        let w = lsp_cos_table[i];
211        let mut j = 1usize;
212        while j < NB_LSP_COEFS {
213            q *= w - lsp[j - 1];
214            p *= w - lsp[j];
215            j += 2;
216        }
217        p *= p * (2.0f32 - w);
218        q *= q * (2.0f32 + w);
219        let mut v = p + q;
220        v = pow_m1_4_tables(v, lsp_pow_e_table, lsp_pow_m_table1, lsp_pow_m_table2);
221        if v > val_max {
222            val_max = v;
223        }
224        out[i] = v;
225    }
226    val_max
227}
228
229fn wma_window_apply(
230    out: &mut [f32],
231    output: &[f32],
232    windows: &[Vec<f32>],
233    frame_len_bits: i32,
234    block_len_bits: i32,
235    prev_block_len_bits: i32,
236    next_block_len_bits: i32,
237    block_len: usize,
238) {
239    // Direct translation of upstream `wma_window`, but parameterized to avoid borrowing `self`.
240    let mut in_buf: &[f32] = output;
241
242    // Left part.
243    if block_len_bits <= prev_block_len_bits {
244        let bsize = (frame_len_bits - block_len_bits) as usize;
245        let win = &windows[bsize];
246        for i in 0..block_len {
247            out[i] = in_buf[i] * win[i] + out[i];
248        }
249    } else {
250        let prev_len = 1usize << prev_block_len_bits;
251        let n = (block_len - prev_len) / 2;
252        let bsize = (frame_len_bits - prev_block_len_bits) as usize;
253        let win = &windows[bsize];
254        for i in 0..prev_len {
255            let idx = n + i;
256            out[idx] = in_buf[idx] * win[i] + out[idx];
257        }
258        out[n + prev_len..n + prev_len + n]
259            .copy_from_slice(&in_buf[n + prev_len..n + prev_len + n]);
260    }
261
262    // Right part.
263    let out2 = &mut out[block_len..];
264    in_buf = &in_buf[block_len..];
265
266    if block_len_bits <= next_block_len_bits {
267        let bsize = (frame_len_bits - block_len_bits) as usize;
268        vector_fmul_reverse(
269            &mut out2[..block_len],
270            &in_buf[..block_len],
271            &windows[bsize],
272        );
273    } else {
274        let next_len = 1usize << next_block_len_bits;
275        let n = (block_len - next_len) / 2;
276        let bsize = (frame_len_bits - next_block_len_bits) as usize;
277        out2[n + next_len..n + next_len + n]
278            .copy_from_slice(&in_buf[n + next_len..n + next_len + n]);
279        vector_fmul_reverse(
280            &mut out2[n..n + next_len],
281            &in_buf[n..n + next_len],
282            &windows[bsize],
283        );
284    }
285}
286
287impl WmaDecoder {
288    pub fn new(info: &AudioStreamInfo) -> Result<Self> {
289        let version = match info.format_tag {
290            0x0160 => WmaVersion::V1,
291            0x0161 => WmaVersion::V2,
292            _ => {
293                return Err(DecoderError::Unsupported(format!(
294                    "unsupported WMA format tag: 0x{:04x}",
295                    info.format_tag
296                )))
297            }
298        };
299
300        if info.block_align == 0 {
301            return Err(DecoderError::InvalidData("block_align is not set".into()));
302        }
303
304        let channels = info.channels as usize;
305        if channels == 0 || channels > MAX_CHANNELS {
306            return Err(DecoderError::Unsupported(
307                "only mono/stereo supported".into(),
308            ));
309        }
310
311        // Extract flags2 like upstream.
312        let mut flags2: u16 = 0;
313        let extradata = &info.extra_data;
314        match version {
315            WmaVersion::V1 => {
316                if extradata.len() >= 4 {
317                    flags2 = u16::from_le_bytes([extradata[2], extradata[3]]);
318                }
319            }
320            WmaVersion::V2 => {
321                if extradata.len() >= 6 {
322                    flags2 = u16::from_le_bytes([extradata[4], extradata[5]]);
323                }
324            }
325        }
326
327        let mut use_variable_block_len = (flags2 & 0x0004) != 0;
328        let use_exp_vlc = (flags2 & 0x0001) != 0;
329        let use_bit_reservoir = (flags2 & 0x0002) != 0;
330
331        // upstream quirk (issue1503).
332        if let WmaVersion::V2 = version {
333            if extradata.len() >= 8 {
334                let v = u16::from_le_bytes([extradata[4], extradata[5]]);
335                if v == 0x000d && use_variable_block_len {
336                    use_variable_block_len = false;
337                }
338            }
339        }
340
341        // Pre-init fixed fields.
342        let mut dec = Self {
343            version,
344            channels,
345            sample_rate: info.sample_rate,
346            bit_rate: info.bit_rate,
347            block_align: info.block_align,
348
349            use_exp_vlc,
350            use_bit_reservoir,
351            use_variable_block_len,
352            use_noise_coding: true,
353
354            byte_offset_bits: 0,
355
356            exp_vlc: Vlc::default(),
357            hgain_vlc: Vlc::default(),
358            coef_vlc: [Vlc::default(), Vlc::default()],
359            run_table: [Vec::new(), Vec::new()],
360            level_table: [Vec::new(), Vec::new()],
361
362            frame_len_bits: 0,
363            frame_len: 0,
364            nb_block_sizes: 0,
365
366            reset_block_lengths: true,
367            block_len_bits: 0,
368            next_block_len_bits: 0,
369            prev_block_len_bits: 0,
370            block_len: 0,
371            block_num: 0,
372            block_pos: 0,
373
374            ms_stereo: false,
375            channel_coded: [false; MAX_CHANNELS],
376
377            exponent_sizes: [0usize; BLOCK_NB_SIZES],
378            exponent_bands: [[0u16; 25]; BLOCK_NB_SIZES],
379            high_band_start: [0usize; BLOCK_NB_SIZES],
380            coefs_start: 0,
381            coefs_end: [0usize; BLOCK_NB_SIZES],
382            exponent_high_sizes: [0usize; BLOCK_NB_SIZES],
383            exponent_high_bands: [[0u16; HIGH_BAND_MAX_SIZE]; BLOCK_NB_SIZES],
384
385            high_band_coded: [[false; HIGH_BAND_MAX_SIZE]; MAX_CHANNELS],
386            high_band_values: [[0i32; HIGH_BAND_MAX_SIZE]; MAX_CHANNELS],
387
388            exponents_bsize: [0usize; MAX_CHANNELS],
389            exponents: [vec![0f32; BLOCK_MAX_SIZE], vec![0f32; BLOCK_MAX_SIZE]],
390            max_exponent: [1.0f32; MAX_CHANNELS],
391            coefs1: [vec![0f32; BLOCK_MAX_SIZE], vec![0f32; BLOCK_MAX_SIZE]],
392            coefs: [vec![0f32; BLOCK_MAX_SIZE], vec![0f32; BLOCK_MAX_SIZE]],
393
394            mdct: Vec::new(),
395            windows: Vec::new(),
396            output: vec![0f32; BLOCK_MAX_SIZE * 2],
397            frame_out: [
398                vec![0f32; BLOCK_MAX_SIZE * 2],
399                vec![0f32; BLOCK_MAX_SIZE * 2],
400            ],
401
402            last_superframe: vec![0u8; MAX_CODED_SUPERFRAME_SIZE + 64],
403            last_bitoffset: 0,
404            last_superframe_len: 0,
405            eof_done: false,
406
407            noise_table: vec![0f32; NOISE_TAB_SIZE],
408            noise_index: 0,
409            noise_mult: 0.0,
410
411            lsp_cos_table: vec![0f32; BLOCK_MAX_SIZE],
412            lsp_pow_e_table: [0f32; 256],
413            lsp_pow_m_table1: [0f32; 1 << LSP_POW_BITS],
414            lsp_pow_m_table2: [0f32; 1 << LSP_POW_BITS],
415
416            exponents_initialized: [false; MAX_CHANNELS],
417        };
418
419        // Full init = ff_wma_init + wma_decode_init bits.
420        dec.ff_wma_init(flags2 as i32)?;
421        dec.wma_decode_init(flags2 as i32)?;
422
423        Ok(dec)
424    }
425
426    pub fn sample_rate(&self) -> u32 {
427        self.sample_rate
428    }
429
430    pub fn channels(&self) -> u16 {
431        self.channels as u16
432    }
433
434    pub fn frame_len(&self) -> usize {
435        self.frame_len
436    }
437
438    /// Decode one ASF packet payload (usually `block_align` bytes).
439    pub fn decode_packet(&mut self, pkt: &[u8], pts_ms: u32) -> Result<Option<PcmFrameF32>> {
440        if pkt.is_empty() {
441            if self.eof_done {
442                return Ok(None);
443            }
444            // Flush delayed samples.
445            self.eof_done = true;
446            let mut out = Vec::with_capacity(self.frame_len * self.channels);
447            for i in 0..self.frame_len {
448                for ch in 0..self.channels {
449                    out.push(self.frame_out[ch][i]);
450                }
451            }
452            self.last_superframe_len = 0;
453            return Ok(Some(PcmFrameF32 {
454                pts_ms,
455                sample_rate: self.sample_rate,
456                channels: self.channels as u16,
457                samples: out,
458            }));
459        }
460
461        if pkt.len() < self.block_align as usize {
462            return Err(DecoderError::InvalidData(format!(
463                "Input packet size too small ({} < {})",
464                pkt.len(),
465                self.block_align
466            )));
467        }
468
469        let buf = &pkt[..self.block_align as usize];
470
471        let mut gb = GetBitContext::new(buf);
472
473        let mut nb_frames: i32;
474
475        if self.use_bit_reservoir {
476            // super frame header
477            gb.skip_bits(4)?; // super frame index
478            let mut nf = gb.get_bits(4)? as i32;
479            nf -= if self.last_superframe_len <= 0 { 1 } else { 0 };
480            nb_frames = nf;
481            if nb_frames <= 0 {
482                let is_error = nb_frames < 0 || gb.bits_left() <= 8;
483                if is_error {
484                    return Err(DecoderError::InvalidData(format!(
485                        "nb_frames is {nb_frames} bits left {}",
486                        gb.bits_left()
487                    )));
488                }
489
490                if self.last_superframe_len + buf.len() - 1 > MAX_CODED_SUPERFRAME_SIZE {
491                    return Err(DecoderError::InvalidData("bit reservoir overflow".into()));
492                }
493
494                let mut q = self.last_superframe_len;
495                let mut len = buf.len() - 1;
496                while len > 0 {
497                    let b = gb.get_bits(8)? as u8;
498                    self.last_superframe[q] = b;
499                    q += 1;
500                    len -= 1;
501                }
502
503                self.last_superframe_len += 8 * buf.len() - 8;
504                return Ok(None);
505            }
506        } else {
507            nb_frames = 1;
508        }
509
510        // Planar output like upstream, then interleave.
511        let mut samples: [Vec<f32>; MAX_CHANNELS] = [Vec::new(), Vec::new()];
512        for ch in 0..self.channels {
513            samples[ch].resize(nb_frames as usize * self.frame_len, 0f32);
514        }
515        let mut samples_offset: usize = 0;
516
517        if self.use_bit_reservoir {
518            let bit_offset = gb.get_bits((self.byte_offset_bits + 3) as usize)? as usize;
519            if bit_offset as isize > gb.bits_left() {
520                return Err(DecoderError::InvalidData(
521                    "Invalid last frame bit offset".into(),
522                ));
523            }
524
525            if self.last_superframe_len > 0 {
526                // Add `bit_offset` bits to last frame.
527                let add_bytes = (bit_offset + 7) >> 3;
528                if self.last_superframe_len + add_bytes > MAX_CODED_SUPERFRAME_SIZE {
529                    return Err(DecoderError::InvalidData("bit reservoir overflow".into()));
530                }
531
532                let mut q = self.last_superframe_len;
533                let mut len = bit_offset;
534                while len > 7 {
535                    self.last_superframe[q] = gb.get_bits(8)? as u8;
536                    q += 1;
537                    len -= 8;
538                }
539                if len > 0 {
540                    self.last_superframe[q] = (gb.get_bits(len)? as u8) << (8 - len);
541                }
542
543                // Decode the previous frame.
544                let total_bits = self.last_superframe_len * 8 + bit_offset;
545                let need_bytes = (total_bits + 7) / 8;
546                // Avoid borrowing `self` across the decode call.
547                let sf_bytes: Vec<u8> = self.last_superframe[..need_bytes].to_vec();
548                let mut gb2 = GetBitContext::new(&sf_bytes);
549                if self.last_bitoffset > 0 {
550                    gb2.skip_bits(self.last_bitoffset)?;
551                }
552                self.reset_block_lengths = true;
553                self.wma_decode_frame(&mut gb2, &mut samples, samples_offset)?;
554                samples_offset += self.frame_len;
555                nb_frames -= 1;
556            }
557
558            // Read each frame starting from bit_offset.
559            let pos = bit_offset + 4 + 4 + (self.byte_offset_bits as usize) + 3;
560            if pos >= MAX_CODED_SUPERFRAME_SIZE * 8 || pos > buf.len() * 8 {
561                return Err(DecoderError::InvalidData("invalid superframe pos".into()));
562            }
563
564            let start_byte = pos >> 3;
565            let mut gb3 = GetBitContext::new(&buf[start_byte..]);
566            let rem = pos & 7;
567            if rem > 0 {
568                gb3.skip_bits(rem)?;
569            }
570
571            self.reset_block_lengths = true;
572            for _ in 0..nb_frames {
573                self.wma_decode_frame(&mut gb3, &mut samples, samples_offset)?;
574                samples_offset += self.frame_len;
575            }
576
577            // Copy end of frame into last frame buffer.
578            let consumed_bits = gb3.bits_read();
579            let mut pos2 =
580                consumed_bits + ((bit_offset + 4 + 4 + (self.byte_offset_bits as usize) + 3) & !7);
581            self.last_bitoffset = pos2 & 7;
582            pos2 >>= 3;
583            let len = buf.len().saturating_sub(pos2);
584            if len > MAX_CODED_SUPERFRAME_SIZE {
585                return Err(DecoderError::InvalidData("invalid reservoir len".into()));
586            }
587            self.last_superframe_len = len;
588            self.last_superframe[..len].copy_from_slice(&buf[pos2..pos2 + len]);
589        } else {
590            self.reset_block_lengths = true;
591            self.wma_decode_frame(&mut gb, &mut samples, samples_offset)?;
592            samples_offset += self.frame_len;
593        }
594
595        // Interleave.
596        let total_samples = samples_offset * self.channels;
597        let mut out = Vec::with_capacity(total_samples);
598        for i in 0..samples_offset {
599            for ch in 0..self.channels {
600                out.push(samples[ch][i]);
601            }
602        }
603
604        Ok(Some(PcmFrameF32 {
605            pts_ms,
606            sample_rate: self.sample_rate,
607            channels: self.channels as u16,
608            samples: out,
609        }))
610    }
611
612    fn wma_decode_init(&mut self, flags2: i32) -> Result<()> {
613        // Initialize MDCT contexts (naive) like wma_decode_init.
614        let scale = 1.0f64 / 32768.0f64;
615        self.mdct.clear();
616        for i in 0..self.nb_block_sizes {
617            let len = 1usize << (self.frame_len_bits - i as i32);
618            self.mdct.push(MdctNaive::new(len, scale));
619        }
620
621        // Noise/hgain VLC.
622        if self.use_noise_coding {
623            let flat: &[u8] = unsafe {
624                std::slice::from_raw_parts(
625                    tables::FF_WMA_HGAIN_HUFFTAB.as_ptr() as *const u8,
626                    tables::FF_WMA_HGAIN_HUFFTAB.len() * 2,
627                )
628            };
629            let lens: &[i8] = unsafe {
630                std::slice::from_raw_parts(flat.as_ptr().add(1) as *const i8, flat.len() - 1)
631            };
632            ff_vlc_init_from_lengths(
633                &mut self.hgain_vlc,
634                HGAINVLCBITS,
635                tables::FF_WMA_HGAIN_HUFFTAB.len(),
636                lens,
637                2,
638                Some(flat),
639                2,
640                1,
641                -18,
642                0,
643            )?;
644        }
645
646        // Exponent VLC.
647        if self.use_exp_vlc {
648            let bits = &tables::FF_AAC_SCALEFACTOR_BITS;
649            let codes_u32 = &tables::FF_AAC_SCALEFACTOR_CODE;
650            let codes_bytes: &[u8] = unsafe {
651                std::slice::from_raw_parts(codes_u32.as_ptr() as *const u8, codes_u32.len() * 4)
652            };
653
654            ff_vlc_init_sparse(
655                &mut self.exp_vlc,
656                EXPVLCBITS,
657                bits.len(),
658                bits,
659                1,
660                1,
661                codes_bytes,
662                4,
663                4,
664                None,
665                0,
666                0,
667                0,
668            )?;
669        } else {
670            self.wma_lsp_to_curve_init(self.frame_len);
671        }
672
673        // Flags and defaults.
674        let _ = flags2;
675        Ok(())
676    }
677
678    fn ff_wma_init(&mut self, flags2: i32) -> Result<()> {
679        // Validate stream params.
680        if self.sample_rate > 50000 || self.channels > 2 || self.bit_rate == 0 {
681            return Err(DecoderError::InvalidData("invalid audio params".into()));
682        }
683
684        let version_id = self.version.id();
685
686        // Compute MDCT block size.
687        self.frame_len_bits = ff_wma_get_frame_len_bits(self.sample_rate as i32, version_id, 0);
688        self.next_block_len_bits = self.frame_len_bits;
689        self.prev_block_len_bits = self.frame_len_bits;
690        self.block_len_bits = self.frame_len_bits;
691
692        self.frame_len = 1usize << self.frame_len_bits;
693        if self.use_variable_block_len {
694            let mut nb = ((flags2 >> 3) & 3) + 1;
695            if (self.bit_rate / self.channels as u32) >= 32000 {
696                nb += 2;
697            }
698            let nb_max = self.frame_len_bits - BLOCK_MIN_BITS;
699            if nb > nb_max {
700                nb = nb_max;
701            }
702            self.nb_block_sizes = (nb + 1) as usize;
703        } else {
704            self.nb_block_sizes = 1;
705        }
706
707        // Rate dependent params.
708        self.use_noise_coding = true;
709        let mut high_freq = self.sample_rate as f32 * 0.5f32;
710
711        // Version 2 normalized rates.
712        let mut sample_rate1 = self.sample_rate as i32;
713        if version_id == 2 {
714            if sample_rate1 >= 44100 {
715                sample_rate1 = 44100;
716            } else if sample_rate1 >= 22050 {
717                sample_rate1 = 22050;
718            } else if sample_rate1 >= 16000 {
719                sample_rate1 = 16000;
720            } else if sample_rate1 >= 11025 {
721                sample_rate1 = 11025;
722            } else if sample_rate1 >= 8000 {
723                sample_rate1 = 8000;
724            }
725        }
726
727        let bps = (self.bit_rate as f32) / ((self.channels as f32) * (self.sample_rate as f32));
728        let mut bps1 = bps;
729        if self.channels == 2 {
730            bps1 = bps * 1.6f32;
731        }
732
733        let x = (bps * (self.frame_len as f32) / 8.0 + 0.5) as u32;
734        self.byte_offset_bits = ilog2_u32(x.max(1)) + 2;
735
736        // Compute high frequency and noise coding.
737        if sample_rate1 == 44100 {
738            if bps1 >= 0.61 {
739                self.use_noise_coding = false;
740            } else {
741                high_freq *= 0.4;
742            }
743        } else if sample_rate1 == 22050 {
744            if bps1 >= 1.16 {
745                self.use_noise_coding = false;
746            } else if bps1 >= 0.72 {
747                high_freq *= 0.7;
748            } else {
749                high_freq *= 0.6;
750            }
751        } else if sample_rate1 == 16000 {
752            if bps > 0.5 {
753                high_freq *= 0.5;
754            } else {
755                high_freq *= 0.3;
756            }
757        } else if sample_rate1 == 11025 {
758            high_freq *= 0.7;
759        } else if sample_rate1 == 8000 {
760            if bps <= 0.625 {
761                high_freq *= 0.5;
762            } else if bps > 0.75 {
763                self.use_noise_coding = false;
764            } else {
765                high_freq *= 0.65;
766            }
767        } else {
768            if bps >= 0.8 {
769                high_freq *= 0.75;
770            } else if bps >= 0.6 {
771                high_freq *= 0.6;
772            } else {
773                high_freq *= 0.5;
774            }
775        }
776
777        // Compute scale factor band sizes.
778        self.coefs_start = if version_id == 1 { 3 } else { 0 };
779
780        for k in 0..self.nb_block_sizes {
781            let block_len = self.frame_len >> k;
782
783            if version_id == 1 {
784                let mut lpos = 0usize;
785                let mut i = 0usize;
786                for idx in 0..25 {
787                    let a = tables::FF_WMA_CRITICAL_FREQS[idx] as usize;
788                    let b = self.sample_rate as usize;
789                    let mut pos = ((block_len * 2 * a) + (b >> 1)) / b;
790                    if pos > block_len {
791                        pos = block_len;
792                    }
793                    self.exponent_bands[0][idx] = (pos - lpos) as u16;
794                    if pos >= block_len {
795                        i = idx + 1;
796                        break;
797                    }
798                    lpos = pos;
799                    i = idx + 1;
800                }
801                self.exponent_sizes[0] = i;
802            } else {
803                // Hardcoded tables.
804                let a = self.frame_len_bits - BLOCK_MIN_BITS - (k as i32);
805                let mut table_row: Option<&[u8; 25]> = None;
806                if a < 3 {
807                    if self.sample_rate >= 44100 {
808                        table_row = Some(&tables::EXPONENT_BAND_44100[a as usize]);
809                    } else if self.sample_rate >= 32000 {
810                        table_row = Some(&tables::EXPONENT_BAND_32000[a as usize]);
811                    } else if self.sample_rate >= 22050 {
812                        table_row = Some(&tables::EXPONENT_BAND_22050[a as usize]);
813                    }
814                }
815
816                if let Some(row) = table_row {
817                    let n = row[0] as usize;
818                    for i in 0..n {
819                        self.exponent_bands[k][i] = row[1 + i] as u16;
820                    }
821                    self.exponent_sizes[k] = n;
822                } else {
823                    let mut j = 0usize;
824                    let mut lpos = 0usize;
825                    for idx in 0..25 {
826                        let a = tables::FF_WMA_CRITICAL_FREQS[idx] as usize;
827                        let b = self.sample_rate as usize;
828                        let mut pos = ((block_len * 2 * a) + (b << 1)) / (4 * b);
829                        pos <<= 2;
830                        if pos > block_len {
831                            pos = block_len;
832                        }
833                        if pos > lpos {
834                            self.exponent_bands[k][j] = (pos - lpos) as u16;
835                            j += 1;
836                        }
837                        if pos >= block_len {
838                            break;
839                        }
840                        lpos = pos;
841                    }
842                    self.exponent_sizes[k] = j;
843                }
844            }
845
846            self.coefs_end[k] = (self.frame_len - ((self.frame_len * 9) / 100)) >> k;
847            self.high_band_start[k] =
848                (((block_len as f32) * 2.0 * high_freq) / (self.sample_rate as f32) + 0.5) as usize;
849
850            let n = self.exponent_sizes[k];
851            let mut j = 0usize;
852            let mut pos = 0usize;
853            for i in 0..n {
854                let start0 = pos;
855                pos += self.exponent_bands[k][i] as usize;
856                let end0 = pos;
857                let mut start = start0;
858                let mut end = end0;
859                if start < self.high_band_start[k] {
860                    start = self.high_band_start[k];
861                }
862                if end > self.coefs_end[k] {
863                    end = self.coefs_end[k];
864                }
865                if end > start {
866                    self.exponent_high_bands[k][j] = (end - start) as u16;
867                    j += 1;
868                }
869            }
870            self.exponent_high_sizes[k] = j;
871        }
872
873        // Init MDCT windows.
874        self.windows.clear();
875        for i in 0..self.nb_block_sizes {
876            let half = 1usize << (self.frame_len_bits - i as i32);
877            self.windows.push(sine_window_init(half));
878        }
879
880        self.reset_block_lengths = true;
881
882        // Noise table.
883        if self.use_noise_coding {
884            self.noise_mult = if self.use_exp_vlc { 0.02 } else { 0.04 };
885            let mut seed: u32 = 1;
886            let norm = (1.0 / ((1u64 << 31) as f32)) * 3.0f32.sqrt() * self.noise_mult;
887            for i in 0..NOISE_TAB_SIZE {
888                seed = seed.wrapping_mul(314159).wrapping_add(1);
889                self.noise_table[i] = (seed as i32 as f32) * norm;
890            }
891        }
892
893        // Choose coef VLC tables.
894        let mut coef_vlc_table = 2;
895        if self.sample_rate >= 32000 {
896            if bps1 < 0.72 {
897                coef_vlc_table = 0;
898            } else if bps1 < 1.16 {
899                coef_vlc_table = 1;
900            }
901        }
902        let t0 = &tables::COEF_VLCS[coef_vlc_table * 2];
903        let t1 = &tables::COEF_VLCS[coef_vlc_table * 2 + 1];
904
905        self.init_coef_vlc(0, t0)?;
906        self.init_coef_vlc(1, t1)?;
907
908        Ok(())
909    }
910
911    fn init_coef_vlc(&mut self, idx: usize, tbl: &tables::CoefVlcTable) -> Result<()> {
912        // vlc_init(vlc, VLCBITS, n, table_bits, 1, 1, table_codes, 4, 4, 0)
913        let bits = tbl.huffbits;
914        let codes_u32 = tbl.huffcodes;
915        let codes_bytes: &[u8] = unsafe {
916            std::slice::from_raw_parts(codes_u32.as_ptr() as *const u8, codes_u32.len() * 4)
917        };
918        ff_vlc_init_sparse(
919            &mut self.coef_vlc[idx],
920            VLCBITS,
921            tbl.n,
922            bits,
923            1,
924            1,
925            codes_bytes,
926            4,
927            4,
928            None,
929            0,
930            0,
931            0,
932        )?;
933
934        // Build run/level tables like init_coef_vlc.
935        let n = tbl.n;
936        let levels_table = tbl.levels;
937
938        let mut run_table = vec![0u16; n];
939        let mut flevel_table = vec![0f32; n];
940        let mut int_table = vec![0u16; n];
941
942        let mut i = 2usize;
943        let mut level = 1usize;
944        let mut k = 0usize;
945        while i < n {
946            int_table[k] = i as u16;
947            let l = levels_table[k] as usize;
948            k += 1;
949            for j in 0..l {
950                run_table[i] = j as u16;
951                flevel_table[i] = level as f32;
952                i += 1;
953            }
954            level += 1;
955        }
956
957        self.run_table[idx] = run_table;
958        self.level_table[idx] = flevel_table;
959
960        Ok(())
961    }
962
963    fn ff_wma_total_gain_to_bits(total_gain: i32) -> i32 {
964        if total_gain < 15 {
965            13
966        } else if total_gain < 32 {
967            12
968        } else if total_gain < 40 {
969            11
970        } else if total_gain < 45 {
971            10
972        } else {
973            9
974        }
975    }
976
977    fn ff_wma_get_large_val(gb: &mut GetBitContext<'_>) -> Result<u32> {
978        let mut n_bits: usize = 8;
979        if gb.get_bits1()? != 0 {
980            n_bits += 8;
981            if gb.get_bits1()? != 0 {
982                n_bits += 8;
983                if gb.get_bits1()? != 0 {
984                    n_bits += 7;
985                }
986            }
987        }
988        gb.get_bits_long(n_bits)
989    }
990
991    #[allow(clippy::too_many_arguments)]
992    fn ff_wma_run_level_decode(
993        gb: &mut GetBitContext<'_>,
994        vlc: &[VlcElem],
995        level_table: &[f32],
996        run_table: &[u16],
997        version: i32,
998        ptr: &mut [f32],
999        mut offset: i32,
1000        num_coefs: i32,
1001        block_len: usize,
1002        frame_len_bits: i32,
1003        coef_nb_bits: i32,
1004    ) -> Result<()> {
1005        let coef_mask = (block_len as i32) - 1;
1006        while offset < num_coefs {
1007            let code = get_vlc2(gb, vlc, VLCBITS, VLCMAX)?;
1008            if code > 1 {
1009                offset += run_table[code as usize] as i32;
1010                let sign = gb.get_bits1()? as i32 - 1;
1011                let lvl_bits = level_table[code as usize].to_bits();
1012                let signed_bits = lvl_bits ^ ((sign as u32) & 0x8000_0000);
1013                ptr[(offset & coef_mask) as usize] = f32::from_bits(signed_bits);
1014            } else if code == 1 {
1015                break;
1016            } else {
1017                let level: i32;
1018                if version == 0 {
1019                    level = gb.get_bits(coef_nb_bits as usize)? as i32;
1020                    offset += gb.get_bits(frame_len_bits as usize)? as i32;
1021                } else {
1022                    level = Self::ff_wma_get_large_val(gb)? as i32;
1023                    if gb.get_bits1()? != 0 {
1024                        if gb.get_bits1()? != 0 {
1025                            if gb.get_bits1()? != 0 {
1026                                return Err(DecoderError::InvalidData(
1027                                    "broken escape sequence".into(),
1028                                ));
1029                            } else {
1030                                offset += gb.get_bits(frame_len_bits as usize)? as i32 + 4;
1031                            }
1032                        } else {
1033                            offset += gb.get_bits(2)? as i32 + 1;
1034                        }
1035                    }
1036                }
1037                let sign = gb.get_bits1()? as i32 - 1;
1038                let v = (level ^ sign) - sign;
1039                ptr[(offset & coef_mask) as usize] = v as f32;
1040            }
1041            offset += 1;
1042        }
1043
1044        if offset > num_coefs {
1045            return Err(DecoderError::InvalidData("overflow in spectral RLE".into()));
1046        }
1047
1048        Ok(())
1049    }
1050
1051    fn wma_lsp_to_curve_init(&mut self, frame_len: usize) {
1052        let wdel = std::f32::consts::PI / (frame_len as f32);
1053        for i in 0..frame_len {
1054            self.lsp_cos_table[i] = 2.0f32 * (wdel * (i as f32)).cos();
1055        }
1056
1057        for i in 0..256 {
1058            let e = (i as i32) - 126;
1059            self.lsp_pow_e_table[i] = (e as f32 * -0.25).exp2();
1060        }
1061
1062        let mut b = 1.0f32;
1063        for i in (0..(1 << LSP_POW_BITS)).rev() {
1064            let m = (1 << LSP_POW_BITS) + i;
1065            let mut a = (m as f32) * (0.5f32 / (1 << LSP_POW_BITS) as f32);
1066            a = 1.0f32 / a.sqrt().sqrt();
1067            self.lsp_pow_m_table1[i] = 2.0f32 * a - b;
1068            self.lsp_pow_m_table2[i] = b - a;
1069            b = a;
1070        }
1071    }
1072
1073    fn decode_exp_lsp(&mut self, gb: &mut GetBitContext<'_>, ch: usize) -> Result<()> {
1074        // upstream wmadec.c: decode_exp_lsp()
1075        let mut lsp: [f32; NB_LSP_COEFS] = [0.0; NB_LSP_COEFS];
1076        for i in 0..NB_LSP_COEFS {
1077            let val = if i == 0 || i >= 8 {
1078                gb.get_bits(3)? as usize
1079            } else {
1080                gb.get_bits(4)? as usize
1081            };
1082            lsp[i] = tables::FF_WMA_LSP_CODEBOOK[i][val];
1083        }
1084
1085        let cos = &self.lsp_cos_table;
1086        let e = &self.lsp_pow_e_table;
1087        let m1 = &self.lsp_pow_m_table1;
1088        let m2 = &self.lsp_pow_m_table2;
1089        let out = &mut self.exponents[ch];
1090        let vmax = wma_lsp_to_curve_tables(out, self.block_len, &lsp, cos, e, m1, m2);
1091        self.max_exponent[ch] = vmax;
1092        Ok(())
1093    }
1094
1095    fn decode_exp_vlc(&mut self, gb: &mut GetBitContext<'_>, ch: usize) -> Result<()> {
1096        let mut last_exp: i32;
1097        let mut max_scale: f32 = 0.0;
1098        let ptab = &tables::POW_TAB[60..];
1099
1100        let bsize = (self.frame_len_bits - self.block_len_bits) as usize;
1101        let bands = &self.exponent_bands[bsize];
1102
1103        let mut q = 0usize;
1104        let q_end = self.block_len;
1105
1106        if self.version.id() == 1 {
1107            last_exp = gb.get_bits(5)? as i32 + 10;
1108            let v = ptab[last_exp as usize];
1109            max_scale = v;
1110            let n = bands[0] as usize;
1111            for _ in 0..n {
1112                self.exponents[ch][q] = v;
1113                q += 1;
1114            }
1115        } else {
1116            last_exp = 36;
1117        }
1118
1119        let mut ptr_idx = 0usize;
1120        if self.version.id() == 1 {
1121            ptr_idx = 1;
1122        }
1123
1124        while q < q_end {
1125            let code = get_vlc2(gb, &self.exp_vlc.table, EXPVLCBITS, EXPMAX)?;
1126            last_exp += code - 60;
1127            if (last_exp as i32 + 60) as usize >= tables::POW_TAB.len() {
1128                return Err(DecoderError::InvalidData(format!(
1129                    "Exponent out of range: {last_exp}"
1130                )));
1131            }
1132            let v = ptab[last_exp as usize];
1133            if v > max_scale {
1134                max_scale = v;
1135            }
1136            let n = bands[ptr_idx] as usize;
1137            ptr_idx += 1;
1138            for _ in 0..n {
1139                self.exponents[ch][q] = v;
1140                q += 1;
1141            }
1142        }
1143
1144        self.max_exponent[ch] = max_scale;
1145        Ok(())
1146    }
1147
1148    fn wma_decode_block(&mut self, gb: &mut GetBitContext<'_>) -> Result<bool> {
1149        // Returns Ok(true) if last block of frame.
1150        // Translated from wma_decode_block.
1151
1152        // Compute current block length.
1153        if self.use_variable_block_len {
1154            let n = ilog2_u32((self.nb_block_sizes - 1) as u32) + 1;
1155            if self.reset_block_lengths {
1156                self.reset_block_lengths = false;
1157                let v = gb.get_bits(n as usize)? as usize;
1158                if v >= self.nb_block_sizes {
1159                    return Err(DecoderError::InvalidData(
1160                        "prev_block_len_bits out of range".into(),
1161                    ));
1162                }
1163                self.prev_block_len_bits = self.frame_len_bits - v as i32;
1164                let v = gb.get_bits(n as usize)? as usize;
1165                if v >= self.nb_block_sizes {
1166                    return Err(DecoderError::InvalidData(
1167                        "block_len_bits out of range".into(),
1168                    ));
1169                }
1170                self.block_len_bits = self.frame_len_bits - v as i32;
1171            } else {
1172                self.prev_block_len_bits = self.block_len_bits;
1173                self.block_len_bits = self.next_block_len_bits;
1174            }
1175            let v = gb.get_bits(n as usize)? as usize;
1176            if v >= self.nb_block_sizes {
1177                return Err(DecoderError::InvalidData(
1178                    "next_block_len_bits out of range".into(),
1179                ));
1180            }
1181            self.next_block_len_bits = self.frame_len_bits - v as i32;
1182        } else {
1183            self.next_block_len_bits = self.frame_len_bits;
1184            self.prev_block_len_bits = self.frame_len_bits;
1185            self.block_len_bits = self.frame_len_bits;
1186        }
1187
1188        let bsize = (self.frame_len_bits - self.block_len_bits) as usize;
1189        if (self.frame_len_bits - self.block_len_bits) as usize >= self.nb_block_sizes {
1190            return Err(DecoderError::InvalidData(
1191                "block_len_bits not initialized".into(),
1192            ));
1193        }
1194
1195        self.block_len = 1usize << self.block_len_bits;
1196        if self.block_pos + self.block_len > self.frame_len {
1197            return Err(DecoderError::InvalidData("frame_len overflow".into()));
1198        }
1199
1200        if self.channels == 2 {
1201            self.ms_stereo = gb.get_bits1()? != 0;
1202        }
1203
1204        let mut v_any = false;
1205        for ch in 0..self.channels {
1206            let a = gb.get_bits1()? != 0;
1207            self.channel_coded[ch] = a;
1208            v_any |= a;
1209        }
1210
1211        if !v_any {
1212            return self.wma_decode_block_next(gb, bsize);
1213        }
1214
1215        // Total gain.
1216        let mut total_gain: i32 = 1;
1217        loop {
1218            if gb.bits_left() < 7 {
1219                return Err(DecoderError::InvalidData("total_gain overread".into()));
1220            }
1221            let a = gb.get_bits(7)? as i32;
1222            total_gain += a;
1223            if a != 127 {
1224                break;
1225            }
1226        }
1227
1228        let coef_nb_bits = Self::ff_wma_total_gain_to_bits(total_gain);
1229
1230        // Number of coefficients.
1231        let ncoefs = (self.coefs_end[bsize] as i32) - (self.coefs_start as i32);
1232        let mut nb_coefs = [0i32; MAX_CHANNELS];
1233        for ch in 0..self.channels {
1234            nb_coefs[ch] = ncoefs;
1235        }
1236
1237        // Noise coding.
1238        if self.use_noise_coding {
1239            for ch in 0..self.channels {
1240                if self.channel_coded[ch] {
1241                    let n1 = self.exponent_high_sizes[bsize];
1242                    for i in 0..n1 {
1243                        let a = gb.get_bits1()? != 0;
1244                        self.high_band_coded[ch][i] = a;
1245                        if a {
1246                            nb_coefs[ch] -= self.exponent_high_bands[bsize][i] as i32;
1247                        }
1248                    }
1249                }
1250            }
1251            for ch in 0..self.channels {
1252                if self.channel_coded[ch] {
1253                    let n1 = self.exponent_high_sizes[bsize];
1254                    let mut val: i32 = 0x8000_0000u32 as i32;
1255                    for i in 0..n1 {
1256                        if self.high_band_coded[ch][i] {
1257                            if val == (0x8000_0000u32 as i32) {
1258                                val = gb.get_bits(7)? as i32 - 19;
1259                            } else {
1260                                val += get_vlc2(gb, &self.hgain_vlc.table, HGAINVLCBITS, HGAINMAX)?;
1261                            }
1262                            self.high_band_values[ch][i] = val;
1263                        }
1264                    }
1265                }
1266            }
1267        }
1268
1269        // Exponents can be reused in short blocks.
1270        let reuse = (self.block_len_bits == self.frame_len_bits) || (gb.get_bits1()? != 0);
1271        if reuse {
1272            for ch in 0..self.channels {
1273                if self.channel_coded[ch] {
1274                    if self.use_exp_vlc {
1275                        self.decode_exp_vlc(gb, ch)?;
1276                    } else {
1277                        self.decode_exp_lsp(gb, ch)?;
1278                    }
1279                    self.exponents_bsize[ch] = bsize;
1280                    self.exponents_initialized[ch] = true;
1281                }
1282            }
1283        }
1284
1285        for ch in 0..self.channels {
1286            if self.channel_coded[ch] && !self.exponents_initialized[ch] {
1287                return Err(DecoderError::InvalidData(
1288                    "exponents not initialized".into(),
1289                ));
1290            }
1291        }
1292
1293        // Parse spectral coefficients.
1294        for ch in 0..self.channels {
1295            if self.channel_coded[ch] {
1296                let tindex = (ch == 1 && self.ms_stereo) as usize;
1297                for v in &mut self.coefs1[ch][..self.block_len] {
1298                    *v = 0.0;
1299                }
1300                // Decode into coefs1 (upstream WMACoef).
1301                Self::ff_wma_run_level_decode(
1302                    gb,
1303                    &self.coef_vlc[tindex].table,
1304                    &self.level_table[tindex],
1305                    &self.run_table[tindex],
1306                    0,
1307                    &mut self.coefs1[ch],
1308                    0,
1309                    nb_coefs[ch],
1310                    self.block_len,
1311                    self.frame_len_bits,
1312                    coef_nb_bits,
1313                )?;
1314            }
1315            if self.version.id() == 1 && self.channels >= 2 {
1316                gb.align_to_byte();
1317            }
1318        }
1319
1320        // Normalize.
1321        let n4 = self.block_len / 2;
1322        let mut mdct_norm = 1.0f32 / (n4 as f32);
1323        if self.version.id() == 1 {
1324            mdct_norm *= (n4 as f32).sqrt();
1325        }
1326
1327        // Compute MDCT coefficients.
1328        for ch in 0..self.channels {
1329            if !self.channel_coded[ch] {
1330                continue;
1331            }
1332
1333            let esize = self.exponents_bsize[ch];
1334            let mult = ff_exp10f(total_gain as f32 * 0.05f32) / self.max_exponent[ch] * mdct_norm;
1335
1336            let mut coefs_pos = 0usize;
1337
1338            if self.use_noise_coding {
1339                // very low freqs: noise
1340                for i in 0..self.coefs_start {
1341                    let exp_idx = ((i << bsize) >> esize) as usize;
1342                    let noise = self.noise_table[self.noise_index];
1343                    self.noise_index = (self.noise_index + 1) & (NOISE_TAB_SIZE - 1);
1344                    self.coefs[ch][coefs_pos] = noise * self.exponents[ch][exp_idx] * mult;
1345                    coefs_pos += 1;
1346                }
1347
1348                let n1 = self.exponent_high_sizes[bsize];
1349
1350                // compute power of high bands
1351                let mut exp_power = [0f32; HIGH_BAND_MAX_SIZE];
1352                let mut exponents_ptr = (self.high_band_start[bsize] << bsize) >> esize;
1353                let mut last_high_band: usize = 0;
1354                for j in 0..n1 {
1355                    let n = self.exponent_high_bands[bsize][j] as usize;
1356                    if self.high_band_coded[ch][j] {
1357                        let mut e2: f32 = 0.0;
1358                        for i in 0..n {
1359                            let v = self.exponents[ch][exponents_ptr + ((i << bsize) >> esize)];
1360                            e2 += v * v;
1361                        }
1362                        exp_power[j] = e2 / (n as f32);
1363                        last_high_band = j;
1364                    }
1365                    exponents_ptr += (n << bsize) >> esize;
1366                }
1367
1368                // main freqs and high freqs
1369                let mut exponents_ptr = (self.coefs_start << bsize) >> esize;
1370                let mut coef1_idx = 0usize;
1371
1372                for j in (-1i32)..(n1 as i32) {
1373                    let n = if j < 0 {
1374                        self.high_band_start[bsize].saturating_sub(self.coefs_start)
1375                    } else {
1376                        self.exponent_high_bands[bsize][j as usize] as usize
1377                    };
1378
1379                    if j >= 0 && self.high_band_coded[ch][j as usize] {
1380                        let mut mult1 = (exp_power[j as usize] / exp_power[last_high_band]).sqrt();
1381                        mult1 *= ff_exp10f(self.high_band_values[ch][j as usize] as f32 * 0.05f32);
1382                        mult1 /= self.max_exponent[ch] * self.noise_mult;
1383                        mult1 *= mdct_norm;
1384
1385                        for i in 0..n {
1386                            let noise = self.noise_table[self.noise_index];
1387                            self.noise_index = (self.noise_index + 1) & (NOISE_TAB_SIZE - 1);
1388                            let exp = self.exponents[ch][exponents_ptr + ((i << bsize) >> esize)];
1389                            self.coefs[ch][coefs_pos] = noise * exp * mult1;
1390                            coefs_pos += 1;
1391                        }
1392                        exponents_ptr += (n << bsize) >> esize;
1393                    } else {
1394                        for i in 0..n {
1395                            let noise = self.noise_table[self.noise_index];
1396                            self.noise_index = (self.noise_index + 1) & (NOISE_TAB_SIZE - 1);
1397                            let exp = self.exponents[ch][exponents_ptr + ((i << bsize) >> esize)];
1398                            let coef1 = self.coefs1[ch][coef1_idx];
1399                            coef1_idx += 1;
1400                            self.coefs[ch][coefs_pos] = (coef1 + noise) * exp * mult;
1401                            coefs_pos += 1;
1402                        }
1403                        exponents_ptr += (n << bsize) >> esize;
1404                    }
1405                }
1406
1407                // very high freqs: noise
1408                let n = self.block_len - self.coefs_end[bsize];
1409                let exp_last =
1410                    self.exponents[ch][((exponents_ptr as i32 - (1 << bsize)) >> esize) as usize];
1411                let mult1 = mult * exp_last;
1412                for _ in 0..n {
1413                    let noise = self.noise_table[self.noise_index];
1414                    self.noise_index = (self.noise_index + 1) & (NOISE_TAB_SIZE - 1);
1415                    self.coefs[ch][coefs_pos] = noise * mult1;
1416                    coefs_pos += 1;
1417                }
1418            } else {
1419                for _ in 0..self.coefs_start {
1420                    self.coefs[ch][coefs_pos] = 0.0;
1421                    coefs_pos += 1;
1422                }
1423
1424                let n = nb_coefs[ch] as usize;
1425                for i in 0..n {
1426                    let exp = self.exponents[ch][((i << bsize) >> esize)];
1427                    let coef1 = self.coefs1[ch][i];
1428                    self.coefs[ch][coefs_pos] = coef1 * exp * mult;
1429                    coefs_pos += 1;
1430                }
1431                let tail = self.block_len - self.coefs_end[bsize];
1432                for _ in 0..tail {
1433                    self.coefs[ch][coefs_pos] = 0.0;
1434                    coefs_pos += 1;
1435                }
1436            }
1437        }
1438
1439        if self.ms_stereo && self.channel_coded[1] {
1440            if !self.channel_coded[0] {
1441                for v in &mut self.coefs[0][..self.block_len] {
1442                    *v = 0.0;
1443                }
1444                self.channel_coded[0] = true;
1445            }
1446            let (c0, c1) = self.coefs.split_at_mut(1);
1447            let v0 = &mut c0[0][..self.block_len];
1448            let v1 = &mut c1[0][..self.block_len];
1449            butterflies_float(v0, v1);
1450        }
1451
1452        self.wma_decode_block_next(gb, bsize)
1453    }
1454
1455    fn wma_decode_block_next(&mut self, _gb: &mut GetBitContext<'_>, bsize: usize) -> Result<bool> {
1456        // MDCT + window add.
1457        for ch in 0..self.channels {
1458            let n4 = self.block_len / 2;
1459            if self.channel_coded[ch] {
1460                self.mdct[bsize].imdct_full(
1461                    &mut self.output[..self.block_len * 2],
1462                    &self.coefs[ch][..self.block_len],
1463                );
1464            } else if !(self.ms_stereo && ch == 1) {
1465                for v in &mut self.output[..self.block_len * 2] {
1466                    *v = 0.0;
1467                }
1468            }
1469
1470            let index = (self.frame_len / 2) + self.block_pos - n4;
1471            // frame_out has length 2*BLOCK_MAX_SIZE.
1472            let frame_len_bits = self.frame_len_bits;
1473            let block_len_bits = self.block_len_bits;
1474            let prev_block_len_bits = self.prev_block_len_bits;
1475            let next_block_len_bits = self.next_block_len_bits;
1476            let block_len = self.block_len;
1477            let windows = &self.windows;
1478            let output = &self.output;
1479            let out_slice = &mut self.frame_out[ch][index..index + block_len * 2];
1480            wma_window_apply(
1481                out_slice,
1482                output,
1483                windows,
1484                frame_len_bits,
1485                block_len_bits,
1486                prev_block_len_bits,
1487                next_block_len_bits,
1488                block_len,
1489            );
1490        }
1491
1492        self.block_num += 1;
1493        self.block_pos += self.block_len;
1494        Ok(self.block_pos >= self.frame_len)
1495    }
1496
1497    fn wma_decode_frame(
1498        &mut self,
1499        gb: &mut GetBitContext<'_>,
1500        samples: &mut [Vec<f32>; MAX_CHANNELS],
1501        samples_offset: usize,
1502    ) -> Result<()> {
1503        self.block_num = 0;
1504        self.block_pos = 0;
1505        loop {
1506            let last = self.wma_decode_block(gb)?;
1507            if last {
1508                break;
1509            }
1510        }
1511
1512        for ch in 0..self.channels {
1513            samples[ch][samples_offset..samples_offset + self.frame_len]
1514                .copy_from_slice(&self.frame_out[ch][..self.frame_len]);
1515            // Shift for overlap.
1516            let tail = self.frame_out[ch][self.frame_len..self.frame_len * 2].to_vec();
1517            self.frame_out[ch][..self.frame_len].copy_from_slice(&tail);
1518        }
1519
1520        Ok(())
1521    }
1522}