1use crate::audio::{DEFAULT_OUTPUT_FREQUENCY, RESAMPLE_SCALING_FACTOR};
2use bincode::{Decode, Encode};
3use multiversion::multiversion;
4use std::array;
5use std::collections::VecDeque;
6use std::marker::PhantomData;
7use std::ops::Deref;
8
9// This is different from VecDeque in that the samples are guaranteed to always be contiguous in
10// memory, which is important for performance when N is large
11#[derive(Debug, Clone, Encode, Decode)]
12struct RingBuffer<const N: usize> {
13    buffer: Vec<f64>,
14    idx: usize,
15    len: usize,
16}
17
18impl<const N: usize> RingBuffer<N> {
19    const fn capacity() -> usize {
20        32 * N
21    }
22
23    fn new() -> Self {
24        Self { buffer: vec![0.0; Self::capacity()], idx: Self::capacity(), len: 0 }
25    }
26
27    fn push(&mut self, sample: f64) {
28        if self.len < N {
29            self.idx -= 1;
30            self.buffer[self.idx] = sample;
31            self.len += 1;
32            return;
33        }
34
35        if self.idx == 0 {
36            for i in 1..N {
37                self.buffer[Self::capacity() - N + i] = self.buffer[i - 1];
38            }
39            self.idx = Self::capacity() - N;
40            self.buffer[self.idx] = sample;
41            return;
42        }
43
44        self.idx -= 1;
45        self.buffer[self.idx] = sample;
46    }
47}
48
49// Force coefficients to be aligned to a 64-byte boundary in order to support AVX512 aligned loads
50#[derive(Debug, Clone)]
51#[repr(C, align(64))]
52pub struct LpfCoefficients<const LPF_TAPS: usize>(pub [f64; LPF_TAPS]);
53
54impl<const LPF_TAPS: usize> Deref for LpfCoefficients<LPF_TAPS> {
55    type Target = [f64; LPF_TAPS];
56
57    #[inline]
58    fn deref(&self) -> &Self::Target {
59        &self.0
60    }
61}
62
63// Would be nicer if `LPF_TAPS` was an associated const, but `[f64; Self::LPF_TAPS]` doesn't compile
64// without nightly features related to const generic expressions
65pub trait FirKernel<const LPF_TAPS: usize> {
66    fn lpf_coefficients() -> &'static LpfCoefficients<LPF_TAPS>;
67}
68
69#[derive(Debug, Clone, Encode, Decode)]
70pub struct FirResampler<const CHANNELS: usize, const LPF_TAPS: usize, Kernel: FirKernel<LPF_TAPS>> {
71    input: [RingBuffer<LPF_TAPS>; CHANNELS],
72    output: VecDeque<[f64; CHANNELS]>,
73    sample_count_product: u64,
74    scaled_output_frequency: u64,
75    scaled_source_frequency: u64,
76    // Required for this struct to compile with the generic Kernel type
77    _marker: PhantomData<Kernel>,
78}
79
80impl<const CHANNELS: usize, const LPF_TAPS: usize, Kernel: FirKernel<LPF_TAPS>>
81    FirResampler<CHANNELS, LPF_TAPS, Kernel>
82{
83    #[must_use]
84    pub fn new(source_frequency: f64, output_frequency: u64) -> Self {
85        Self {
86            input: array::from_fn(|_| RingBuffer::new()),
87            output: VecDeque::with_capacity((DEFAULT_OUTPUT_FREQUENCY / 30) as usize),
88            sample_count_product: 0,
89            scaled_output_frequency: output_frequency * RESAMPLE_SCALING_FACTOR,
90            scaled_source_frequency: Self::scale_frequency(source_frequency),
91            _marker: PhantomData,
92        }
93    }
94
95    fn scale_frequency(source_frequency: f64) -> u64 {
96        (source_frequency * RESAMPLE_SCALING_FACTOR as f64).round() as u64
97    }
98
99    #[inline]
100    pub fn collect(&mut self, samples: [f64; CHANNELS]) {
101        for (ch, sample) in samples.into_iter().enumerate() {
102            self.input[ch].push(sample);
103        }
104
105        self.sample_count_product += self.scaled_output_frequency;
106        while self.sample_count_product >= self.scaled_source_frequency {
107            self.sample_count_product -= self.scaled_source_frequency;
108
109            let output_samples =
110                apply_fir_filter(array::from_fn(|ch| &self.input[ch]), Kernel::lpf_coefficients());
111            self.output.push_back(output_samples);
112        }
113    }
114
115    #[inline]
116    #[must_use]
117    pub fn output_buffer_len(&self) -> usize {
118        self.output.len()
119    }
120
121    #[inline]
122    pub fn output_buffer_pop_front(&mut self) -> Option<[f64; CHANNELS]> {
123        self.output.pop_front()
124    }
125
126    #[inline]
127    pub fn update_output_frequency(&mut self, output_frequency: f64) {
128        self.scaled_output_frequency = Self::scale_frequency(output_frequency);
129    }
130
131    #[inline]
132    pub fn update_source_frequency(&mut self, source_frequency: f64) {
133        self.scaled_source_frequency = Self::scale_frequency(source_frequency);
134    }
135}
136
137#[multiversion(targets("x86_64+sse4.2", "x86_64+avx2+fma", "x86_64+avx512f"))]
138fn apply_fir_filter<const N: usize, const CHANNELS: usize>(
139    samples: [&RingBuffer<N>; CHANNELS],
140    coefficients: &LpfCoefficients<N>,
141) -> [f64; CHANNELS] {
142    if samples[0].len >= N {
143        let mut sums = [0.0_f64; CHANNELS];
144        for (i, coefficient) in coefficients.iter().copied().enumerate() {
145            let input_idx = samples[0].idx + i;
146            for (ch, sum) in sums.iter_mut().enumerate() {
147                *sum = sum.algebraic_add(coefficient.algebraic_mul(samples[ch].buffer[input_idx]));
148            }
149        }
150        sums
151    } else {
152        let mut sums = [0.0_f64; CHANNELS];
153        for i in N - samples[0].len..N {
154            let input_idx = samples[0].idx + i - (N - samples[0].len);
155            for (ch, sum) in sums.iter_mut().enumerate() {
156                *sum =
157                    sum.algebraic_add(coefficients[i].algebraic_mul(samples[ch].buffer[input_idx]));
158            }
159        }
160        sums
161    }
162}
163
164pub type MonoFirResampler<const LPF_TAPS: usize, Kernel> = FirResampler<1, LPF_TAPS, Kernel>;
165pub type StereoFirResampler<const LPF_TAPS: usize, Kernel> = FirResampler<2, LPF_TAPS, Kernel>;
166
167#[cfg(test)]
168mod tests {
169    use super::*;
170
171    #[test]
172    fn ring_buffer_basic() {
173        let mut buffer = RingBuffer::<3>::new();
174        assert_eq!(buffer.idx, buffer.buffer.len());
175        assert_eq!(buffer.len, 0);
176
177        buffer.push(3.0);
178        assert_eq!(buffer.idx, buffer.buffer.len() - 1);
179        assert_eq!(buffer.len, 1);
180        assert_eq!(buffer.buffer[buffer.idx], 3.0);
181
182        buffer.push(5.0);
183        assert_eq!(buffer.idx, buffer.buffer.len() - 2);
184        assert_eq!(buffer.len, 2);
185        assert_eq!(&buffer.buffer[buffer.idx..buffer.idx + 2], &[5.0, 3.0]);
186
187        buffer.push(7.0);
188        assert_eq!(buffer.idx, buffer.buffer.len() - 3);
189        assert_eq!(buffer.len, 3);
190        assert_eq!(&buffer.buffer[buffer.idx..buffer.idx + 3], &[7.0, 5.0, 3.0]);
191
192        // Buffer is now full; next push should move the starting point but not increase length
193        buffer.push(9.0);
194        assert_eq!(buffer.idx, buffer.buffer.len() - 4);
195        assert_eq!(buffer.len, 3);
196        assert_eq!(&buffer.buffer[buffer.idx..buffer.idx + 3], &[9.0, 7.0, 5.0]);
197
198        // Push one more
199        buffer.push(11.0);
200        assert_eq!(buffer.idx, buffer.buffer.len() - 5);
201        assert_eq!(buffer.len, 3);
202        assert_eq!(&buffer.buffer[buffer.idx..buffer.idx + 3], &[11.0, 9.0, 7.0]);
203    }
204
205    #[test]
206    fn ring_buffer_wrap() {
207        const N: usize = 4;
208
209        let mut buffer = RingBuffer::<N>::new();
210        for i in 0..buffer.buffer.len() {
211            buffer.buffer[i] = (i + 5) as f64;
212        }
213        buffer.idx = 1;
214        buffer.len = N;
215
216        let current: [f64; N] = buffer.buffer[1..=N].try_into().unwrap();
217
218        // Last push before buffer is full
219        buffer.push(54321.0);
220        assert_eq!(buffer.idx, 0);
221        assert_eq!(buffer.len, N);
222        assert_eq!(&buffer.buffer[0..N], &[54321.0, current[0], current[1], current[2]]);
223
224        // Push while buffer is full should copy contents to the end of the buffer
225        buffer.push(56789.0);
226        assert_eq!(buffer.idx, buffer.buffer.len() - N);
227        assert_eq!(buffer.len, N);
228        assert_eq!(&buffer.buffer[buffer.idx..], &[56789.0, 54321.0, current[0], current[1]]);
229
230        buffer.push(12345.0);
231        assert_eq!(buffer.idx, buffer.buffer.len() - N - 1);
232        assert_eq!(buffer.len, N);
233        assert_eq!(
234            &buffer.buffer[buffer.idx..buffer.idx + N],
235            &[12345.0, 56789.0, 54321.0, current[0]]
236        );
237    }
238}