Skip to main content

micro_wakeword/
audio.rs

1use std::sync::atomic::{AtomicUsize, Ordering};
2use std::sync::{Arc, mpsc::SyncSender};
3
4use cpal::traits::{DeviceTrait, HostTrait};
5use cpal::{Device, SampleFormat, Stream, StreamConfig};
6use rubato::audioadapter_buffers::direct::InterleavedSlice;
7use rubato::{Fft, FixedSync, Resampler};
8
9use crate::{AUDIO_BLOCK_SAMPLES, Error, Result, SAMPLE_RATE};
10
11#[derive(Clone, Debug, PartialEq, Eq)]
12pub struct AudioDevice {
13    pub index: usize,
14    pub name: String,
15}
16
17pub fn available_input_devices() -> Result<Vec<AudioDevice>> {
18    cpal::default_host()
19        .input_devices()
20        .map_err(|e| Error::Audio(e.to_string()))?
21        .enumerate()
22        .map(|(index, device)| {
23            Ok(AudioDevice {
24                index,
25                name: device.name().map_err(|e| Error::Audio(e.to_string()))?,
26            })
27        })
28        .collect()
29}
30
31// Keeping samples inline avoids one heap allocation for every 10 ms audio block.
32#[allow(clippy::large_enum_variant)]
33pub(crate) enum AudioEvent {
34    Samples([i16; AUDIO_BLOCK_SAMPLES]),
35    Error(String),
36}
37
38pub(crate) fn open_input(
39    selector: Option<&str>,
40    sender: SyncSender<AudioEvent>,
41    dropped: Arc<AtomicUsize>,
42) -> Result<Stream> {
43    let host = cpal::default_host();
44    let device = if let Some(selector) = selector {
45        let devices: Vec<_> = host
46            .input_devices()
47            .map_err(|e| Error::Audio(e.to_string()))?
48            .collect();
49        if let Ok(index) = selector.parse::<usize>() {
50            devices
51                .into_iter()
52                .nth(index)
53                .ok_or_else(|| Error::Audio("microphone index is out of range".into()))?
54        } else {
55            let needle = selector.to_lowercase();
56            devices
57                .into_iter()
58                .find(|d| d.name().is_ok_and(|n| n.to_lowercase().contains(&needle)))
59                .ok_or_else(|| Error::Audio(format!("no microphone name contains {selector:?}")))?
60        }
61    } else {
62        host.default_input_device()
63            .ok_or_else(|| Error::Audio("there is no default microphone".into()))?
64    };
65    let supported = device
66        .default_input_config()
67        .map_err(|e| Error::Audio(e.to_string()))?;
68    let format = supported.sample_format();
69    let config = supported.config();
70    match format {
71        SampleFormat::I16 => build_stream(&device, &config, sender, dropped, |s: i16| {
72            s as f32 / 32768.0
73        }),
74        SampleFormat::U16 => build_stream(&device, &config, sender, dropped, |s: u16| {
75            (s as f32 - 32768.0) / 32768.0
76        }),
77        SampleFormat::F32 => build_stream(&device, &config, sender, dropped, |s: f32| s),
78        other => Err(Error::Audio(format!(
79            "unsupported microphone sample format: {other:?}"
80        ))),
81    }
82}
83
84struct Pipeline {
85    channels: usize,
86    input_frames: usize,
87    pending_input: Vec<f32>,
88    pending_output: Vec<f32>,
89    resampler: Option<Fft<f32>>,
90    sender: SyncSender<AudioEvent>,
91    dropped: Arc<AtomicUsize>,
92}
93
94impl Pipeline {
95    fn new(
96        rate: u32,
97        channels: usize,
98        sender: SyncSender<AudioEvent>,
99        dropped: Arc<AtomicUsize>,
100    ) -> Result<Self> {
101        if channels == 0 || rate % 100 != 0 {
102            return Err(Error::Audio(format!(
103                "unsupported microphone format: {rate} Hz, {channels} channels"
104            )));
105        }
106        let input_frames = rate as usize / 100;
107        let resampler = if rate == SAMPLE_RATE {
108            None
109        } else {
110            Some(
111                Fft::new(
112                    rate as usize,
113                    SAMPLE_RATE as usize,
114                    input_frames,
115                    1,
116                    1,
117                    FixedSync::Input,
118                )
119                .map_err(|e| Error::Audio(e.to_string()))?,
120            )
121        };
122        Ok(Self {
123            channels,
124            input_frames,
125            pending_input: Vec::new(),
126            pending_output: Vec::new(),
127            resampler,
128            sender,
129            dropped,
130        })
131    }
132
133    fn push<T: Copy>(&mut self, data: &[T], convert: impl Fn(T) -> f32) -> Result<()> {
134        for frame in data.chunks_exact(self.channels) {
135            self.pending_input
136                .push(frame.iter().copied().map(&convert).sum::<f32>() / self.channels as f32);
137        }
138        while self.pending_input.len() >= self.input_frames {
139            if let Some(resampler) = &mut self.resampler {
140                let input = InterleavedSlice::new(
141                    &self.pending_input[..self.input_frames],
142                    1,
143                    self.input_frames,
144                )
145                .map_err(|e| Error::Audio(e.to_string()))?;
146                self.pending_output.extend(
147                    resampler
148                        .process(&input, 0, None)
149                        .map_err(|e| Error::Audio(e.to_string()))?
150                        .take_data(),
151                );
152            } else {
153                self.pending_output
154                    .extend_from_slice(&self.pending_input[..self.input_frames]);
155            }
156            self.pending_input.drain(..self.input_frames);
157            while self.pending_output.len() >= AUDIO_BLOCK_SAMPLES {
158                let mut block = [0_i16; AUDIO_BLOCK_SAMPLES];
159                for (out, sample) in block.iter_mut().zip(&self.pending_output) {
160                    *out = (sample.clamp(-1.0, 1.0) * i16::MAX as f32) as i16;
161                }
162                self.pending_output.drain(..AUDIO_BLOCK_SAMPLES);
163                if self.sender.try_send(AudioEvent::Samples(block)).is_err() {
164                    self.dropped.fetch_add(1, Ordering::Relaxed);
165                }
166            }
167        }
168        Ok(())
169    }
170}
171
172fn build_stream<T, F>(
173    device: &Device,
174    config: &StreamConfig,
175    sender: SyncSender<AudioEvent>,
176    dropped: Arc<AtomicUsize>,
177    convert: F,
178) -> Result<Stream>
179where
180    T: cpal::SizedSample + Copy,
181    F: Fn(T) -> f32 + Send + 'static,
182{
183    let mut pipeline = Pipeline::new(
184        config.sample_rate.0,
185        config.channels as usize,
186        sender.clone(),
187        dropped,
188    )?;
189    let error_sender = sender;
190    device
191        .build_input_stream(
192            config,
193            move |data: &[T], _| {
194                if let Err(error) = pipeline.push(data, &convert) {
195                    let _ = pipeline
196                        .sender
197                        .try_send(AudioEvent::Error(error.to_string()));
198                }
199            },
200            move |error| {
201                let _ = error_sender.try_send(AudioEvent::Error(error.to_string()));
202            },
203            None,
204        )
205        .map_err(|e| Error::Audio(e.to_string()))
206}
207
208#[cfg(test)]
209mod tests {
210    use super::*;
211
212    #[test]
213    fn downmixes_stereo_into_one_ten_millisecond_block() {
214        let (sender, receiver) = std::sync::mpsc::sync_channel(1);
215        let dropped = Arc::new(AtomicUsize::new(0));
216        let mut pipeline = Pipeline::new(SAMPLE_RATE, 2, sender, dropped).unwrap();
217        pipeline
218            .push(&[0.5_f32; AUDIO_BLOCK_SAMPLES * 2], |sample| sample)
219            .unwrap();
220        let AudioEvent::Samples(samples) = receiver.try_recv().unwrap() else {
221            panic!("expected audio")
222        };
223        assert!(samples.iter().all(|sample| *sample == 16_383));
224    }
225
226    #[test]
227    fn reports_queue_overflow() {
228        let (sender, _receiver) = std::sync::mpsc::sync_channel(0);
229        let dropped = Arc::new(AtomicUsize::new(0));
230        let mut pipeline = Pipeline::new(SAMPLE_RATE, 1, sender, dropped.clone()).unwrap();
231        pipeline
232            .push(&[0.0_f32; AUDIO_BLOCK_SAMPLES], |sample| sample)
233            .unwrap();
234        assert_eq!(dropped.load(Ordering::Relaxed), 1);
235    }
236}