zk-audio 0.1.0

Audio processing library for voice recording and enhancement
Documentation
use crate::contracts::AudioProcessor;
use crate::core::{AudioError, AudioFrame, AudioProfile, AudioResult, AudioSpec};
use crate::processors::{analyze_frame, classify_frame, FrameClass};
use crate::profiles::profile_tuning;
use biquad::{Biquad, Coefficients, DirectForm1, Hertz, ToHertz, Type, Q_BUTTERWORTH_F32};

pub struct PresenceBoostProcessor {
    profile: AudioProfile,
    filter: Option<DirectForm1<f32>>,
}

impl PresenceBoostProcessor {
    pub fn new(profile: AudioProfile) -> Self {
        Self {
            profile,
            filter: None,
        }
    }
}

impl AudioProcessor for PresenceBoostProcessor {
    fn name(&self) -> &'static str {
        "presence_boost"
    }

    fn prepare(&mut self, spec: AudioSpec) -> AudioResult<()> {
        let gain_db = profile_tuning(self.profile).presence_gain_db;
        let coeffs = Coefficients::<f32>::from_params(
            Type::HighShelf(gain_db),
            Hertz::<f32>::from_hz(spec.sample_rate as f32)
                .map_err(|_| AudioError::new("Invalid sample rate for presence boost"))?,
            2_400.0f32.hz(),
            Q_BUTTERWORTH_F32,
        )
        .map_err(|_| AudioError::new("Invalid presence boost configuration"))?;
        self.filter = Some(DirectForm1::<f32>::new(coeffs));
        Ok(())
    }

    fn process(&mut self, frame: &mut AudioFrame) -> AudioResult<()> {
        let filter = self
            .filter
            .as_mut()
            .ok_or_else(|| AudioError::new("Presence boost not prepared"))?;
        for sample in &mut frame.samples {
            *sample = filter.run(*sample);
        }
        Ok(())
    }
}

pub struct LowShelfCutProcessor {
    profile: AudioProfile,
    filter: Option<DirectForm1<f32>>,
}

impl LowShelfCutProcessor {
    pub fn new(profile: AudioProfile) -> Self {
        Self {
            profile,
            filter: None,
        }
    }
}

impl AudioProcessor for LowShelfCutProcessor {
    fn name(&self) -> &'static str {
        "low_shelf_cut"
    }

    fn prepare(&mut self, spec: AudioSpec) -> AudioResult<()> {
        let tuning = profile_tuning(self.profile);
        let gain_db = tuning.low_shelf_gain_db;
        let cutoff_hz = tuning.low_shelf_cutoff_hz;
        let coeffs = Coefficients::<f32>::from_params(
            Type::LowShelf(gain_db),
            Hertz::<f32>::from_hz(spec.sample_rate as f32)
                .map_err(|_| AudioError::new("Invalid sample rate for low-shelf cut"))?,
            cutoff_hz.hz(),
            Q_BUTTERWORTH_F32,
        )
        .map_err(|_| AudioError::new("Invalid low-shelf cut configuration"))?;
        self.filter = Some(DirectForm1::<f32>::new(coeffs));
        Ok(())
    }

    fn process(&mut self, frame: &mut AudioFrame) -> AudioResult<()> {
        let filter = self
            .filter
            .as_mut()
            .ok_or_else(|| AudioError::new("Low-shelf cut not prepared"))?;
        for sample in &mut frame.samples {
            *sample = filter.run(*sample);
        }
        Ok(())
    }
}

pub struct AdaptiveToneShaperProcessor {
    profile: AudioProfile,
    detector_low_pass: Option<DirectForm1<f32>>,
    detector_high_pass: Option<DirectForm1<f32>>,
    apply_low_pass: Option<DirectForm1<f32>>,
    tilt_db: f32,
    noise_floor: f32,
    speech_memory: f32,
    last_rms: f32,
}

impl AdaptiveToneShaperProcessor {
    pub fn new(profile: AudioProfile) -> Self {
        Self {
            profile,
            detector_low_pass: None,
            detector_high_pass: None,
            apply_low_pass: None,
            tilt_db: 0.0,
            noise_floor: 0.01,
            speech_memory: 0.0,
            last_rms: 0.0,
        }
    }
}

impl AudioProcessor for AdaptiveToneShaperProcessor {
    fn name(&self) -> &'static str {
        "adaptive_tone_shaper"
    }

    fn prepare(&mut self, spec: AudioSpec) -> AudioResult<()> {
        let tuning = profile_tuning(self.profile);
        let sample_rate = Hertz::<f32>::from_hz(spec.sample_rate as f32)
            .map_err(|_| AudioError::new("Invalid sample rate for adaptive tone shaper"))?;
        let detector_low = Coefficients::<f32>::from_params(
            Type::LowPass,
            sample_rate,
            tuning.tone_shaper_low_detect_hz.hz(),
            Q_BUTTERWORTH_F32,
        )
        .map_err(|_| AudioError::new("Invalid low detector configuration for tone shaper"))?;
        let detector_high = Coefficients::<f32>::from_params(
            Type::HighPass,
            sample_rate,
            tuning.tone_shaper_high_detect_hz.hz(),
            Q_BUTTERWORTH_F32,
        )
        .map_err(|_| AudioError::new("Invalid high detector configuration for tone shaper"))?;
        let apply_low = Coefficients::<f32>::from_params(
            Type::LowPass,
            sample_rate,
            tuning.tone_shaper_pivot_hz.hz(),
            Q_BUTTERWORTH_F32,
        )
        .map_err(|_| AudioError::new("Invalid apply split configuration for tone shaper"))?;

        self.detector_low_pass = Some(DirectForm1::<f32>::new(detector_low));
        self.detector_high_pass = Some(DirectForm1::<f32>::new(detector_high));
        self.apply_low_pass = Some(DirectForm1::<f32>::new(apply_low));
        self.tilt_db = 0.0;
        self.noise_floor = 0.01;
        self.speech_memory = 0.0;
        self.last_rms = 0.0;
        Ok(())
    }

    fn process(&mut self, frame: &mut AudioFrame) -> AudioResult<()> {
        if frame.samples.is_empty() {
            return Ok(());
        }

        let detector_low = self
            .detector_low_pass
            .as_mut()
            .ok_or_else(|| AudioError::new("Tone shaper low detector not prepared"))?;
        let detector_high = self
            .detector_high_pass
            .as_mut()
            .ok_or_else(|| AudioError::new("Tone shaper high detector not prepared"))?;
        let apply_low = self
            .apply_low_pass
            .as_mut()
            .ok_or_else(|| AudioError::new("Tone shaper apply split not prepared"))?;

        let mut low_energy = 0.0;
        let mut high_energy = 0.0;
        for sample in &frame.samples {
            let low = detector_low.run(*sample);
            let high = detector_high.run(*sample);
            low_energy += low * low;
            high_energy += high * high;
        }
        low_energy = (low_energy / frame.samples.len() as f32).sqrt();
        high_energy = (high_energy / frame.samples.len() as f32).sqrt();

        let features = analyze_frame(&frame.samples, self.noise_floor);
        let frame_class = classify_frame(self.profile, features, self.noise_floor);
        let rms = features.rms;
        let peak = features.peak;
        if peak < self.noise_floor * 3.0 && rms < self.noise_floor * 2.0 {
            self.noise_floor = self.noise_floor * 0.93 + rms.max(0.0005) * 0.07;
        } else {
            self.noise_floor = self.noise_floor * 0.997 + rms.max(0.0005) * 0.003;
        }

        let tuning = profile_tuning(self.profile);
        let spectral_balance = ((low_energy + 1e-4) / (high_energy + 1e-4)).ln();
        let target_balance = match self.profile {
            AudioProfile::VoiceClean => 0.08,
            AudioProfile::VoiceNoisyRoom => 0.18,
            AudioProfile::VoiceHvac => 0.26,
            AudioProfile::Raw => 0.0,
        };
        let raw_tilt = ((target_balance - spectral_balance) * 3.2).clamp(
            -tuning.tone_shaper_max_tilt_db,
            tuning.tone_shaper_max_tilt_db,
        );
        let speech_weight = match frame_class {
            FrameClass::SpeechLike => 1.0,
            FrameClass::Transitional => 0.45,
            FrameClass::NoiseOnly => 0.0,
        };
        let sibilance_guard =
            ((high_energy - low_energy * 0.85) / (low_energy + 1e-4)).clamp(0.0, 1.0);
        let falling_energy = if self.last_rms > 1e-4 {
            ((self.last_rms - rms) / self.last_rms.max(rms)).clamp(0.0, 1.0)
        } else {
            0.0
        };
        let tail_brightness =
            ((high_energy - low_energy * 0.72) / (high_energy + 1e-4)).clamp(0.0, 1.0);
        let tail_guard =
            if !matches!(frame_class, FrameClass::SpeechLike) && self.speech_memory > 0.08 {
                (self.speech_memory * 0.55 + falling_energy * 0.85 + tail_brightness * 0.65)
                    .clamp(0.0, 1.0)
            } else {
                0.0
            };
        let target_tilt =
            (raw_tilt * speech_weight * (1.0 - 0.40 * sibilance_guard) * (1.0 - 0.65 * tail_guard))
                .clamp(
                    -tuning.tone_shaper_max_tilt_db,
                    tuning.tone_shaper_max_tilt_db,
                );
        let base_smoothing = if target_tilt.abs() > self.tilt_db.abs() {
            tuning.tone_shaper_attack
        } else {
            tuning.tone_shaper_release
        };
        let smoothing = if tail_guard > 0.0 {
            base_smoothing.max(0.18 + tail_guard * 0.28)
        } else {
            base_smoothing
        };
        let start_tilt = self.tilt_db;
        self.tilt_db = start_tilt * (1.0 - smoothing) + target_tilt * smoothing;
        if tail_guard > 0.0 {
            self.tilt_db *= (1.0 - 0.55 * tail_guard).clamp(0.0, 1.0);
        }

        let len = frame.samples.len().max(1) as f32;
        let contour_high_gain = 1.0 - 0.26 * tail_guard;
        for (index, sample) in frame.samples.iter_mut().enumerate() {
            let t = index as f32 / len;
            let tilt_db = start_tilt + (self.tilt_db - start_tilt) * t;
            let low_gain = 10.0f32.powf((tilt_db * 0.5) / 20.0);
            let high_gain = 10.0f32.powf((-tilt_db * 0.5) / 20.0);
            let low = apply_low.run(*sample);
            let high = *sample - low;
            *sample = low * low_gain + high * high_gain * contour_high_gain;
        }

        self.speech_memory = match frame_class {
            FrameClass::SpeechLike => (self.speech_memory * 0.82 + 0.24).clamp(0.0, 1.0),
            FrameClass::Transitional => self.speech_memory * 0.84,
            FrameClass::NoiseOnly => self.speech_memory * 0.68,
        };
        self.last_rms = rms;

        Ok(())
    }
}

#[cfg(test)]
mod tests {
    use super::*;

    fn test_spec() -> AudioSpec {
        AudioSpec {
            sample_rate: 44_100,
            channels: 1,
        }
    }

    #[test]
    fn low_shelf_cut_reduces_low_frequency_energy() {
        let spec = test_spec();
        let mut processor = LowShelfCutProcessor::new(AudioProfile::VoiceNoisyRoom);
        processor.prepare(spec).unwrap();

        let mut frame = AudioFrame {
            samples: vec![0.2; 480],
            spec,
        };
        let before =
            frame.samples.iter().map(|s| s.abs()).sum::<f32>() / frame.samples.len() as f32;
        processor.process(&mut frame).unwrap();
        let after = frame.samples.iter().map(|s| s.abs()).sum::<f32>() / frame.samples.len() as f32;
        assert!(after < before);
    }

    #[test]
    fn adaptive_tone_shaper_warms_bright_signal() {
        let spec = test_spec();
        let mut processor = AdaptiveToneShaperProcessor::new(AudioProfile::VoiceHvac);
        processor.prepare(spec).unwrap();

        let mut frame = AudioFrame {
            samples: (0..512)
                .map(|i| {
                    let t = i as f32 / spec.sample_rate as f32;
                    (std::f32::consts::TAU * 2_800.0 * t).sin() * 0.10
                        + (std::f32::consts::TAU * 180.0 * t).sin() * 0.03
                })
                .collect(),
            spec,
        };

        for _ in 0..6 {
            processor.process(&mut frame).unwrap();
        }

        assert!(processor.tilt_db > 0.1);
    }

    #[test]
    fn adaptive_tone_shaper_brightens_dull_signal() {
        let spec = test_spec();
        let mut processor = AdaptiveToneShaperProcessor::new(AudioProfile::VoiceClean);
        processor.prepare(spec).unwrap();

        let mut frame = AudioFrame {
            samples: (0..512)
                .map(|i| {
                    let t = i as f32 / spec.sample_rate as f32;
                    (std::f32::consts::TAU * 180.0 * t).sin() * 0.12
                })
                .collect(),
            spec,
        };

        for _ in 0..6 {
            processor.process(&mut frame).unwrap();
        }

        assert!(processor.tilt_db < -0.05);
    }

    #[test]
    fn adaptive_tone_shaper_stays_idle_on_noise_only_frames() {
        let spec = test_spec();
        let mut processor = AdaptiveToneShaperProcessor::new(AudioProfile::VoiceHvac);
        processor.prepare(spec).unwrap();

        let mut frame = AudioFrame {
            samples: (0..512)
                .map(|i| if i % 2 == 0 { 0.01 } else { -0.01 })
                .collect(),
            spec,
        };
        processor.process(&mut frame).unwrap();

        assert!(processor.tilt_db.abs() < 0.02);
    }
}