Skip to main content

opus_rs/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};
6use crate::silk::decode_pulses::silk_decode_pulses;
7use crate::silk::decoder_structs::SilkDecoderState;
8use crate::silk::define::*;
9use crate::silk::init_decoder::{silk_decoder_set_fs, silk_init_decoder};
10use crate::silk::stereo_ms_to_lr::{StereoDecState, silk_stereo_ms_to_lr};
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    /// Persistent stereo M/S decoder state.
23    pub s_stereo: StereoDecState,
24
25    /// Per-channel temp buffers for M/S → L/R conversion (frame_length + 2).
26    w_silk_buf: [[i16; MAX_FRAME_LENGTH + 2]; 2],
27}
28
29impl Default for SilkDecoder {
30    fn default() -> Self {
31        Self::new()
32    }
33}
34
35impl SilkDecoder {
36    pub fn new() -> Self {
37        let mut dec = Self {
38            channel_state: [SilkDecoderState::default(), SilkDecoderState::default()],
39            n_channels_api: 1,
40            n_channels_internal: 1,
41            prev_decode_only_middle: 0,
42            s_stereo: StereoDecState::default(),
43            w_silk_buf: [[0; MAX_FRAME_LENGTH + 2]; 2],
44        };
45        silk_init_decoder(&mut dec.channel_state[0]);
46        silk_init_decoder(&mut dec.channel_state[1]);
47        dec
48    }
49
50    pub fn init(&mut self, sample_rate_hz: i32, channels: i32) -> i32 {
51        let fs_khz = sample_rate_hz / 1000;
52        let ret = silk_decoder_set_fs(&mut self.channel_state[0], fs_khz, sample_rate_hz);
53        if ret < 0 {
54            return ret;
55        }
56        if channels == 2 {
57            let ret = silk_decoder_set_fs(&mut self.channel_state[1], fs_khz, sample_rate_hz);
58            if ret < 0 {
59                return ret;
60            }
61        }
62
63        self.channel_state[0].n_frames_per_packet = 1;
64        self.n_channels_api = channels;
65        self.n_channels_internal = channels;
66        ret
67    }
68
69    pub fn decode(
70        &mut self,
71        range_dec: &mut RangeCoder,
72        output: &mut [i16],
73        lost_flag: i32,
74        new_packet: bool,
75        payload_size_ms: i32,
76        internal_sample_rate: i32,
77    ) -> i32 {
78        if new_packet {
79            self.channel_state[0].n_frames_decoded = 0;
80            self.channel_state[1].n_frames_decoded = 0;
81        }
82
83        if self.channel_state[0].n_frames_decoded == 0 {
84            let n_channels = self.n_channels_internal as usize;
85            for n in 0..n_channels {
86                match payload_size_ms {
87                    0 | 10 => {
88                        self.channel_state[n].n_frames_per_packet = 1;
89                        self.channel_state[n].nb_subfr = 2;
90                    }
91                    20 => {
92                        self.channel_state[n].n_frames_per_packet = 1;
93                        self.channel_state[n].nb_subfr = MAX_NB_SUBFR as i32;
94                    }
95                    40 => {
96                        self.channel_state[n].n_frames_per_packet = 2;
97                        self.channel_state[n].nb_subfr = MAX_NB_SUBFR as i32;
98                    }
99                    60 => {
100                        self.channel_state[n].n_frames_per_packet = 3;
101                        self.channel_state[n].nb_subfr = MAX_NB_SUBFR as i32;
102                    }
103                    _ => return -1,
104                }
105            }
106
107            let fs_khz_dec = (internal_sample_rate >> 10) + 1;
108            if fs_khz_dec != 8 && fs_khz_dec != 12 && fs_khz_dec != 16 {
109                return -1;
110            }
111            let api_sample_rate = self.channel_state[0].fs_api_hz;
112            for n in 0..n_channels {
113                let ret = silk_decoder_set_fs(&mut self.channel_state[n], fs_khz_dec, api_sample_rate);
114                if ret < 0 {
115                    return ret;
116                }
117                if payload_size_ms == 10 {
118                    self.channel_state[n].nb_subfr = 2;
119                    self.channel_state[n].frame_length = self.channel_state[n].subfr_length * 2;
120                }
121            }
122        }
123
124        // Initialise second channel and stereo state on mono→stereo transition.
125        if self.n_channels_internal == 2 && self.prev_decode_only_middle == 0 && self.s_stereo.pred_prev_q13 == [0, 0] {
126            // Fresh stereo init: clear stereo state.
127            self.s_stereo = StereoDecState::default();
128        }
129
130        if lost_flag != FLAG_PACKET_LOST && self.channel_state[0].n_frames_decoded == 0 {
131            let n_frames_per_packet = self.channel_state[0].n_frames_per_packet.max(1);
132            let n_channels = self.n_channels_internal as usize;
133
134            for n in 0..n_channels {
135                for i in 0..n_frames_per_packet as usize {
136                    let vad = range_dec.decode_bit_logp(1);
137                    self.channel_state[n].vad_flags[i] = if vad { 1 } else { 0 };
138                }
139                let lbrr = range_dec.decode_bit_logp(1);
140                self.channel_state[n].lbrr_flag = if lbrr { 1 } else { 0 };
141            }
142
143            for n in 0..n_channels {
144                self.channel_state[n].lbrr_flags.fill(0);
145                if self.channel_state[n].lbrr_flag != 0 {
146                    if n_frames_per_packet == 1 {
147                        self.channel_state[n].lbrr_flags[0] = 1;
148                    } else {
149                        let lbrr_icdf = match n_frames_per_packet {
150                            2 => &SILK_LBRR_FLAGS_2_ICDF[..],
151                            3 => &SILK_LBRR_FLAGS_3_ICDF[..],
152                            _ => &SILK_LBRR_FLAGS_2_ICDF[..],
153                        };
154                        let lbrr_symbol = range_dec.decode_icdf(lbrr_icdf, 8) + 1;
155                        for i in 0..n_frames_per_packet as usize {
156                            self.channel_state[n].lbrr_flags[i] = (lbrr_symbol >> i) & 1;
157                        }
158                    }
159                }
160            }
161
162            // Skip LBRR data for all frames/channels
163            if lost_flag == FLAG_DECODE_NORMAL {
164                for i in 0..n_frames_per_packet as usize {
165                    for n in 0..n_channels {
166                        if self.channel_state[n].lbrr_flags[i] != 0 {
167                            if n_channels == 2 && n == 0 {
168                                let _ = silk_stereo_decode_pred(range_dec);
169                                if self.channel_state[1].lbrr_flags[i] == 0 {
170                                    let _ = silk_stereo_decode_mid_only(range_dec);
171                                }
172                            }
173                            let cond_coding =
174                                if i > 0 && self.channel_state[n].lbrr_flags[i - 1] != 0 {
175                                    CODE_CONDITIONALLY
176                                } else {
177                                    CODE_INDEPENDENTLY
178                                };
179                            silk_decode_indices(
180                                &mut self.channel_state[n],
181                                range_dec,
182                                i as i32,
183                                1,
184                                cond_coding,
185                            );
186                            let mut pulses = [0i16; MAX_FRAME_LENGTH];
187                            silk_decode_pulses(
188                                range_dec,
189                                &mut pulses,
190                                self.channel_state[n].indices.signal_type as i32,
191                                self.channel_state[n].indices.quant_offset_type as i32,
192                                self.channel_state[n].frame_length,
193                            );
194                        }
195                    }
196                }
197            }
198        }
199
200        // Decode M/S predictors for stereo.
201        let mut ms_pred_q13 = [0i32; 2];
202        let mut decode_only_middle = 0i32;
203        let frame_index = self.channel_state[0].n_frames_decoded as usize;
204        if self.n_channels_internal == 2 {
205            if lost_flag == FLAG_DECODE_NORMAL {
206                ms_pred_q13 = silk_stereo_decode_pred(range_dec);
207                if self.channel_state[1].vad_flags[frame_index] == 0 {
208                    decode_only_middle = if silk_stereo_decode_mid_only(range_dec) { 1 } else { 0 };
209                }
210            } else {
211                ms_pred_q13 = self.s_stereo.pred_prev_q13;
212            }
213        }
214
215        // Reset side-channel prediction memory when transitioning from mid-only to has-side.
216        if self.n_channels_internal == 2
217            && decode_only_middle == 0
218            && self.prev_decode_only_middle == 1
219        {
220            let ch1 = &mut self.channel_state[1];
221            ch1.out_buf.fill(0);
222            ch1.s_lpc_q14_buf.fill(0);
223            ch1.lag_prev = 100;
224            ch1.last_gain_index = 10;
225            ch1.prev_signal_type = TYPE_NO_VOICE_ACTIVITY;
226            ch1.first_frame_after_reset = 1;
227        }
228
229        let has_side = decode_only_middle == 0;
230        let n_channels = self.n_channels_internal as usize;
231        let frame_length = self.channel_state[0].frame_length as usize;
232
233        // Decode each channel into w_silk_buf[n][2..2+frame_length].
234        let mut n_samples_out: i32 = 0;
235        for n in 0..n_channels {
236            if n == 0 || has_side {
237                let frame_idx = self.channel_state[0].n_frames_decoded - n as i32;
238                let cond_coding = if frame_idx <= 0 {
239                    CODE_INDEPENDENTLY
240                } else if n > 0 && self.prev_decode_only_middle != 0 {
241                    CODE_INDEPENDENTLY_NO_LTP_SCALING
242                } else {
243                    CODE_CONDITIONALLY
244                };
245
246                // Clear the overlap region before decode.
247                self.w_silk_buf[n][0] = 0;
248                self.w_silk_buf[n][1] = 0;
249
250                let mut ns = 0i32;
251                let ret = silk_decode_frame(
252                    &mut self.channel_state[n],
253                    range_dec,
254                    &mut self.w_silk_buf[n][2..],
255                    &mut ns,
256                    lost_flag,
257                    cond_coding,
258                );
259                if ret < 0 {
260                    return ret;
261                }
262                if n == 0 {
263                    n_samples_out = ns;
264                }
265            } else {
266                // Mid-only: zero the side channel.
267                for i in 0..frame_length {
268                    self.w_silk_buf[n][2 + i] = 0;
269                }
270            }
271            self.channel_state[n].n_frames_decoded += 1;
272        }
273
274        // Convert Mid/Side → Left/Right for stereo, or do mono buffering.
275        if n_channels == 2 {
276            let (buf_l, buf_r) = self.w_silk_buf.split_at_mut(1);
277            silk_stereo_ms_to_lr(
278                &mut self.s_stereo,
279                &mut buf_l[0],
280                &mut buf_r[0],
281                &ms_pred_q13,
282                self.channel_state[0].fs_khz,
283                n_samples_out,
284            );
285        } else {
286            // Mono buffering: save overlap for next frame.
287            self.w_silk_buf[0][0] = self.s_stereo.s_mid[0];
288            self.w_silk_buf[0][1] = self.s_stereo.s_mid[1];
289            self.s_stereo.s_mid[0] = self.w_silk_buf[0][n_samples_out as usize];
290            self.s_stereo.s_mid[1] = self.w_silk_buf[0][n_samples_out as usize + 1];
291        }
292
293        // Copy planar output: ch0 at output[0..fl], ch1 at output[fl..2*fl].
294        // For stereo, MS_to_LR outputs at [1..1+fl] (with overlap compensation).
295        // For mono, decoded data is at [2..2+fl] — use it directly (no offset)
296        // to match the pre-stereo-rewrite behaviour.
297        let fl = n_samples_out as usize;
298        if n_channels == 2 {
299            if output.len() < 2 * fl {
300                return -1;
301            }
302            output[..fl].copy_from_slice(&self.w_silk_buf[0][1..1 + fl]);
303            output[fl..2 * fl].copy_from_slice(&self.w_silk_buf[1][1..1 + fl]);
304        } else {
305            if output.len() < fl {
306                return -1;
307            }
308            output[..fl].copy_from_slice(&self.w_silk_buf[0][2..2 + fl]);
309        }
310
311        self.prev_decode_only_middle = decode_only_middle;
312
313        if n_samples_out < 0 { -1 } else { n_samples_out }
314    }
315
316    pub fn decode_bytes(&mut self, data: &[u8], output: &mut [i16], new_packet: bool) -> i32 {
317        let mut range_dec = RangeCoder::new_decoder(data);
318        let internal_rate = self.channel_state[0].fs_khz * 1000;
319        let payload_ms = if self.channel_state[0].nb_subfr == 2 {
320            10
321        } else {
322            20
323        };
324        self.decode(
325            &mut range_dec,
326            output,
327            FLAG_DECODE_NORMAL,
328            new_packet,
329            payload_ms,
330            internal_rate,
331        )
332    }
333
334    pub fn reset(&mut self) {
335        silk_init_decoder(&mut self.channel_state[0]);
336        silk_init_decoder(&mut self.channel_state[1]);
337        self.prev_decode_only_middle = 0;
338        self.s_stereo = StereoDecState::default();
339    }
340
341    pub fn frame_length(&self) -> i32 {
342        self.channel_state[0].frame_length
343    }
344
345    pub fn sample_rate(&self) -> i32 {
346        self.channel_state[0].fs_khz * 1000
347    }
348}
349
350#[cfg(test)]
351mod tests {
352    use super::*;
353
354    #[test]
355    fn test_decoder_creation() {
356        let dec = SilkDecoder::new();
357        assert_eq!(dec.n_channels_api, 1);
358        assert_eq!(dec.n_channels_internal, 1);
359    }
360
361    #[test]
362    fn test_decoder_init() {
363        let mut dec = SilkDecoder::new();
364
365        let ret = dec.init(16000, 1);
366        assert_eq!(ret, 0);
367        assert_eq!(dec.sample_rate(), 16000);
368    }
369
370    #[test]
371    fn test_decoder_16khz() {
372        let mut dec = SilkDecoder::new();
373        let ret = dec.init(16000, 1);
374        assert_eq!(ret, 0);
375        assert_eq!(dec.frame_length(), 320);
376    }
377}