use num_traits::FloatConst;
use crate::operations::iir_filtering::IirFilter;
use crate::operations::traits::AudioParametricEq;
use crate::operations::types::{EqBand, EqBandType, ParametricEq, ThreeBandEqConfig};
use crate::traits::StandardSample;
use crate::utils::audio_math::db_to_amplitude as db_to_linear;
use crate::{AudioData, AudioSampleError, LayoutError};
use crate::{AudioSampleResult, AudioSamples, ConvertTo};
impl<T> AudioParametricEq for AudioSamples<'_, T>
where
T: StandardSample,
{
#[inline]
fn apply_parametric_eq_in_place(&mut self, eq: &ParametricEq) -> AudioSampleResult<()> {
let sample_rate = self.sample_rate_hz();
if eq.is_bypassed() {
return Ok(());
}
let eq = eq.clone().validate(sample_rate)?;
for band in &eq.bands {
if band.is_enabled() {
self.apply_eq_band_in_place(band)?;
}
}
if eq.output_gain_db != 0.0 {
let output_gain_linear = db_to_linear(eq.output_gain_db);
self.apply_linear_gain(output_gain_linear);
}
Ok(())
}
#[inline]
fn apply_eq_band_in_place(&mut self, band: &EqBand) -> AudioSampleResult<()> {
let sample_rate = self.sample_rate_hz();
if !band.is_enabled() {
return Ok(());
}
band.validate(sample_rate)?;
let (b_coeffs, a_coeffs) = design_eq_band_filter(band, sample_rate);
let mut filter = IirFilter::new(b_coeffs, a_coeffs);
match self.data_mut() {
AudioData::Mono(_) => {
let Some(working_samples) = self.as_slice_mut() else {
return Err(AudioSampleError::Layout(LayoutError::NonContiguous {
operation: "parametric EQ".to_string(),
layout_type: "non-contiguous mono samples".to_string(),
}));
};
for sample in working_samples.iter_mut() {
let x: f64 = (*sample).convert_to();
*sample = filter.process_sample(x).convert_to();
}
}
AudioData::Multi(samples) => {
let mut view = samples.view_mut();
for mut channel in view.outer_iter_mut() {
for sample in channel.iter_mut() {
let x: f64 = (*sample).convert_to();
*sample = filter.process_sample(x).convert_to();
}
filter.reset();
}
}
}
Ok(())
}
#[inline]
fn apply_peak_filter_in_place(
&mut self,
frequency: f64,
gain_db: f64,
q_factor: f64,
) -> AudioSampleResult<()> {
let band = EqBand::peak(frequency, gain_db, q_factor);
self.apply_eq_band_in_place(&band)
}
#[inline]
fn apply_low_shelf_in_place(
&mut self,
frequency: f64,
gain_db: f64,
q_factor: f64,
) -> AudioSampleResult<()> {
let band = EqBand::low_shelf(frequency, gain_db, q_factor);
self.apply_eq_band_in_place(&band)
}
#[inline]
fn apply_high_shelf_in_place(
&mut self,
frequency: f64,
gain_db: f64,
q_factor: f64,
) -> AudioSampleResult<()> {
let band = EqBand::high_shelf(frequency, gain_db, q_factor);
self.apply_eq_band_in_place(&band)
}
#[inline]
fn apply_three_band_eq_in_place(
&mut self,
config: &ThreeBandEqConfig,
) -> AudioSampleResult<()> {
config.validate()?;
let eq = ParametricEq::three_band(
config.low_freq,
config.low_gain,
config.mid_freq,
config.mid_gain,
config.mid_q,
config.high_freq,
config.high_gain,
);
self.apply_parametric_eq_in_place(&eq)
}
#[inline]
fn eq_frequency_response(
&self,
eq: &ParametricEq,
frequencies: &[f64],
) -> AudioSampleResult<(Vec<f64>, Vec<f64>)> {
let mut combined_magnitude = vec![1.0; frequencies.len()];
let mut combined_phase = vec![0.0; frequencies.len()];
let sample_rate = self.sample_rate().get();
for band in &eq.bands {
if !band.is_enabled() {
continue;
}
let (b_coeffs, a_coeffs) = design_eq_band_filter(band, f64::from(sample_rate));
let filter = IirFilter::new(b_coeffs, a_coeffs);
let (magnitude, phase) = filter.frequency_response(frequencies, f64::from(sample_rate));
for i in 0..frequencies.len() {
combined_magnitude[i] *= magnitude[i];
combined_phase[i] += phase[i];
}
}
if eq.output_gain_db != 0.0 {
let output_gain_linear = db_to_linear(eq.output_gain_db);
for magnitude in &mut combined_magnitude {
*magnitude *= output_gain_linear;
}
}
Ok((combined_magnitude, combined_phase))
}
}
fn design_eq_band_filter(band: &EqBand, sample_rate: f64) -> (Vec<f64>, Vec<f64>) {
match band.band_type {
EqBandType::Peak => {
design_peak_filter(band.frequency, band.gain_db, band.q_factor, sample_rate)
}
EqBandType::LowShelf => {
design_low_shelf_filter(band.frequency, band.gain_db, band.q_factor, sample_rate)
}
EqBandType::HighShelf => {
design_high_shelf_filter(band.frequency, band.gain_db, band.q_factor, sample_rate)
}
EqBandType::LowPass => design_lowpass_filter(band.frequency, band.q_factor, sample_rate),
EqBandType::HighPass => design_highpass_filter(band.frequency, band.q_factor, sample_rate),
EqBandType::BandPass => design_bandpass_filter(band.frequency, band.q_factor, sample_rate),
EqBandType::BandStop => design_bandstop_filter(band.frequency, band.q_factor, sample_rate),
}
}
fn design_peak_filter(
frequency: f64,
gain_db: f64,
q_factor: f64,
sample_rate: f64,
) -> (Vec<f64>, Vec<f64>) {
let a = 10.0f64.powf(gain_db / 40.0); let omega = 2.0 * std::f64::consts::PI * frequency / sample_rate;
let (sin_omega, cos_omega) = omega.sin_cos();
let alpha = sin_omega / (2.0 * q_factor);
let b0 = 1.0 + alpha * a;
let b1 = -2.0 * cos_omega;
let b2 = 1.0 - alpha * a;
let a0 = 1.0 + alpha / a;
let a1 = -2.0 * cos_omega;
let a2 = 1.0 - alpha / a;
let b_coeffs = vec![b0 / a0, b1 / a0, b2 / a0];
let a_coeffs = vec![1.0, a1 / a0, a2 / a0];
(b_coeffs, a_coeffs)
}
fn design_low_shelf_filter(
frequency: f64,
gain_db: f64,
q_factor: f64,
sample_rate: f64,
) -> (Vec<f64>, Vec<f64>) {
let a = 10.0f64.powf(gain_db / 40.0); let omega = 2.0 * std::f64::consts::PI * frequency / sample_rate;
let (sin_omega, cos_omega) = omega.sin_cos();
let alpha = sin_omega / (2.0 * q_factor);
let sqrt_2a = (2.0 * a).sqrt();
let b0 = a * ((a - 1.0).mul_add(-cos_omega, a + 1.0) + sqrt_2a * alpha);
let b1 = 2.0 * a * (a + 1.0).mul_add(-cos_omega, a - 1.0);
let b2 = a * ((a - 1.0).mul_add(-cos_omega, a + 1.0) - sqrt_2a * alpha);
let a0 = (a - 1.0).mul_add(cos_omega, a + 1.0) + sqrt_2a * alpha;
let a1 = -2.0 * (a + 1.0).mul_add(cos_omega, a - 1.0);
let a2 = (a - 1.0).mul_add(cos_omega, a + 1.0) - sqrt_2a * alpha;
let b_coeffs = vec![b0 / a0, b1 / a0, b2 / a0];
let a_coeffs = vec![1.0, a1 / a0, a2 / a0];
(b_coeffs, a_coeffs)
}
fn design_high_shelf_filter(
frequency: f64,
gain_db: f64,
q_factor: f64,
sample_rate: f64,
) -> (Vec<f64>, Vec<f64>) {
let a = 10.0f64.powf(gain_db / 40.0); let omega = 2.0 * f64::PI() * frequency / sample_rate;
let (sin_omega, cos_omega) = omega.sin_cos();
let alpha = sin_omega / (2.0 * q_factor);
let sqrt_2a = (2.0 * a).sqrt();
let b0 = a * ((a - 1.0).mul_add(cos_omega, a + 1.0) + sqrt_2a * alpha);
let b1 = -2.0 * a * (a + 1.0).mul_add(cos_omega, a - 1.0);
let b2 = a * ((a - 1.0).mul_add(cos_omega, a + 1.0) - sqrt_2a * alpha);
let a0 = (a - 1.0).mul_add(-cos_omega, a + 1.0) + sqrt_2a * alpha;
let a1 = 2.0 * (a + 1.0).mul_add(-cos_omega, a - 1.0);
let a2 = (a - 1.0).mul_add(-cos_omega, a + 1.0) - sqrt_2a * alpha;
let b_coeffs = vec![b0 / a0, b1 / a0, b2 / a0];
let a_coeffs = vec![1.0, a1 / a0, a2 / a0];
(b_coeffs, a_coeffs)
}
fn design_lowpass_filter(frequency: f64, q_factor: f64, sample_rate: f64) -> (Vec<f64>, Vec<f64>) {
let omega = 2.0 * std::f64::consts::PI * frequency / sample_rate;
let (sin_omega, cos_omega) = omega.sin_cos();
let alpha = sin_omega / (2.0 * q_factor);
let b0 = (1.0 - cos_omega) / 2.0;
let b1 = 1.0 - cos_omega;
let b2 = (1.0 - cos_omega) / 2.0;
let a0 = 1.0 + alpha;
let a1 = -2.0 * cos_omega;
let a2 = 1.0 - alpha;
let b_coeffs = vec![b0 / a0, b1 / a0, b2 / a0];
let a_coeffs = vec![1.0, a1 / a0, a2 / a0];
(b_coeffs, a_coeffs)
}
fn design_highpass_filter(frequency: f64, q_factor: f64, sample_rate: f64) -> (Vec<f64>, Vec<f64>) {
let omega = 2.0 * std::f64::consts::PI * frequency / sample_rate;
let (sin_omega, cos_omega) = omega.sin_cos();
let alpha = sin_omega / (2.0 * q_factor);
let b0 = f64::midpoint(1.0, cos_omega);
let b1 = -(1.0 + cos_omega);
let b2 = f64::midpoint(1.0, cos_omega);
let a0 = 1.0 + alpha;
let a1 = -2.0 * cos_omega;
let a2 = 1.0 - alpha;
let b_coeffs = vec![b0 / a0, b1 / a0, b2 / a0];
let a_coeffs = vec![1.0, a1 / a0, a2 / a0];
(b_coeffs, a_coeffs)
}
fn design_bandpass_filter(frequency: f64, q_factor: f64, sample_rate: f64) -> (Vec<f64>, Vec<f64>) {
let omega = 2.0 * f64::PI() * frequency / sample_rate;
let (sin_omega, cos_omega) = omega.sin_cos();
let alpha = sin_omega / (2.0 * q_factor);
let b0 = alpha;
let b1 = 0.0;
let b2 = -alpha;
let a0 = 1.0 + alpha;
let a1 = -2.0 * cos_omega;
let a2 = 1.0 - alpha;
let b_coeffs = vec![b0 / a0, b1 / a0, b2 / a0];
let a_coeffs = vec![1.0, a1 / a0, a2 / a0];
(b_coeffs, a_coeffs)
}
fn design_bandstop_filter(frequency: f64, q_factor: f64, sample_rate: f64) -> (Vec<f64>, Vec<f64>) {
let omega = 2.0 * std::f64::consts::PI * frequency / sample_rate;
let (sin_omega, cos_omega) = omega.sin_cos();
let alpha = sin_omega / (2.0 * q_factor);
let b0 = 1.0;
let b1 = -2.0 * cos_omega;
let b2 = 1.0;
let a0 = 1.0 + alpha;
let a1 = -2.0 * cos_omega;
let a2 = 1.0 - alpha;
let b_coeffs = vec![b0 / a0, b1 / a0, b2 / a0];
let a_coeffs = vec![1.0, a1 / a0, a2 / a0];
(b_coeffs, a_coeffs)
}
impl<T> AudioSamples<'_, T>
where
T: StandardSample,
{
fn apply_linear_gain(&mut self, gain: f64) {
self.apply(|x| {
let x_f: f64 = T::cast_into(x);
let y_f = x_f * gain;
T::cast_from(y_f)
});
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::operations::traits::AudioParametricEq;
use crate::sample_rate;
use crate::utils::audio_math::amplitude_to_db as linear_to_db;
use non_empty_slice::{NonEmptyVec, non_empty_vec};
use std::f64::consts::PI;
#[test]
fn test_apply_eq_band_multichannel_matches_per_channel_mono() {
let sr = 44100.0;
let n = 256usize;
let ch0: Vec<f32> = (0..n)
.map(|i| (2.0 * PI * 440.0 * (i as f64 / sr)).sin() as f32)
.collect();
let ch1: Vec<f32> = (0..n)
.map(|i| 0.5 * (2.0 * PI * 880.0 * (i as f64 / sr)).sin() as f32)
.collect();
let stereo_arr =
ndarray::Array2::from_shape_fn((2, n), |(c, i)| if c == 0 { ch0[i] } else { ch1[i] });
let mut stereo = AudioSamples::new_multi_channel(stereo_arr, sample_rate!(44100)).unwrap();
let band = EqBand::peak(880.0, 6.0, 2.0);
stereo.apply_eq_band_in_place(&band).unwrap();
let mut mono0: AudioSamples<'_, f32> =
AudioSamples::from_mono_vec(NonEmptyVec::new(ch0).unwrap(), sample_rate!(44100));
let mut mono1: AudioSamples<'_, f32> =
AudioSamples::from_mono_vec(NonEmptyVec::new(ch1).unwrap(), sample_rate!(44100));
mono0.apply_eq_band_in_place(&band).unwrap();
mono1.apply_eq_band_in_place(&band).unwrap();
let out = stereo.as_multi_channel().unwrap();
let r0 = mono0.as_slice().unwrap();
let r1 = mono1.as_slice().unwrap();
for i in 0..n {
assert!((out[(0, i)] - r0[i]).abs() < 1e-6, "ch0 sample {i}");
assert!((out[(1, i)] - r1[i]).abs() < 1e-6, "ch1 sample {i}");
}
}
#[test]
fn test_peak_filter() {
let sample_rate = 44100.0;
let duration = 0.1;
let samples_count = (sample_rate * duration) as usize;
let mut samples = Vec::new();
for i in 0..samples_count {
let t = i as f64 / sample_rate;
let value = (2.0 * PI * 440.0 * t).sin()
+ (2.0 * PI * 880.0 * t).sin()
+ (2.0 * PI * 1760.0 * t).sin();
samples.push(value as f32);
}
let samples = NonEmptyVec::new(samples).unwrap();
let mut audio: AudioSamples<'_, f32> =
AudioSamples::from_mono_vec(samples, sample_rate!(44100));
let result = audio.apply_peak_filter_in_place(880.0, 6.0, 2.0);
assert!(result.is_ok());
}
#[test]
fn test_low_shelf_filter() {
let sample_rate = 44100.0;
let duration = 0.1;
let samples_count = (sample_rate * duration) as usize;
let mut samples = Vec::new();
for i in 0..samples_count {
let t = i as f64 / sample_rate;
let value = (2.0 * PI * 100.0 * t).sin()
+ (2.0 * PI * 1000.0 * t).sin()
+ (2.0 * PI * 5000.0 * t).sin();
samples.push(value as f32);
}
let samples = NonEmptyVec::new(samples).unwrap();
let mut audio: AudioSamples<'_, f32> =
AudioSamples::from_mono_vec(samples, sample_rate!(44100));
let result = audio.apply_low_shelf_in_place(500.0, -3.0, 0.707);
assert!(result.is_ok());
}
#[test]
fn test_high_shelf_filter() {
let sample_rate = 44100.0;
let duration = 0.1;
let samples_count = (sample_rate * duration) as usize;
let mut samples = Vec::new();
for i in 0..samples_count {
let t = i as f64 / sample_rate;
let value = (2.0 * PI * 100.0 * t).sin()
+ (2.0 * PI * 1000.0 * t).sin()
+ (2.0 * PI * 5000.0 * t).sin();
samples.push(value as f32);
}
let samples = NonEmptyVec::new(samples).unwrap();
let mut audio: AudioSamples<'_, f32> =
AudioSamples::from_mono_vec(samples, sample_rate!(44100));
let result = audio.apply_high_shelf_in_place(2000.0, 4.0, 0.707);
assert!(result.is_ok());
}
#[test]
fn test_three_band_eq() {
let sample_rate = 44100.0;
let duration = 0.1;
let samples_count = (sample_rate * duration) as usize;
let mut samples = Vec::new();
for i in 0..samples_count {
let t = i as f64 / sample_rate;
let value = (2.0 * PI * 100.0 * t).sin()
+ (2.0 * PI * 1000.0 * t).sin()
+ (2.0 * PI * 5000.0 * t).sin();
samples.push(value as f32);
}
let samples = NonEmptyVec::new(samples).unwrap();
let mut audio: AudioSamples<'_, f32> =
AudioSamples::from_mono_vec(samples, sample_rate!(44100));
let config = ThreeBandEqConfig::new(200.0, -2.0, 1000.0, 3.0, 2.0, 4000.0, 1.0);
let result = audio.apply_three_band_eq_in_place(&config);
assert!(result.is_ok());
}
#[test]
fn test_three_band_eq_config_validate_rejects_bad_input() {
assert!(
ThreeBandEqConfig::new(0.0, 0.0, 1000.0, 0.0, 1.0, 4000.0, 0.0)
.validate()
.is_err()
);
assert!(
ThreeBandEqConfig::new(2000.0, 0.0, 1000.0, 0.0, 1.0, 4000.0, 0.0)
.validate()
.is_err()
);
assert!(
ThreeBandEqConfig::new(200.0, 0.0, 1000.0, 0.0, 0.0, 4000.0, 0.0)
.validate()
.is_err()
);
assert!(ThreeBandEqConfig::flat().validate().is_ok());
}
#[test]
fn test_parametric_eq_configuration() {
let mut audio: AudioSamples<'_, f32> =
AudioSamples::from_mono_vec(non_empty_vec![1.0f32, 0.0, -1.0], sample_rate!(44100));
let mut eq = ParametricEq::new();
eq.add_band(EqBand::peak(1000.0, 3.0, 2.0));
eq.add_band(EqBand::low_shelf(100.0, -2.0, 0.707));
eq.set_output_gain(1.0);
let result = audio.apply_parametric_eq_in_place(&eq);
assert!(result.is_ok());
assert_eq!(eq.band_count(), 2);
assert_eq!(eq.output_gain_db, 1.0);
assert!(!eq.is_bypassed());
}
#[test]
fn test_eq_band_validation() {
let sample_rate = 44100.0;
let valid_band = EqBand::peak(1000.0, 3.0, 2.0);
assert!(valid_band.validate(sample_rate).is_ok());
let invalid_band = EqBand::peak(sample_rate, 3.0, 2.0);
assert!(invalid_band.validate(sample_rate).is_err());
let invalid_band = EqBand::peak(1000.0, 3.0, 0.0);
assert!(invalid_band.validate(sample_rate).is_err());
let extreme_band = EqBand::peak(1000.0, 50.0, 2.0);
assert!(extreme_band.validate(sample_rate).is_err());
}
#[test]
fn test_eq_band_enable_disable() {
let mut band = EqBand::peak(1000.0, 3.0, 2.0);
assert!(band.is_enabled());
band.set_enabled(false);
assert!(!band.is_enabled());
band.set_enabled(true);
assert!(band.is_enabled());
}
#[test]
fn test_parametric_eq_bypass() {
let mut audio: AudioSamples<'_, f32> =
AudioSamples::from_mono_vec(non_empty_vec![1.0f32, 0.5, -0.5], sample_rate!(44100));
let original_samples = audio.data().clone();
let mut eq = ParametricEq::new();
eq.add_band(EqBand::peak(1000.0, 10.0, 2.0)); eq.set_bypassed(true);
let result = audio.apply_parametric_eq_in_place(&eq);
assert!(result.is_ok());
match (audio.data(), &original_samples) {
(AudioData::Mono(new), AudioData::Mono(orig)) => {
assert_eq!(new, orig);
}
_ => panic!("Expected mono audio"),
}
}
#[test]
fn test_db_linear_conversion() {
assert!((db_to_linear(0.0_f64) - 1.0).abs() < 1e-10);
assert!((db_to_linear(20.0_f64) - 10.0).abs() < 1e-10);
assert!((db_to_linear(-20.0_f64) - 0.1).abs() < 1e-10);
assert!((linear_to_db(1.0_f64) - 0.0).abs() < 1e-10);
assert!((linear_to_db(10.0_f64) - 20.0).abs() < 1e-10);
assert!((linear_to_db(0.1_f64) - (-20.0)).abs() < 1e-10);
}
#[test]
fn test_peak_filter_gain_base_is_10() {
let frequency = 1000.0_f64;
let gain_db = 12.0_f64;
let q_factor = 1.0_f64;
let sample_rate = 48000.0_f64;
let expected_a = 10.0_f64.powf(gain_db / 40.0);
let buggy_a = 10.04_f64.powf(gain_db / 40.0);
assert!(
(expected_a - buggy_a).abs() > 1e-4,
"test premise: correct and buggy A must differ measurably"
);
let (b, a) = design_peak_filter(frequency, gain_db, q_factor, sample_rate);
let omega = 2.0 * PI * frequency / sample_rate;
let (sin_omega, cos_omega) = omega.sin_cos();
let alpha = sin_omega / (2.0 * q_factor);
let a0_unnorm = -2.0 * cos_omega / a[1];
let recovered_a = alpha / (a0_unnorm - 1.0);
assert!(
(recovered_a - expected_a).abs() < 1e-9,
"recovered A = {recovered_a}, expected {expected_a} (10.0 base)"
);
let expected_b0 = (1.0 + alpha * expected_a) / a0_unnorm;
assert!((b[0] - expected_b0).abs() < 1e-9);
}
#[test]
fn test_five_band_eq() {
let eq = ParametricEq::five_band();
assert_eq!(eq.band_count(), 5);
assert_eq!(eq.get_band(0).unwrap().frequency, 100.0);
assert_eq!(eq.get_band(1).unwrap().frequency, 300.0);
assert_eq!(eq.get_band(2).unwrap().frequency, 1000.0);
assert_eq!(eq.get_band(3).unwrap().frequency, 3000.0);
assert_eq!(eq.get_band(4).unwrap().frequency, 8000.0);
}
#[test]
fn test_apply_peak_filter_dual_variant() {
let sample_rate = 44100.0;
let samples_count = 1024;
let mut samples = Vec::with_capacity(samples_count);
for i in 0..samples_count {
let t = i as f64 / sample_rate;
samples.push((2.0 * PI * 880.0 * t).sin() as f32);
}
let original: AudioSamples<'_, f32> = AudioSamples::from_mono_vec(
NonEmptyVec::new(samples.clone()).unwrap(),
sample_rate!(44100),
);
let filtered = original.apply_peak_filter(880.0, 6.0, 2.0).unwrap();
let mut in_place = original.clone();
in_place
.apply_peak_filter_in_place(880.0, 6.0, 2.0)
.unwrap();
assert_eq!(
filtered.as_slice().unwrap(),
in_place.as_slice().unwrap(),
"non-mutating and in-place variants must produce equal results"
);
let pristine: AudioSamples<'_, f32> =
AudioSamples::from_mono_vec(NonEmptyVec::new(samples).unwrap(), sample_rate!(44100));
assert_eq!(
original.as_slice().unwrap(),
pristine.as_slice().unwrap(),
"non-mutating variant must not modify the original"
);
}
}