use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::{Arc, mpsc::SyncSender};
use cpal::traits::{DeviceTrait, HostTrait};
use cpal::{Device, SampleFormat, Stream, StreamConfig};
use rubato::audioadapter_buffers::direct::InterleavedSlice;
use rubato::{Fft, FixedSync, Resampler};
use crate::{AUDIO_BLOCK_SAMPLES, Error, Result, SAMPLE_RATE};
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct AudioDevice {
pub index: usize,
pub name: String,
}
pub fn available_input_devices() -> Result<Vec<AudioDevice>> {
cpal::default_host()
.input_devices()
.map_err(|e| Error::Audio(e.to_string()))?
.enumerate()
.map(|(index, device)| {
Ok(AudioDevice {
index,
name: device.name().map_err(|e| Error::Audio(e.to_string()))?,
})
})
.collect()
}
#[allow(clippy::large_enum_variant)]
pub(crate) enum AudioEvent {
Samples([i16; AUDIO_BLOCK_SAMPLES]),
Error(String),
}
pub(crate) fn open_input(
selector: Option<&str>,
sender: SyncSender<AudioEvent>,
dropped: Arc<AtomicUsize>,
) -> Result<Stream> {
let host = cpal::default_host();
let device = if let Some(selector) = selector {
let devices: Vec<_> = host
.input_devices()
.map_err(|e| Error::Audio(e.to_string()))?
.collect();
if let Ok(index) = selector.parse::<usize>() {
devices
.into_iter()
.nth(index)
.ok_or_else(|| Error::Audio("microphone index is out of range".into()))?
} else {
let needle = selector.to_lowercase();
devices
.into_iter()
.find(|d| d.name().is_ok_and(|n| n.to_lowercase().contains(&needle)))
.ok_or_else(|| Error::Audio(format!("no microphone name contains {selector:?}")))?
}
} else {
host.default_input_device()
.ok_or_else(|| Error::Audio("there is no default microphone".into()))?
};
let supported = device
.default_input_config()
.map_err(|e| Error::Audio(e.to_string()))?;
let format = supported.sample_format();
let config = supported.config();
match format {
SampleFormat::I16 => build_stream(&device, &config, sender, dropped, |s: i16| {
s as f32 / 32768.0
}),
SampleFormat::U16 => build_stream(&device, &config, sender, dropped, |s: u16| {
(s as f32 - 32768.0) / 32768.0
}),
SampleFormat::F32 => build_stream(&device, &config, sender, dropped, |s: f32| s),
other => Err(Error::Audio(format!(
"unsupported microphone sample format: {other:?}"
))),
}
}
struct Pipeline {
channels: usize,
input_frames: usize,
pending_input: Vec<f32>,
pending_output: Vec<f32>,
resampler: Option<Fft<f32>>,
sender: SyncSender<AudioEvent>,
dropped: Arc<AtomicUsize>,
}
impl Pipeline {
fn new(
rate: u32,
channels: usize,
sender: SyncSender<AudioEvent>,
dropped: Arc<AtomicUsize>,
) -> Result<Self> {
if channels == 0 || rate % 100 != 0 {
return Err(Error::Audio(format!(
"unsupported microphone format: {rate} Hz, {channels} channels"
)));
}
let input_frames = rate as usize / 100;
let resampler = if rate == SAMPLE_RATE {
None
} else {
Some(
Fft::new(
rate as usize,
SAMPLE_RATE as usize,
input_frames,
1,
1,
FixedSync::Input,
)
.map_err(|e| Error::Audio(e.to_string()))?,
)
};
Ok(Self {
channels,
input_frames,
pending_input: Vec::new(),
pending_output: Vec::new(),
resampler,
sender,
dropped,
})
}
fn push<T: Copy>(&mut self, data: &[T], convert: impl Fn(T) -> f32) -> Result<()> {
for frame in data.chunks_exact(self.channels) {
self.pending_input
.push(frame.iter().copied().map(&convert).sum::<f32>() / self.channels as f32);
}
while self.pending_input.len() >= self.input_frames {
if let Some(resampler) = &mut self.resampler {
let input = InterleavedSlice::new(
&self.pending_input[..self.input_frames],
1,
self.input_frames,
)
.map_err(|e| Error::Audio(e.to_string()))?;
self.pending_output.extend(
resampler
.process(&input, 0, None)
.map_err(|e| Error::Audio(e.to_string()))?
.take_data(),
);
} else {
self.pending_output
.extend_from_slice(&self.pending_input[..self.input_frames]);
}
self.pending_input.drain(..self.input_frames);
while self.pending_output.len() >= AUDIO_BLOCK_SAMPLES {
let mut block = [0_i16; AUDIO_BLOCK_SAMPLES];
for (out, sample) in block.iter_mut().zip(&self.pending_output) {
*out = (sample.clamp(-1.0, 1.0) * i16::MAX as f32) as i16;
}
self.pending_output.drain(..AUDIO_BLOCK_SAMPLES);
if self.sender.try_send(AudioEvent::Samples(block)).is_err() {
self.dropped.fetch_add(1, Ordering::Relaxed);
}
}
}
Ok(())
}
}
fn build_stream<T, F>(
device: &Device,
config: &StreamConfig,
sender: SyncSender<AudioEvent>,
dropped: Arc<AtomicUsize>,
convert: F,
) -> Result<Stream>
where
T: cpal::SizedSample + Copy,
F: Fn(T) -> f32 + Send + 'static,
{
let mut pipeline = Pipeline::new(
config.sample_rate.0,
config.channels as usize,
sender.clone(),
dropped,
)?;
let error_sender = sender;
device
.build_input_stream(
config,
move |data: &[T], _| {
if let Err(error) = pipeline.push(data, &convert) {
let _ = pipeline
.sender
.try_send(AudioEvent::Error(error.to_string()));
}
},
move |error| {
let _ = error_sender.try_send(AudioEvent::Error(error.to_string()));
},
None,
)
.map_err(|e| Error::Audio(e.to_string()))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn downmixes_stereo_into_one_ten_millisecond_block() {
let (sender, receiver) = std::sync::mpsc::sync_channel(1);
let dropped = Arc::new(AtomicUsize::new(0));
let mut pipeline = Pipeline::new(SAMPLE_RATE, 2, sender, dropped).unwrap();
pipeline
.push(&[0.5_f32; AUDIO_BLOCK_SAMPLES * 2], |sample| sample)
.unwrap();
let AudioEvent::Samples(samples) = receiver.try_recv().unwrap() else {
panic!("expected audio")
};
assert!(samples.iter().all(|sample| *sample == 16_383));
}
#[test]
fn reports_queue_overflow() {
let (sender, _receiver) = std::sync::mpsc::sync_channel(0);
let dropped = Arc::new(AtomicUsize::new(0));
let mut pipeline = Pipeline::new(SAMPLE_RATE, 1, sender, dropped.clone()).unwrap();
pipeline
.push(&[0.0_f32; AUDIO_BLOCK_SAMPLES], |sample| sample)
.unwrap();
assert_eq!(dropped.load(Ordering::Relaxed), 1);
}
}