1use crate::iir::IirFilter;
2use num::complex::{Complex64, ComplexFloat};
3use num::{One, Zero};
4use std::array;
5
6#[derive(Debug, Clone, Copy, PartialEq, Eq)]
7pub enum FilterType {
8    LowPass,
9    HighPass,
10}
11
12#[must_use]
13pub fn butterworth<const N: usize>(fc: f64, fs: f64, filter: FilterType) -> IirFilter<N> {
14    let wc = fc / (fs / 2.0);
15    if !(0.0..=1.0).contains(&wc) {
16        log::error!(
17            "Attempted to design order {N} {filter:?} Butterworth filter with invalid frequencies, replacing with identity filter: fc={fc}, fs={fs}"
18        );
19        return IirFilter::identity();
20    }
21
22    let (b, a) = butterworth_coefficients::<N>(fc, fs, filter);
23    IirFilter::new(&b, &a)
24}
25
26#[allow(clippy::many_single_char_names)]
27fn butterworth_coefficients<const N: usize>(
28    fc: f64,
29    fs: f64,
30    filter: FilterType,
31) -> (Vec<f64>, Vec<f64>) {
32    use std::f64::consts::{E, PI};
33
34    let n = N as f64;
35    let j = Complex64::i();
36
37    // Compute Butterworth poles for low-pass prototype
38    let poles: [_; N] = array::from_fn(|i| {
39        let k = (i + 1) as f64;
40        E.powc(j * PI * (2.0 * k + n - 1.0) / (2.0 * n))
41    });
42
43    // Warp analog frequency and convert low-pass prototype poles to poles for desired filter type
44    let wc = fc / (fs / 2.0);
45    let warp = 2.0 * (wc * PI / 2.0).tan();
46    let poles = match filter {
47        FilterType::LowPass => poles.map(|p| warp * p),
48        FilterType::HighPass => poles.map(|p| warp / p),
49    };
50
51    // Perform bilinear transform
52    let poles = poles.map(|p| (1.0 + p / 2.0) / (1.0 - p / 2.0));
53
54    // Compute base feedforward coefficients
55    let zeroes = match filter {
56        FilterType::LowPass => [Complex64::new(-1.0, 0.0); N],
57        FilterType::HighPass => [Complex64::new(1.0, 0.0); N],
58    };
59    let b = polynomial_coefficients(zeroes);
60
61    // Compute feedback coefficients
62    let a = polynomial_coefficients(poles);
63
64    // Normalize feedforward coefficients
65    let k = match filter {
66        FilterType::LowPass => a.iter().copied().sum::<f64>() / b.iter().copied().sum::<f64>(),
67        FilterType::HighPass => {
68            let high_pass_sum = |arr: &[f64]| {
69                arr.iter()
70                    .copied()
71                    .enumerate()
72                    .map(|(i, n)| (-1.0_f64).powi(i as i32) * n)
73                    .sum::<f64>()
74            };
75
76            high_pass_sum(&a) / high_pass_sum(&b)
77        }
78    };
79    let b: Vec<_> = b.into_iter().map(|b| b * k).collect();
80
81    log::debug!("Filter for fc={fc}, fs={fs}, type {filter:?}:");
82    log::debug!("  b={b:?}");
83    log::debug!("  a={a:?}");
84
85    (b, a)
86}
87
88fn polynomial_coefficients<const N: usize>(roots: [Complex64; N]) -> Vec<f64> {
89    (0..=N)
90        .map(|i| {
91            let sign = (-1.0_f64).powi(i as i32);
92            (sign * sum_combinations(Complex64::one(), &roots, i)).re
93        })
94        .collect()
95}
96
97fn sum_combinations(product: Complex64, roots: &[Complex64], len: usize) -> Complex64 {
98    if len == 0 {
99        return product;
100    }
101
102    if roots.len() < len {
103        return Complex64::zero();
104    }
105
106    (0..roots.len()).map(|i| sum_combinations(product * roots[i], &roots[i + 1..], len - 1)).sum()
107}
108
109#[cfg(test)]
110mod tests {
111    use super::*;
112    use std::iter;
113
114    fn float_slice_equal(a: &[f64], b: &[f64]) -> bool {
115        if a.len() != b.len() {
116            return false;
117        }
118
119        iter::zip(a, b).all(|(&a_elem, &b_elem)| (a_elem - b_elem).abs() < 1e-9)
120    }
121
122    fn assert_float_slice_eq(a: &[f64], b: &[f64]) {
123        assert!(float_slice_equal(a, b), "float slices not equal: {a:?} {b:?}");
124    }
125
126    #[test]
127    fn butterworth_low_pass() {
128        // Expected filters generated using Python w/ scipy:
129        //   b, a = butter(n, 3390 / (53693175 / 7 / 6 / 24 / 2), btype="lowpass")
130
131        const B1: &[f64] = &[0.1684983368367697, 0.1684983368367697];
132        const A1: &[f64] = &[1.0, -0.6630033263264605];
133
134        const B2: &[f64] = &[0.030930211590861196, 0.06186042318172239, 0.030930211590861196];
135        const A2: &[f64] = &[1.0, -1.4445658935949237, 0.5682867399583684];
136
137        const B3: &[f64] = &[
138            0.005563425334839113,
139            0.016690276004517342,
140            0.016690276004517342,
141            0.005563425334839113,
142        ];
143        const A3: &[f64] = &[1.0, -2.2050627555803564, 1.696520707984555, -0.44695054972548554];
144
145        let fc = 3390.0;
146        let fs = 53693175.0 / 7.0 / 6.0 / 24.0;
147
148        let (b1, a1) = butterworth_coefficients::<1>(fc, fs, FilterType::LowPass);
149        assert_float_slice_eq(&b1, B1);
150        assert_float_slice_eq(&a1, A1);
151
152        let (b2, a2) = butterworth_coefficients::<2>(fc, fs, FilterType::LowPass);
153        assert_float_slice_eq(&b2, B2);
154        assert_float_slice_eq(&a2, A2);
155
156        let (b3, a3) = butterworth_coefficients::<3>(fc, fs, FilterType::LowPass);
157        assert_float_slice_eq(&b3, B3);
158        assert_float_slice_eq(&a3, A3);
159    }
160
161    #[test]
162    fn butterworth_high_pass() {
163        // Expected filters generated using Python w/ scipy:
164        //   b, a = butter(n, 3390 / (53693175 / 7 / 6 / 24 / 2), btype="highpass")
165
166        const B1: &[f64] = &[0.8315016631632303, -0.8315016631632303];
167        const A1: &[f64] = &[1.0, -0.6630033263264605];
168
169        const B2: &[f64] = &[0.753213158388323, -1.506426316776646, 0.753213158388323];
170        const A2: &[f64] = &[1.0, -1.4445658935949237, 0.5682867399583684];
171
172        const B3: &[f64] =
173            &[0.6685667516612996, -2.005700254983899, 2.005700254983899, -0.6685667516612996];
174        const A3: &[f64] = &[1.0, -2.2050627555803564, 1.696520707984555, -0.4469505497254855];
175
176        let fc = 3390.0;
177        let fs = 53693175.0 / 7.0 / 6.0 / 24.0;
178
179        let (b1, a1) = butterworth_coefficients::<1>(fc, fs, FilterType::HighPass);
180        assert_float_slice_eq(&b1, B1);
181        assert_float_slice_eq(&a1, A1);
182
183        let (b2, a2) = butterworth_coefficients::<2>(fc, fs, FilterType::HighPass);
184        assert_float_slice_eq(&b2, B2);
185        assert_float_slice_eq(&a2, A2);
186
187        let (b3, a3) = butterworth_coefficients::<3>(fc, fs, FilterType::HighPass);
188        assert_float_slice_eq(&b3, B3);
189        assert_float_slice_eq(&a3, A3);
190    }
191}