use std::fmt;
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::sync::{Arc, Mutex, Weak};
use fixed_resample::{ResamplingChannelConfig, ResamplingCons, ResamplingProd, resampling_channel};
use sonora::config::{EchoCanceller, GainController2, NoiseSuppression, NoiseSuppressionLevel};
use sonora::{AudioProcessing, StreamConfig};
use crate::Error;
use crate::playback::{self, BUS_CHANNELS};
const REFERENCE_RATE: u32 = 48_000;
const REFERENCE_LATENCY: f64 = 0.01;
const REFERENCE_CAPACITY: f64 = 0.2;
const REFERENCE_FRAME: usize = REFERENCE_RATE as usize / 100;
const MAX_CALLBACK: usize = 4096;
#[derive(Clone, Debug)]
#[non_exhaustive]
pub struct Config {
pub noise_suppression: bool,
pub auto_gain: bool,
}
impl Default for Config {
fn default() -> Self {
Self {
noise_suppression: true,
auto_gain: true,
}
}
}
impl Config {
fn build(&self) -> sonora::Config {
sonora::Config {
echo_canceller: Some(EchoCanceller::default()),
noise_suppression: self.noise_suppression.then(|| NoiseSuppression {
level: NoiseSuppressionLevel::High,
..Default::default()
}),
gain_controller2: self.auto_gain.then(GainController2::default),
..Default::default()
}
}
}
#[derive(Clone)]
pub struct Canceller {
inner: Arc<Inner>,
}
impl Canceller {
pub(crate) fn new(shared: Arc<playback::Shared>, config: Config) -> Self {
static NEXT_ID: AtomicU64 = AtomicU64::new(0);
let inner = Arc::new(Inner {
id: NEXT_ID.fetch_add(1, Ordering::Relaxed),
enabled: AtomicBool::new(true),
config,
state: Arc::new(Mutex::new(State::default())),
shared: shared.clone(),
});
shared.set_reference(Reference {
id: inner.id,
alive: Arc::downgrade(&inner),
state: inner.state.clone(),
pending: None,
});
Self { inner }
}
pub fn set_enabled(&self, enabled: bool) {
self.inner.enabled.store(enabled, Ordering::Relaxed);
}
pub fn enabled(&self) -> bool {
self.inner.enabled.load(Ordering::Relaxed)
}
pub(crate) fn open(&self, sample_rate: u32, channels: u32) -> Result<(), Error> {
if !(8_000..=384_000).contains(&sample_rate) {
return Err(Error::Unsupported(format!(
"echo cancellation needs a microphone between 8 and 384 kHz (got {sample_rate})"
)));
}
if channels == 0 || channels > BUS_CHANNELS as u32 {
return Err(Error::Unsupported(format!(
"echo cancellation accepts a mono or stereo microphone (got {channels} channels)"
)));
}
let capture = StreamConfig::new(sample_rate, channels as u16);
let processor = AudioProcessing::builder()
.capture_config(capture)
.render_config(reference_config())
.config(self.inner.config.build())
.build();
let mut state = self.inner.state.lock().unwrap();
if state.processor.is_none() || state.capture != capture {
state.processor = Some(processor);
state.capture = capture;
state.resize();
} else {
state.reset();
}
if let Some(reference) = &mut state.reference {
reference.discard_frames(reference.available_frames());
}
Ok(())
}
pub(crate) fn process(&self, buf: &mut [f32]) {
let enabled = self.enabled();
self.inner.state.lock().unwrap().process(buf, enabled);
}
}
impl fmt::Debug for Canceller {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Canceller")
.field("enabled", &self.enabled())
.finish_non_exhaustive()
}
}
struct Inner {
id: u64,
enabled: AtomicBool,
config: Config,
state: Arc<Mutex<State>>,
shared: Arc<playback::Shared>,
}
impl Drop for Inner {
fn drop(&mut self) {
self.shared.clear_reference(self.id);
}
}
struct State {
processor: Option<AudioProcessing>,
capture: StreamConfig,
reference: Option<ResamplingCons<f32>>,
running: bool,
pending: Vec<f32>,
processed: Vec<f32>,
reference_frame: Vec<f32>,
render_in: Vec<f32>,
render_out: Vec<f32>,
capture_in: Vec<f32>,
capture_out: Vec<f32>,
}
impl Default for State {
fn default() -> Self {
Self {
processor: None,
capture: reference_config(),
reference: None,
running: false,
pending: Vec::new(),
processed: Vec::new(),
reference_frame: Vec::new(),
render_in: Vec::new(),
render_out: Vec::new(),
capture_in: Vec::new(),
capture_out: Vec::new(),
}
}
}
impl State {
fn resize(&mut self) {
let frame = self.capture.num_frames();
let channels = self.capture.num_channels() as usize;
let headroom = (frame + MAX_CALLBACK) * channels;
self.pending = Vec::with_capacity(headroom);
self.processed = Vec::with_capacity(headroom);
self.reference_frame = vec![0.0; REFERENCE_FRAME * BUS_CHANNELS];
self.render_in = vec![0.0; REFERENCE_FRAME * BUS_CHANNELS];
self.render_out = vec![0.0; REFERENCE_FRAME * BUS_CHANNELS];
self.capture_in = vec![0.0; frame * channels];
self.capture_out = vec![0.0; frame * channels];
self.reset();
}
fn reset(&mut self) {
self.pending.clear();
self.processed.clear();
self.running = false;
}
fn process(&mut self, buf: &mut [f32], enabled: bool) {
let Self {
processor,
capture,
reference,
running,
pending,
processed,
reference_frame,
render_in,
render_out,
capture_in,
capture_out,
} = self;
let Some(processor) = processor else { return };
let channels = capture.num_channels() as usize;
let stride = capture.num_frames() * channels;
if !buf.len().is_multiple_of(channels) {
return;
}
if !enabled {
if *running {
pending.clear();
processed.clear();
*running = false;
}
if let Some(reference) = reference {
reference.discard_frames(reference.available_frames());
}
return;
}
*running = true;
if let Some(reference) = reference {
while reference.available_frames() >= REFERENCE_FRAME {
reference.read_interleaved(reference_frame, false);
deinterleave(reference_frame, render_in, BUS_CHANNELS);
let _ = process_render(processor, render_in, render_out);
}
}
pending.extend_from_slice(buf);
while pending.len() >= stride {
deinterleave(&pending[..stride], capture_in, channels);
match process_capture(processor, capture, capture_in, capture_out) {
Ok(()) => interleave(capture_out, channels, processed),
Err(_) => processed.extend_from_slice(&pending[..stride]),
}
pending.drain(..stride);
}
let ready = processed.len().min(buf.len());
buf[..ready].copy_from_slice(&processed[..ready]);
buf[ready..].fill(0.0);
processed.drain(..ready);
}
}
fn reference_config() -> StreamConfig {
StreamConfig::new(REFERENCE_RATE, BUS_CHANNELS as u16)
}
fn process_render(processor: &mut AudioProcessing, input: &[f32], output: &mut [f32]) -> Result<(), sonora::Error> {
let config = reference_config();
let (left, right) = input.split_at(REFERENCE_FRAME);
let (left_out, right_out) = output.split_at_mut(REFERENCE_FRAME);
processor.process_render_f32_with_config(&[left, right], &config, &config, &mut [left_out, right_out])
}
fn process_capture(
processor: &mut AudioProcessing,
config: &StreamConfig,
input: &[f32],
output: &mut [f32],
) -> Result<(), sonora::Error> {
let frame = config.num_frames();
match config.num_channels() {
1 => processor.process_capture_f32_with_config(&[input], config, config, &mut [output]),
_ => {
let (left, right) = input.split_at(frame);
let (left_out, right_out) = output.split_at_mut(frame);
processor.process_capture_f32_with_config(&[left, right], config, config, &mut [left_out, right_out])
}
}
}
fn deinterleave(src: &[f32], dest: &mut [f32], channels: usize) {
let frames = src.len() / channels;
for (channel, samples) in dest.chunks_exact_mut(frames).enumerate() {
for (frame, sample) in samples.iter_mut().enumerate() {
*sample = src[frame * channels + channel];
}
}
}
fn interleave(src: &[f32], channels: usize, dest: &mut Vec<f32>) {
let frames = src.len() / channels;
for frame in 0..frames {
for channel in 0..channels {
dest.push(src[channel * frames + frame]);
}
}
}
pub(crate) struct Reference {
id: u64,
alive: Weak<Inner>,
state: Arc<Mutex<State>>,
pending: Option<ResamplingProd<f32>>,
}
impl Reference {
pub(crate) fn owned_by(&self, id: u64) -> bool {
self.id == id
}
pub(crate) fn attached(&self) -> bool {
self.pending.is_none()
}
pub(crate) fn take(&mut self) -> Option<ResamplingProd<f32>> {
self.pending.take()
}
pub(crate) fn restore(&mut self, prod: ResamplingProd<f32>) {
self.pending = Some(prod);
}
pub(crate) fn rebuild(&mut self, rate: u32) -> bool {
if self.alive.strong_count() == 0 {
return false;
}
let (prod, cons) = channel(rate);
self.state.lock().unwrap().reference = Some(cons);
self.pending = Some(prod);
true
}
}
fn channel(rate: u32) -> (ResamplingProd<f32>, ResamplingCons<f32>) {
resampling_channel::<f32>(
BUS_CHANNELS,
rate,
REFERENCE_RATE,
true,
ResamplingChannelConfig {
latency_seconds: REFERENCE_LATENCY,
capacity_seconds: REFERENCE_CAPACITY,
underflow_autocorrect_percent_threshold: None,
..Default::default()
},
)
}
#[cfg(test)]
mod tests {
use std::collections::VecDeque;
use super::*;
const fn frame(rate: u32) -> usize {
rate as usize / 100
}
fn detached() -> Canceller {
Canceller::new(Arc::new(playback::Shared::default()), Config::default())
}
fn opened(sample_rate: u32, channels: u32) -> Canceller {
let canceller = detached();
canceller.open(sample_rate, channels).unwrap();
canceller
}
fn tap(canceller: &Canceller) -> ResamplingProd<f32> {
let (prod, cons) = channel(REFERENCE_RATE);
canceller.inner.state.lock().unwrap().reference = Some(cons);
prod
}
fn play(prod: &mut ResamplingProd<f32>, value: f32) {
prod.push_interleaved(&vec![value; REFERENCE_FRAME * BUS_CHANNELS]);
}
#[test]
fn rejects_formats_the_processor_cannot_take() {
let canceller = detached();
assert!(matches!(canceller.open(4_000, 1), Err(Error::Unsupported(_))));
assert!(matches!(canceller.open(48_000, 0), Err(Error::Unsupported(_))));
assert!(matches!(canceller.open(48_000, 6), Err(Error::Unsupported(_))));
canceller.open(48_000, 1).unwrap();
canceller.open(16_000, 2).unwrap();
}
#[test]
fn returns_exactly_what_it_was_given() {
let canceller = opened(48_000, 2);
for frames in [64, 480, 512, 4096] {
let mut buf = vec![0.25f32; frames * 2];
canceller.process(&mut buf);
assert_eq!(buf.len(), frames * 2, "{frames} frames");
}
}
#[test]
fn whole_frames_are_not_delayed() {
let canceller = opened(48_000, 1);
let frame = frame(48_000);
let mut loud = false;
for _ in 0..20 {
let mut buf = vec![0.5f32; frame];
canceller.process(&mut buf);
loud |= buf.iter().any(|s| s.abs() > 0.01);
}
assert!(loud, "the microphone came back silent");
}
#[test]
fn partial_frames_cost_the_leftover_once() {
let canceller = opened(48_000, 1);
let frame = frame(48_000);
let mut first = vec![0.5f32; frame * 3 / 2];
canceller.process(&mut first);
let silent = first.iter().rev().take_while(|s| **s == 0.0).count();
assert_eq!(silent, frame / 2, "expected the leftover half-frame to be held back");
let mut second = vec![0.5f32; frame * 3 / 2];
canceller.process(&mut second);
assert_ne!(second.last(), Some(&0.0), "the leftover was never made up");
}
#[test]
fn keeps_up_with_callbacks_that_straddle_frames() {
let canceller = opened(48_000, 1);
let chunk = frame(48_000) * 3 / 2;
let mut short = 0;
for round in 0..20 {
let mut buf = vec![0.5f32; chunk];
canceller.process(&mut buf);
if round > 2 {
short += buf.iter().rev().take_while(|s| **s == 0.0).count();
}
}
assert_eq!(short, 0, "samples went missing after the pipeline filled");
}
#[test]
fn disabled_passes_the_microphone_straight_through() {
let canceller = opened(48_000, 2);
canceller.set_enabled(false);
assert!(!canceller.enabled());
let mut buf = vec![0.5f32; 960 * 2];
canceller.process(&mut buf);
assert!(buf.iter().all(|s| *s == 0.5), "passthrough altered the samples");
}
#[test]
fn toggling_off_drops_buffered_samples() {
let canceller = opened(48_000, 1);
let frame = frame(48_000);
let mut buf = vec![0.5f32; frame * 3 / 2];
canceller.process(&mut buf);
assert!(!canceller.inner.state.lock().unwrap().pending.is_empty());
canceller.set_enabled(false);
let mut buf = vec![0.5f32; frame / 2];
canceller.process(&mut buf);
assert!(buf.iter().all(|s| *s == 0.5), "passthrough altered the samples");
let state = canceller.inner.state.lock().unwrap();
assert!(state.pending.is_empty(), "a partial frame survived the toggle");
assert!(state.processed.is_empty(), "processed samples survived the toggle");
}
#[test]
fn a_replaced_canceller_leaves_the_tap_alone() {
let shared = Arc::new(playback::Shared::default());
let first = Canceller::new(shared.clone(), Config::default());
let second = Canceller::new(shared.clone(), Config::default());
assert!(shared.has_reference());
drop(first);
assert!(shared.has_reference(), "the replacement lost its tap");
drop(second);
assert!(!shared.has_reference(), "the tap outlived every canceller");
}
#[test]
fn a_canceller_without_a_microphone_leaves_the_buffer_alone() {
let canceller = detached();
let mut buf = vec![0.5f32; 480];
canceller.process(&mut buf);
assert!(buf.iter().all(|s| *s == 0.5));
}
#[test]
fn a_pass_drains_the_whole_tap() {
let canceller = opened(48_000, 1);
let mut prod = tap(&canceller);
let frame = frame(48_000);
canceller.process(&mut vec![0.0f32; frame]);
play(&mut prod, 0.5);
canceller.process(&mut vec![0.0f32; frame]);
for _ in 0..10 {
play(&mut prod, 0.5);
}
canceller.process(&mut vec![0.0f32; frame]);
let queued = canceller
.inner
.state
.lock()
.unwrap()
.reference
.as_ref()
.map_or(0, |r| r.available_frames());
assert!(
queued < REFERENCE_FRAME,
"{queued} reference frames were left queued as standing delay"
);
}
#[test]
fn survives_a_missing_reference() {
let canceller = opened(48_000, 2);
assert!(canceller.inner.state.lock().unwrap().reference.is_none());
let mut buf = vec![0.5f32; frame(48_000) * 2 * 2];
canceller.process(&mut buf);
canceller.process(&mut buf);
assert!(buf.iter().any(|s| s.abs() > 0.01), "audio was dropped");
}
struct Room {
canceller: Canceller,
prod: ResamplingProd<f32>,
noise: Noise,
history: VecDeque<Vec<f32>>,
delay: usize,
frame: usize,
}
const MAX_DELAY: usize = 12;
const ATTENUATION: f32 = 0.5;
impl Room {
fn new(delay: usize) -> Self {
let canceller = detached();
let capture = StreamConfig::new(48_000, 1);
{
let mut state = canceller.inner.state.lock().unwrap();
state.processor = Some(
AudioProcessing::builder()
.capture_config(capture)
.render_config(reference_config())
.config(sonora::Config {
echo_canceller: Some(EchoCanceller::default()),
..Default::default()
})
.build(),
);
state.capture = capture;
state.resize();
}
let prod = tap(&canceller);
Self {
canceller,
prod,
noise: Noise::default(),
history: VecDeque::with_capacity(MAX_DELAY + 1),
delay,
frame: frame(48_000),
}
}
fn round(&mut self) -> (f64, f64) {
let played: Vec<f32> = (0..self.frame).map(|_| self.noise.next()).collect();
let mut reference = vec![0.0f32; self.frame * BUS_CHANNELS];
for (i, sample) in played.iter().enumerate() {
reference[i * BUS_CHANNELS] = *sample;
reference[i * BUS_CHANNELS + 1] = *sample;
}
self.prod.push_interleaved(&reference);
self.history.push_back(played);
if self.history.len() > MAX_DELAY + 1 {
self.history.pop_front();
}
let echo = self.history.len().saturating_sub(1 + self.delay);
let mut heard: Vec<f32> = self.history[echo].iter().map(|s| s * ATTENUATION).collect();
let before = energy(&heard);
self.canceller.process(&mut heard);
(before, energy(&heard))
}
}
fn energy(samples: &[f32]) -> f64 {
samples.iter().map(|s| (*s as f64).powi(2)).sum()
}
#[test]
fn cancels_the_echo_of_what_was_played() {
let mut room = Room::new(2);
let (mut heard, mut left, mut measured) = (0.0f64, 0.0f64, 0);
for round in 0..200 {
let (before, after) = room.round();
if round >= 150 {
heard += before;
left += after;
measured += 1;
}
}
assert!(measured > 0);
let attenuation = 10.0 * (heard / left.max(f64::MIN_POSITIVE)).log10();
assert!(attenuation > 20.0, "echo was only attenuated by {attenuation:.1} dB");
}
#[test]
fn survives_a_moving_echo_delay() {
let mut room = Room::new(MAX_DELAY - 4);
for _ in 0..400 {
room.round();
}
for delay in [1, MAX_DELAY - 4, 2, MAX_DELAY - 2, 1] {
room.delay = delay;
for _ in 0..200 {
room.round();
}
}
let (mut heard, mut left) = (0.0f64, 0.0f64);
for _ in 0..100 {
let (before, after) = room.round();
heard += before;
left += after;
}
let attenuation = 10.0 * (heard / left.max(f64::MIN_POSITIVE)).log10();
assert!(attenuation > 10.0, "the filter never re-converged: {attenuation:.1} dB");
}
#[derive(Default)]
struct Noise(u32);
impl Noise {
fn next(&mut self) -> f32 {
if self.0 == 0 {
self.0 = 0x1234_5678;
}
self.0 ^= self.0 << 13;
self.0 ^= self.0 >> 17;
self.0 ^= self.0 << 5;
self.0 as f32 / u32::MAX as f32 - 0.5
}
}
#[tokio::test]
#[ignore]
async fn taps_a_real_output_device() {
let engine = playback::Engine::open(playback::Config::default())
.await
.expect("an output device");
let canceller = engine.canceller(Config::default());
canceller.open(48_000, 1).expect("a mono microphone");
let mut sink = engine
.sink(playback::Input {
sample_rate: 48_000,
channels: 2,
..Default::default()
})
.expect("a sink");
let frames = 48_000 / 10;
let mut tone = Vec::with_capacity(frames * 2 * 4);
for frame in 0..frames {
let value = (std::f32::consts::TAU * 440.0 * frame as f32 / 48_000.0).sin() * 0.5;
for _ in 0..BUS_CHANNELS {
tone.extend_from_slice(&value.to_le_bytes());
}
}
let mut energy = 0.0f64;
let mut buf = vec![0.0f32; REFERENCE_FRAME * BUS_CHANNELS];
for _ in 0..20 {
sink.write(&tone).expect("write");
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
let mut state = canceller.inner.state.lock().unwrap();
while let Some(reference) = state.reference.as_mut()
&& reference.available_frames() >= REFERENCE_FRAME
{
reference.read_interleaved(&mut buf, false);
energy += buf.iter().map(|s| (*s as f64).powi(2)).sum::<f64>();
}
}
assert!(energy > 1.0, "the mix never reached the echo reference");
}
#[test]
fn interleaving_round_trips() {
let interleaved = [1.0, -1.0, 2.0, -2.0, 3.0, -3.0];
let mut planar = vec![0.0; 6];
deinterleave(&interleaved, &mut planar, 2);
assert_eq!(planar, vec![1.0, 2.0, 3.0, -1.0, -2.0, -3.0]);
let mut round = Vec::new();
interleave(&planar, 2, &mut round);
assert_eq!(round, interleaved);
}
}