use crate::core::io::AudioData;
use thiserror::Error;
#[derive(Error, Debug)]
pub enum AmplitudeError {
#[error("Invalid gain factor: {0}")]
InvalidGain(String),
#[error("Invalid signal: {0}")]
InvalidSignal(String),
}
pub fn amplify(signal: &AudioData, gain: f32) -> Result<AudioData, AmplitudeError> {
if gain <= 0.0 {
return Err(AmplitudeError::InvalidGain(
"Gain must be positive".to_string(),
));
}
let samples = signal.samples.iter().map(|&s| s * gain).collect();
Ok(AudioData {
samples,
sample_rate: signal.sample_rate,
channels: signal.channels,
})
}
pub fn attenuate(signal: &AudioData, gain: f32) -> Result<AudioData, AmplitudeError> {
if gain < 0.0 {
return Err(AmplitudeError::InvalidGain(
"Gain must be non-negative".to_string(),
));
}
let samples = signal.samples.iter().map(|&s| s * gain).collect();
Ok(AudioData {
samples,
sample_rate: signal.sample_rate,
channels: signal.channels,
})
}
pub fn normalize(signal: &AudioData, target: f32, mode: &str) -> Result<AudioData, AmplitudeError> {
if target <= 0.0 {
return Err(AmplitudeError::InvalidGain(
"Target level must be positive".to_string(),
));
}
if signal.samples.is_empty() {
return Err(AmplitudeError::InvalidSignal(
"Signal cannot be empty".to_string(),
));
}
let gain = match mode.to_lowercase().as_str() {
"peak" => {
let max_amplitude = signal
.samples
.iter()
.map(|s| s.abs())
.max_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal))
.unwrap_or(0.0);
if max_amplitude == 0.0 {
return Err(AmplitudeError::InvalidSignal(
"Signal has no amplitude to normalize".to_string(),
));
}
target / max_amplitude
}
"rms" => {
let rms = (signal.samples.iter().map(|&s| s * s).sum::<f32>() / signal.samples.len() as f32)
.sqrt();
if rms == 0.0 {
return Err(AmplitudeError::InvalidSignal(
"Signal has no RMS level to normalize".to_string(),
));
}
target / rms
}
_ => {
return Err(AmplitudeError::InvalidGain(format!(
"Unknown normalization mode: {}",
mode
)))
}
};
let samples = signal.samples.iter().map(|&s| s * gain).collect();
Ok(AudioData {
samples,
sample_rate: signal.sample_rate,
channels: signal.channels,
})
}