use libm::Libm;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum Oversample {
#[default]
Off,
X2,
X4,
}
impl Oversample {
#[inline]
pub fn factor(self) -> usize {
match self {
Oversample::Off => 1,
Oversample::X2 => 2,
Oversample::X4 => 4,
}
}
}
const MAX_TAPS: usize = 63;
const MAX_INPUT_TAPS: usize = 16;
#[derive(Debug, Clone)]
pub struct Oversampler {
factor: usize,
len: usize,
in_len: usize,
h: [f64; MAX_TAPS],
in_ring: [f64; MAX_INPUT_TAPS],
in_pos: usize,
os_ring: [f64; MAX_TAPS],
os_pos: usize,
}
impl Oversampler {
pub fn new(mode: Oversample) -> Self {
let factor = mode.factor();
let (len, cutoff) = match mode {
Oversample::Off => (1, 0.5),
Oversample::X2 => (31, 0.25),
Oversample::X4 => (63, 0.125),
};
let mut h = [0.0; MAX_TAPS];
if factor > 1 {
design_lowpass(&mut h[..len], cutoff);
}
let in_len = if factor > 1 {
len.div_ceil(factor).clamp(1, MAX_INPUT_TAPS)
} else {
1
};
Self {
factor,
len,
in_len,
h,
in_ring: [0.0; MAX_INPUT_TAPS],
in_pos: 0,
os_ring: [0.0; MAX_TAPS],
os_pos: 0,
}
}
#[inline]
pub fn factor(&self) -> usize {
self.factor
}
pub fn reset(&mut self) {
self.in_ring = [0.0; MAX_INPUT_TAPS];
self.os_ring = [0.0; MAX_TAPS];
self.in_pos = 0;
self.os_pos = 0;
}
#[inline]
pub fn process<F: FnMut(f64) -> f64>(&mut self, x: f64, mut f: F) -> f64 {
if self.factor == 1 {
return f(x);
}
let n = self.factor;
let len = self.len;
self.in_pos = (self.in_pos + 1) % self.in_len;
self.in_ring[self.in_pos] = x;
for p in 0..n {
let mut acc = 0.0;
let mut tap = p;
let mut i = 0;
while tap < len {
let idx = (self.in_pos + self.in_len - i) % self.in_len;
acc += self.h[tap] * self.in_ring[idx];
tap += n;
i += 1;
}
let up = acc * (n as f64);
let processed = f(up);
self.os_pos = (self.os_pos + 1) % len;
self.os_ring[self.os_pos] = processed;
}
let center = (self.os_pos + len - (n - 1)) % len;
let mut out = 0.0;
for (j, &hj) in self.h[..len].iter().enumerate() {
let idx = (center + len - j) % len;
out += hj * self.os_ring[idx];
}
out
}
}
fn design_lowpass(h: &mut [f64], cutoff: f64) {
let l = h.len();
let m = (l - 1) as f64 / 2.0;
let pi = core::f64::consts::PI;
let mut sum = 0.0;
for (n, hn) in h.iter_mut().enumerate() {
let x = n as f64 - m;
let sinc = if Libm::<f64>::fabs(x) < 1e-9 {
2.0 * cutoff
} else {
Libm::<f64>::sin(2.0 * pi * cutoff * x) / (pi * x)
};
let nn = n as f64;
let denom = (l - 1) as f64;
let window = 0.42 - 0.5 * Libm::<f64>::cos(2.0 * pi * nn / denom)
+ 0.08 * Libm::<f64>::cos(4.0 * pi * nn / denom);
*hn = sinc * window;
sum += *hn;
}
if sum.abs() > 1e-12 {
for hn in h.iter_mut() {
*hn /= sum;
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_oversample_factor() {
assert_eq!(Oversample::Off.factor(), 1);
assert_eq!(Oversample::X2.factor(), 2);
assert_eq!(Oversample::X4.factor(), 4);
assert_eq!(Oversample::default(), Oversample::Off);
}
#[test]
fn test_off_is_transparent() {
let mut os = Oversampler::new(Oversample::Off);
for &x in &[0.0, 1.0, -0.5, 3.3, -2.2] {
let y = os.process(x, |v| v);
assert!((y - x).abs() < 1e-12);
}
}
#[test]
fn test_dc_gain_unity_2x() {
let mut os = Oversampler::new(Oversample::X2);
let mut last = 0.0;
for _ in 0..2000 {
last = os.process(1.0, |v| v);
}
assert!((last - 1.0).abs() < 1e-3, "2x DC gain not unity: {last}");
}
#[test]
fn test_dc_gain_unity_4x() {
let mut os = Oversampler::new(Oversample::X4);
let mut last = 0.0;
for _ in 0..4000 {
last = os.process(1.0, |v| v);
}
assert!((last - 1.0).abs() < 1e-3, "4x DC gain not unity: {last}");
}
#[test]
fn test_passband_sine_preserved_2x() {
let sr = 44100.0;
let freq = 200.0;
let mut os = Oversampler::new(Oversample::X2);
let n = 4000;
let mut max_out: f64 = 0.0;
for i in 0..n {
let t = i as f64 / sr;
let x = Libm::<f64>::sin(2.0 * core::f64::consts::PI * freq * t);
let y = os.process(x, |v| v);
if i > 2000 {
max_out = max_out.max(y.abs());
}
}
assert!(
(max_out - 1.0).abs() < 0.05,
"passband sine amplitude altered: {max_out}"
);
}
#[test]
fn test_reset_clears_state() {
let mut os = Oversampler::new(Oversample::X4);
for _ in 0..100 {
os.process(1.0, |v| v);
}
os.reset();
let y = os.process(0.0, |v| v);
assert!(y.abs() < 1e-9);
}
}