Skip to main content

rusty_opus/silk/
dec_api.rs

1use crate::range_coder::RangeCoder;
2use crate::silk::decode_frame::{FLAG_DECODE_NORMAL, FLAG_PACKET_LOST, silk_decode_frame};
3use crate::silk::decode_indices::{
4    silk_decode_indices, silk_stereo_decode_mid_only, silk_stereo_decode_pred,
5    silk_stereo_ms_to_lr,
6};
7use crate::silk::decode_pulses::silk_decode_pulses;
8use crate::silk::decoder_structs::SilkDecoderState;
9use crate::silk::define::*;
10use crate::silk::init_decoder::{silk_decoder_set_fs, silk_init_decoder};
11use crate::silk::tables::{SILK_LBRR_FLAGS_2_ICDF, SILK_LBRR_FLAGS_3_ICDF};
12
13pub struct SilkDecoder {
14    pub channel_state: [SilkDecoderState; 2],
15
16    pub n_channels_api: i32,
17
18    pub n_channels_internal: i32,
19
20    pub prev_decode_only_middle: i32,
21
22    // Stereo MS->LR reconstruction state (libopus stereo_dec_state).
23    pub s_stereo_pred_prev_q13: [i32; 2],
24    pub s_stereo_mid: [i16; 2],
25    pub s_stereo_side: [i16; 2],
26
27    // When set, a stereo decode reconstructs L/R via silk_stereo_ms_to_lr and
28    // publishes them (1-sample-delay-line layout, ready to resample) in l_out/
29    // r_out. When clear, stereo decodes emit only the mid (mono downmix) as
30    // before. lib.rs sets this for the pure-SILK stereo path.
31    pub produce_lr: bool,
32    pub l_out: [i16; MAX_FRAME_LENGTH],
33    pub r_out: [i16; MAX_FRAME_LENGTH],
34}
35
36impl Default for SilkDecoder {
37    fn default() -> Self {
38        Self::new()
39    }
40}
41
42impl SilkDecoder {
43    pub fn new() -> Self {
44        let mut dec = Self {
45            channel_state: [SilkDecoderState::default(), SilkDecoderState::default()],
46            n_channels_api: 1,
47            n_channels_internal: 1,
48            prev_decode_only_middle: 0,
49            s_stereo_pred_prev_q13: [0; 2],
50            s_stereo_mid: [0; 2],
51            s_stereo_side: [0; 2],
52            produce_lr: false,
53            l_out: [0; MAX_FRAME_LENGTH],
54            r_out: [0; MAX_FRAME_LENGTH],
55        };
56        silk_init_decoder(&mut dec.channel_state[0]);
57        silk_init_decoder(&mut dec.channel_state[1]);
58        dec
59    }
60
61    pub fn init(&mut self, sample_rate_hz: i32, channels: i32) -> i32 {
62        let fs_khz = sample_rate_hz / 1000;
63        let ret = silk_decoder_set_fs(&mut self.channel_state[0], fs_khz, sample_rate_hz);
64        if ret < 0 {
65            return ret;
66        }
67        if channels == 2 {
68            let ret = silk_decoder_set_fs(&mut self.channel_state[1], fs_khz, sample_rate_hz);
69            if ret < 0 {
70                return ret;
71            }
72        }
73
74        self.channel_state[0].n_frames_per_packet = 1;
75        self.n_channels_api = channels;
76        self.n_channels_internal = channels;
77        ret
78    }
79
80    pub fn decode(
81        &mut self,
82        range_dec: &mut RangeCoder,
83        output: &mut [i16],
84        lost_flag: i32,
85        new_packet: bool,
86        payload_size_ms: i32,
87        internal_sample_rate: i32,
88    ) -> i32 {
89        if new_packet {
90            self.channel_state[0].n_frames_decoded = 0;
91            self.channel_state[1].n_frames_decoded = 0;
92        }
93
94        if self.channel_state[0].n_frames_decoded == 0 {
95            let fs_khz_dec = (internal_sample_rate >> 10) + 1;
96            if fs_khz_dec != 8 && fs_khz_dec != 12 && fs_khz_dec != 16 {
97                return -1;
98            }
99            let api_sample_rate = self.channel_state[0].fs_api_hz;
100            // libopus configures EVERY internal channel here, not just ch0 — the
101            // side channel needs its frame_length/fs set or a stereo side decode
102            // consumes the wrong number of bits and desyncs the range coder.
103            for n in 0..self.n_channels_internal as usize {
104                match payload_size_ms {
105                    0 | 10 => {
106                        self.channel_state[n].n_frames_per_packet = 1;
107                        self.channel_state[n].nb_subfr = 2;
108                    }
109                    20 => {
110                        self.channel_state[n].n_frames_per_packet = 1;
111                        self.channel_state[n].nb_subfr = MAX_NB_SUBFR as i32;
112                    }
113                    40 => {
114                        self.channel_state[n].n_frames_per_packet = 2;
115                        self.channel_state[n].nb_subfr = MAX_NB_SUBFR as i32;
116                    }
117                    60 => {
118                        self.channel_state[n].n_frames_per_packet = 3;
119                        self.channel_state[n].nb_subfr = MAX_NB_SUBFR as i32;
120                    }
121                    _ => return -1,
122                }
123                let ret =
124                    silk_decoder_set_fs(&mut self.channel_state[n], fs_khz_dec, api_sample_rate);
125                if ret < 0 {
126                    return ret;
127                }
128                if payload_size_ms == 10 {
129                    self.channel_state[n].nb_subfr = 2;
130                    self.channel_state[n].frame_length = self.channel_state[n].subfr_length * 2;
131                }
132            }
133        }
134
135        if lost_flag != FLAG_PACKET_LOST && self.channel_state[0].n_frames_decoded == 0 {
136            let n_frames_per_packet = self.channel_state[0].n_frames_per_packet.max(1);
137            let n_channels = self.n_channels_internal as usize;
138
139            for n in 0..n_channels {
140                for i in 0..n_frames_per_packet as usize {
141                    let vad = range_dec.decode_bit_logp(1);
142                    self.channel_state[n].vad_flags[i] = if vad { 1 } else { 0 };
143                }
144                let lbrr = range_dec.decode_bit_logp(1);
145                self.channel_state[n].lbrr_flag = if lbrr { 1 } else { 0 };
146            }
147
148            for n in 0..n_channels {
149                self.channel_state[n].lbrr_flags.fill(0);
150                if self.channel_state[n].lbrr_flag != 0 {
151                    if n_frames_per_packet == 1 {
152                        self.channel_state[n].lbrr_flags[0] = 1;
153                    } else {
154                        let lbrr_icdf = match n_frames_per_packet {
155                            2 => &SILK_LBRR_FLAGS_2_ICDF[..],
156                            3 => &SILK_LBRR_FLAGS_3_ICDF[..],
157                            _ => &SILK_LBRR_FLAGS_2_ICDF[..],
158                        };
159                        let lbrr_symbol = range_dec.decode_icdf(lbrr_icdf, 8) + 1;
160                        for i in 0..n_frames_per_packet as usize {
161                            self.channel_state[n].lbrr_flags[i] = (lbrr_symbol >> i) & 1;
162                        }
163                    }
164                }
165            }
166
167            // Skip LBRR data for all frames/channels
168            if lost_flag == FLAG_DECODE_NORMAL {
169                for i in 0..n_frames_per_packet as usize {
170                    for n in 0..n_channels {
171                        if self.channel_state[n].lbrr_flags[i] != 0 {
172                            if n_channels == 2 && n == 0 {
173                                let _ = silk_stereo_decode_pred(range_dec);
174                                if self.channel_state[1].lbrr_flags[i] == 0 {
175                                    silk_stereo_decode_mid_only(range_dec);
176                                }
177                            }
178                            let cond_coding =
179                                if i > 0 && self.channel_state[n].lbrr_flags[i - 1] != 0 {
180                                    CODE_CONDITIONALLY
181                                } else {
182                                    CODE_INDEPENDENTLY
183                                };
184                            silk_decode_indices(
185                                &mut self.channel_state[n],
186                                range_dec,
187                                i as i32,
188                                1,
189                                cond_coding,
190                            );
191                            let mut pulses = [0i16; MAX_FRAME_LENGTH];
192                            silk_decode_pulses(
193                                range_dec,
194                                &mut pulses,
195                                self.channel_state[n].indices.signal_type as i32,
196                                self.channel_state[n].indices.quant_offset_type as i32,
197                                self.channel_state[n].frame_length,
198                            );
199                        }
200                    }
201                }
202            }
203        }
204
205        let frame_index = self.channel_state[0].n_frames_decoded as usize;
206        let mut decode_only_middle = 0i32;
207        let mut ms_pred_q13 = [0i32; 2];
208        if self.n_channels_internal == 2 && lost_flag == FLAG_DECODE_NORMAL {
209            // Predictors + mid-only flag. Bits must be read to keep the range
210            // coder in sync.
211            ms_pred_q13 = silk_stereo_decode_pred(range_dec);
212            if self.channel_state[1].vad_flags[frame_index] == 0 {
213                decode_only_middle = if silk_stereo_decode_mid_only(range_dec) {
214                    1
215                } else {
216                    0
217                };
218            }
219        }
220
221        // Reset the side channel's prediction memory for the first frame that
222        // codes side after a run of mid-only frames (libopus dec_API.c:249-256).
223        if self.n_channels_internal == 2
224            && decode_only_middle == 0
225            && self.prev_decode_only_middle == 1
226        {
227            let ch1 = &mut self.channel_state[1];
228            ch1.out_buf.fill(0);
229            ch1.s_lpc_q14_buf.fill(0);
230            ch1.lag_prev = 100;
231            ch1.last_gain_index = 10;
232            ch1.prev_signal_type = TYPE_NO_VOICE_ACTIVITY;
233            ch1.first_frame_after_reset = 1;
234        }
235
236        let cond_coding = if frame_index == 0 {
237            CODE_INDEPENDENTLY
238        } else {
239            CODE_CONDITIONALLY
240        };
241
242        let mut n_samples_out: i32 = 0;
243        let ret = silk_decode_frame(
244            &mut self.channel_state[0],
245            range_dec,
246            output,
247            &mut n_samples_out,
248            lost_flag,
249            cond_coding,
250        );
251        self.channel_state[0].n_frames_decoded += 1;
252
253        let mut side_samples = [0i16; MAX_FRAME_LENGTH];
254        let mut side_decoded = false;
255
256        // Stereo: we output only the mid (mono downmix) for now, but the side
257        // frame's bits MUST still be consumed or the range coder desyncs for the
258        // next internal frame of a multi-frame packet (the mid of frames 2/3 then
259        // decodes from garbage). Decode it into a scratch buffer and discard.
260        if self.n_channels_internal == 2 && decode_only_middle == 0 && lost_flag == FLAG_DECODE_NORMAL
261        {
262            // libopus FrameIndex for the side (n=1) = channel_state[0].nFramesDecoded - 1,
263            // evaluated AFTER ch0's own increment, which equals the original
264            // frame_index. (Using frame_index-1 wrongly forced INDEP on the 2nd
265            // internal frame -> wrong bit count -> desync.)
266            let cond1 = if frame_index == 0 {
267                CODE_INDEPENDENTLY
268            } else if self.prev_decode_only_middle == 1 {
269                CODE_INDEPENDENTLY_NO_LTP_SCALING
270            } else {
271                CODE_CONDITIONALLY
272            };
273            let mut ns_side: i32 = 0;
274            let ret1 = silk_decode_frame(
275                &mut self.channel_state[1],
276                range_dec,
277                &mut side_samples,
278                &mut ns_side,
279                lost_flag,
280                cond1,
281            );
282            if ret1 < 0 {
283                return ret1;
284            }
285            side_decoded = true;
286        }
287
288        // Reconstruct L/R for the pure-SILK stereo path. mid is in `output`
289        // ([0..n]); side in side_samples ([0..n], zero if mid-only). ms_to_lr
290        // wants x1/x2 laid out as [hist0, hist1, samples...] and writes L to
291        // x1[1..1+n], R to x2[1..1+n] — exactly the 1-sample delay line the
292        // resampler consumes.
293        if self.produce_lr && self.n_channels_internal == 2 && n_samples_out > 0 {
294            let n = n_samples_out as usize;
295            let mut mid_buf = [0i16; MAX_FRAME_LENGTH + 2];
296            let mut side_buf = [0i16; MAX_FRAME_LENGTH + 2];
297            mid_buf[2..2 + n].copy_from_slice(&output[..n]);
298            if side_decoded {
299                side_buf[2..2 + n].copy_from_slice(&side_samples[..n]);
300            }
301            silk_stereo_ms_to_lr(
302                &mut self.s_stereo_pred_prev_q13,
303                &mut self.s_stereo_mid,
304                &mut self.s_stereo_side,
305                &mut mid_buf,
306                &mut side_buf,
307                &ms_pred_q13,
308                self.channel_state[0].fs_khz,
309                n,
310            );
311            self.l_out[..n].copy_from_slice(&mid_buf[1..1 + n]);
312            self.r_out[..n].copy_from_slice(&side_buf[1..1 + n]);
313        } else if n_samples_out >= 2 {
314            // Mono frame: buffer the last two mid samples so the NEXT stereo
315            // frame's ms_to_lr starts from the correct 2-sample history (libopus
316            // dec_API.c:312-313 — sStereo.sMid is updated even on mono frames).
317            // Without this the first stereo frame after a mono run reconstructs
318            // from stale history (~1/4-magnitude glitch at the switch).
319            let n = n_samples_out as usize;
320            self.s_stereo_mid[0] = output[n - 2];
321            self.s_stereo_mid[1] = output[n - 1];
322        }
323
324        // libopus increments channel_state[1].nFramesDecoded on EVERY internal
325        // frame (even mid-only). The side's silk_decode_frame keys its VAD/
326        // signal-type lookup off this counter; leaving it at 0 made every side
327        // frame read vad_flags[0] -> wrong signal type on frame 2+ -> desync.
328        if self.n_channels_internal == 2 {
329            self.channel_state[1].n_frames_decoded += 1;
330        }
331
332        self.prev_decode_only_middle = decode_only_middle;
333
334        if ret < 0 { ret } else { n_samples_out }
335    }
336
337    pub fn decode_bytes(&mut self, data: &[u8], output: &mut [i16], new_packet: bool) -> i32 {
338        let mut range_dec = RangeCoder::new_decoder(data);
339        let internal_rate = self.channel_state[0].fs_khz * 1000;
340        let payload_ms = if self.channel_state[0].nb_subfr == 2 {
341            10
342        } else {
343            20
344        };
345        self.decode(
346            &mut range_dec,
347            output,
348            FLAG_DECODE_NORMAL,
349            new_packet,
350            payload_ms,
351            internal_rate,
352        )
353    }
354
355    pub fn reset(&mut self) {
356        silk_init_decoder(&mut self.channel_state[0]);
357        silk_init_decoder(&mut self.channel_state[1]);
358        self.prev_decode_only_middle = 0;
359    }
360
361    pub fn frame_length(&self) -> i32 {
362        self.channel_state[0].frame_length
363    }
364
365    pub fn sample_rate(&self) -> i32 {
366        self.channel_state[0].fs_khz * 1000
367    }
368}
369
370#[cfg(test)]
371mod tests {
372    use super::*;
373
374    #[test]
375    fn test_decoder_creation() {
376        let dec = SilkDecoder::new();
377        assert_eq!(dec.n_channels_api, 1);
378        assert_eq!(dec.n_channels_internal, 1);
379    }
380
381    #[test]
382    fn test_decoder_init() {
383        let mut dec = SilkDecoder::new();
384
385        let ret = dec.init(16000, 1);
386        assert_eq!(ret, 0);
387        assert_eq!(dec.sample_rate(), 16000);
388    }
389
390    #[test]
391    fn test_decoder_16khz() {
392        let mut dec = SilkDecoder::new();
393        let ret = dec.init(16000, 1);
394        assert_eq!(ret, 0);
395        assert_eq!(dec.frame_length(), 320);
396    }
397}