#![forbid(unsafe_code)]
use crate::error::{AudioError, AudioResult};
#[inline]
fn db_to_linear(db: f64) -> f64 {
10.0_f64.powf(db / 20.0)
}
#[inline]
fn linear_to_db(linear: f64) -> f64 {
if linear < 1e-6 {
-120.0
} else {
20.0 * linear.log10()
}
}
#[inline]
fn ema_coeff(time_ms: f64, sample_rate: u32) -> f64 {
if time_ms <= 0.0 {
return 1.0; }
let tau_samples = time_ms * 1e-3 * f64::from(sample_rate);
(-1.0 / tau_samples).exp()
}
#[derive(Clone, Debug)]
pub struct AutoGainConfig {
pub target_dbfs: f64,
pub attack_ms: f64,
pub release_ms: f64,
pub gain_smoothing_ms: f64,
pub min_gain_db: f64,
pub max_gain_db: f64,
pub hold_ms: f64,
}
impl Default for AutoGainConfig {
fn default() -> Self {
Self {
target_dbfs: -18.0,
attack_ms: 50.0,
release_ms: 200.0,
gain_smoothing_ms: 100.0,
min_gain_db: -30.0,
max_gain_db: 30.0,
hold_ms: 20.0,
}
}
}
pub struct AutoGainController {
config: AutoGainConfig,
sample_rate: u32,
power_ema: f64,
smooth_gain: f64,
min_gain_linear: f64,
max_gain_linear: f64,
target_linear: f64,
attack_coeff: f64,
release_coeff: f64,
gain_coeff: f64,
hold_samples_remaining: u64,
hold_samples_total: u64,
prev_desired_gain: f64,
}
impl AutoGainController {
pub fn new(config: AutoGainConfig, sample_rate: u32) -> AudioResult<Self> {
if sample_rate == 0 {
return Err(AudioError::InvalidParameter(
"sample_rate must be non-zero".into(),
));
}
if config.min_gain_db > config.max_gain_db {
return Err(AudioError::InvalidParameter(
"min_gain_db must be <= max_gain_db".into(),
));
}
let attack_coeff = ema_coeff(config.attack_ms, sample_rate);
let release_coeff = ema_coeff(config.release_ms, sample_rate);
let gain_coeff = ema_coeff(config.gain_smoothing_ms, sample_rate);
let target_linear = db_to_linear(config.target_dbfs);
let min_gain_linear = db_to_linear(config.min_gain_db);
let max_gain_linear = db_to_linear(config.max_gain_db);
let hold_samples_total = (config.hold_ms * 1e-3 * f64::from(sample_rate)).round() as u64;
Ok(Self {
config,
sample_rate,
power_ema: 0.0,
smooth_gain: 1.0,
min_gain_linear,
max_gain_linear,
target_linear,
attack_coeff,
release_coeff,
gain_coeff,
hold_samples_remaining: 0,
hold_samples_total,
prev_desired_gain: 1.0,
})
}
pub fn process_inplace(&mut self, samples: &mut [f32]) {
for s in samples.iter_mut() {
let x = f64::from(*s);
let power = x * x;
let coeff = if power > self.power_ema {
self.attack_coeff
} else {
self.release_coeff
};
self.power_ema = coeff * self.power_ema + (1.0 - coeff) * power;
let rms = self.power_ema.sqrt();
let desired_gain = if rms < 1e-7 {
self.max_gain_linear
} else {
(self.target_linear / rms).clamp(self.min_gain_linear, self.max_gain_linear)
};
let effective_desired = if desired_gain > self.prev_desired_gain {
if self.hold_samples_remaining > 0 {
self.hold_samples_remaining -= 1;
self.prev_desired_gain } else {
self.hold_samples_remaining = self.hold_samples_total;
desired_gain
}
} else {
self.hold_samples_remaining = 0;
desired_gain
};
self.prev_desired_gain = effective_desired;
self.smooth_gain =
self.gain_coeff * self.smooth_gain + (1.0 - self.gain_coeff) * effective_desired;
*s = (x * self.smooth_gain) as f32;
}
}
#[must_use]
pub fn process(&mut self, samples: &[f32]) -> Vec<f32> {
let mut out = samples.to_vec();
self.process_inplace(&mut out);
out
}
#[must_use]
pub fn current_gain_db(&self) -> f64 {
linear_to_db(self.smooth_gain)
}
#[must_use]
pub fn current_rms_dbfs(&self) -> f64 {
linear_to_db(self.power_ema.sqrt())
}
pub fn reset(&mut self) {
self.power_ema = 0.0;
self.smooth_gain = 1.0;
self.prev_desired_gain = 1.0;
self.hold_samples_remaining = 0;
}
pub fn set_target_dbfs(&mut self, target_dbfs: f64) {
self.config.target_dbfs = target_dbfs;
self.target_linear = db_to_linear(target_dbfs);
}
#[must_use]
pub fn config(&self) -> &AutoGainConfig {
&self.config
}
#[must_use]
pub fn sample_rate(&self) -> u32 {
self.sample_rate
}
}
#[derive(Clone, Debug)]
pub struct BatchNormResult {
pub gain_linear: f64,
pub gain_db: f64,
pub measured_peak_linear: f64,
pub measured_rms_dbfs: f64,
}
pub struct BatchAutoGain {
pub target_rms_dbfs: f64,
pub peak_ceiling_dbfs: f64,
}
impl Default for BatchAutoGain {
fn default() -> Self {
Self {
target_rms_dbfs: -18.0,
peak_ceiling_dbfs: -0.5,
}
}
}
impl BatchAutoGain {
#[must_use]
pub fn new(target_rms_dbfs: f64, peak_ceiling_dbfs: f64) -> Self {
Self {
target_rms_dbfs,
peak_ceiling_dbfs,
}
}
pub fn normalize_inplace(&self, buffer: &mut [f32]) -> AudioResult<BatchNormResult> {
if buffer.is_empty() {
return Err(AudioError::InvalidData("buffer is empty".into()));
}
let sum_sq: f64 = buffer.iter().map(|&s| f64::from(s) * f64::from(s)).sum();
let rms = (sum_sq / buffer.len() as f64).sqrt();
let peak = buffer.iter().map(|s| s.abs()).fold(0.0_f32, f32::max);
let peak_linear = f64::from(peak);
let measured_rms_dbfs = linear_to_db(rms);
let target_rms_linear = db_to_linear(self.target_rms_dbfs);
let gain_for_rms = if rms < 1e-9 {
1.0
} else {
target_rms_linear / rms
};
let peak_ceiling_linear = db_to_linear(self.peak_ceiling_dbfs);
let max_gain_for_peak = if peak_linear < 1e-9 {
gain_for_rms
} else {
peak_ceiling_linear / peak_linear
};
let gain_linear = gain_for_rms.min(max_gain_for_peak).max(0.0);
let gain_db = linear_to_db(gain_linear);
for s in buffer.iter_mut() {
*s = (*s as f64 * gain_linear) as f32;
}
Ok(BatchNormResult {
gain_linear,
gain_db,
measured_peak_linear: peak_linear,
measured_rms_dbfs,
})
}
pub fn normalize(&self, samples: &[f32]) -> AudioResult<(Vec<f32>, BatchNormResult)> {
let mut buf = samples.to_vec();
let result = self.normalize_inplace(&mut buf)?;
Ok((buf, result))
}
}
#[cfg(test)]
mod tests {
use super::*;
fn make_agc() -> AutoGainController {
AutoGainController::new(AutoGainConfig::default(), 48_000).expect("valid config")
}
#[test]
fn test_agc_construction() {
let _ = make_agc();
}
#[test]
fn test_agc_invalid_sample_rate() {
let result = AutoGainController::new(AutoGainConfig::default(), 0);
assert!(result.is_err());
}
#[test]
fn test_agc_invalid_gain_range() {
let config = AutoGainConfig {
min_gain_db: 10.0,
max_gain_db: 5.0, ..Default::default()
};
let result = AutoGainController::new(config, 48_000);
assert!(result.is_err());
}
#[test]
fn test_agc_silence_does_not_panic() {
let mut agc = make_agc();
let mut buf = vec![0.0_f32; 480];
agc.process_inplace(&mut buf);
for s in &buf {
assert!(*s == 0.0 || s.is_finite());
}
}
#[test]
fn test_agc_gain_reduces_loud_signal() {
let mut agc = AutoGainController::new(
AutoGainConfig {
target_dbfs: -18.0,
attack_ms: 1.0,
release_ms: 10.0,
gain_smoothing_ms: 5.0,
hold_ms: 0.0,
..Default::default()
},
48_000,
)
.expect("valid");
let loud: Vec<f32> = vec![0.95_f32; 4800];
agc.process(&loud);
assert!(
agc.smooth_gain < 1.0,
"smooth_gain={} expected < 1.0",
agc.smooth_gain
);
}
#[test]
fn test_agc_reset_clears_state() {
let mut agc = make_agc();
let loud: Vec<f32> = vec![0.9_f32; 1000];
let _ = agc.process(&loud);
agc.reset();
assert!((agc.smooth_gain - 1.0).abs() < 1e-9);
assert!(agc.power_ema < 1e-15);
}
#[test]
fn test_agc_set_target_changes_target() {
let mut agc = make_agc();
agc.set_target_dbfs(-12.0);
assert!((agc.config().target_dbfs - (-12.0)).abs() < 1e-9);
assert!((agc.target_linear - db_to_linear(-12.0)).abs() < 1e-9);
}
#[test]
fn test_agc_current_gain_db_type() {
let agc = make_agc();
let db = agc.current_gain_db();
assert!((db - 0.0).abs() < 1e-6, "db={}", db);
}
#[test]
fn test_batch_normalise_empty_error() {
let bgc = BatchAutoGain::default();
let result = bgc.normalize_inplace(&mut []);
assert!(result.is_err());
}
#[test]
fn test_batch_normalise_rms_target() {
let bgc = BatchAutoGain::new(-18.0, -0.1);
let mut buf = vec![0.5_f32; 4800];
let result = bgc.normalize_inplace(&mut buf).expect("ok");
let rms_after: f64 = {
let sum_sq: f64 = buf.iter().map(|&s| f64::from(s) * f64::from(s)).sum();
(sum_sq / buf.len() as f64).sqrt()
};
let rms_db_after = linear_to_db(rms_after);
assert!(
(rms_db_after - (-18.0)).abs() < 0.5,
"rms_after={:.2} dBFS",
rms_db_after
);
assert!(result.gain_linear > 0.0);
}
#[test]
fn test_batch_normalise_peak_ceiling() {
let bgc = BatchAutoGain::new(-3.0, -1.0); let mut buf = vec![0.9_f32; 4800]; let result = bgc.normalize_inplace(&mut buf).expect("ok");
let peak_after = buf.iter().map(|s| s.abs()).fold(0.0_f32, f32::max);
let peak_ceiling = db_to_linear(-1.0);
assert!(
f64::from(peak_after) <= peak_ceiling + 1e-6,
"peak_after={:.4} limit={:.4}",
peak_after,
peak_ceiling
);
assert!(result.gain_db < 0.0 || result.gain_db.abs() < 1.0);
}
#[test]
fn test_db_linear_roundtrip() {
for db in [-60.0_f64, -18.0, 0.0, 6.0, 20.0] {
let linear = db_to_linear(db);
let back = linear_to_db(linear);
assert!((back - db).abs() < 1e-9, "db={} back={}", db, back);
}
}
}