Skip to main content

micro_wakeword/
audio.rs

1use std::sync::atomic::{AtomicUsize, Ordering};
2use std::sync::{Arc, Mutex, mpsc::SyncSender, mpsc::TrySendError};
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)]
33#[derive(Default)]
34pub(crate) struct AudioStatus {
35    dropped: AtomicUsize,
36    error: Mutex<Option<String>>,
37}
38
39impl AudioStatus {
40    pub(crate) fn dropped(&self) -> usize {
41        self.dropped.swap(0, Ordering::Relaxed)
42    }
43
44    pub(crate) fn report_error(&self, error: String) {
45        let mut pending = self.error.lock().unwrap_or_else(|lock| lock.into_inner());
46        if pending.is_none() {
47            *pending = Some(error);
48        }
49    }
50
51    pub(crate) fn take_error(&self) -> Option<String> {
52        self.error
53            .lock()
54            .unwrap_or_else(|lock| lock.into_inner())
55            .take()
56    }
57}
58
59pub(crate) fn open_input(
60    selector: Option<&str>,
61    sender: SyncSender<[i16; AUDIO_BLOCK_SAMPLES]>,
62    status: Arc<AudioStatus>,
63) -> Result<Stream> {
64    let host = cpal::default_host();
65    let device = if let Some(selector) = selector {
66        let devices: Vec<_> = host
67            .input_devices()
68            .map_err(|e| Error::Audio(e.to_string()))?
69            .collect();
70        if let Ok(index) = selector.parse::<usize>() {
71            devices
72                .into_iter()
73                .nth(index)
74                .ok_or_else(|| Error::Audio("microphone index is out of range".into()))?
75        } else {
76            let needle = selector.to_lowercase();
77            devices
78                .into_iter()
79                .find(|d| d.name().is_ok_and(|n| n.to_lowercase().contains(&needle)))
80                .ok_or_else(|| Error::Audio(format!("no microphone name contains {selector:?}")))?
81        }
82    } else {
83        host.default_input_device()
84            .ok_or_else(|| Error::Audio("there is no default microphone".into()))?
85    };
86    let supported = device
87        .default_input_config()
88        .map_err(|e| Error::Audio(e.to_string()))?;
89    let format = supported.sample_format();
90    let config = supported.config();
91    match format {
92        SampleFormat::I16 => build_stream(&device, &config, sender, status, |s: i16| {
93            s as f32 / 32768.0
94        }),
95        SampleFormat::U16 => build_stream(&device, &config, sender, status, |s: u16| {
96            (s as f32 - 32768.0) / 32768.0
97        }),
98        SampleFormat::F32 => build_stream(&device, &config, sender, status, |s: f32| s),
99        other => Err(Error::Audio(format!(
100            "unsupported microphone sample format: {other:?}"
101        ))),
102    }
103}
104
105struct Pipeline {
106    channels: usize,
107    input_frames: usize,
108    pending_input: Vec<f32>,
109    pending_output: Vec<f32>,
110    resampler: Option<Fft<f32>>,
111    sender: SyncSender<[i16; AUDIO_BLOCK_SAMPLES]>,
112    status: Arc<AudioStatus>,
113}
114
115impl Pipeline {
116    fn new(
117        rate: u32,
118        channels: usize,
119        sender: SyncSender<[i16; AUDIO_BLOCK_SAMPLES]>,
120        status: Arc<AudioStatus>,
121    ) -> Result<Self> {
122        if channels == 0 || rate % 100 != 0 {
123            return Err(Error::Audio(format!(
124                "unsupported microphone format: {rate} Hz, {channels} channels"
125            )));
126        }
127        let input_frames = rate as usize / 100;
128        let resampler = if rate == SAMPLE_RATE {
129            None
130        } else {
131            Some(
132                Fft::new(
133                    rate as usize,
134                    SAMPLE_RATE as usize,
135                    input_frames,
136                    1,
137                    1,
138                    FixedSync::Input,
139                )
140                .map_err(|e| Error::Audio(e.to_string()))?,
141            )
142        };
143        Ok(Self {
144            channels,
145            input_frames,
146            pending_input: Vec::new(),
147            pending_output: Vec::new(),
148            resampler,
149            sender,
150            status,
151        })
152    }
153
154    fn push<T: Copy>(&mut self, data: &[T], convert: impl Fn(T) -> f32) -> Result<()> {
155        for frame in data.chunks_exact(self.channels) {
156            self.pending_input
157                .push(frame.iter().copied().map(&convert).sum::<f32>() / self.channels as f32);
158        }
159        while self.pending_input.len() >= self.input_frames {
160            if let Some(resampler) = &mut self.resampler {
161                let input = InterleavedSlice::new(
162                    &self.pending_input[..self.input_frames],
163                    1,
164                    self.input_frames,
165                )
166                .map_err(|e| Error::Audio(e.to_string()))?;
167                self.pending_output.extend(
168                    resampler
169                        .process(&input, 0, None)
170                        .map_err(|e| Error::Audio(e.to_string()))?
171                        .take_data(),
172                );
173            } else {
174                self.pending_output
175                    .extend_from_slice(&self.pending_input[..self.input_frames]);
176            }
177            self.pending_input.drain(..self.input_frames);
178            while self.pending_output.len() >= AUDIO_BLOCK_SAMPLES {
179                let mut block = [0_i16; AUDIO_BLOCK_SAMPLES];
180                for (out, sample) in block.iter_mut().zip(&self.pending_output) {
181                    *out = (sample.clamp(-1.0, 1.0) * i16::MAX as f32) as i16;
182                }
183                self.pending_output.drain(..AUDIO_BLOCK_SAMPLES);
184                match self.sender.try_send(block) {
185                    Ok(()) => {}
186                    Err(TrySendError::Full(_)) => {
187                        self.status.dropped.fetch_add(1, Ordering::Relaxed);
188                    }
189                    Err(TrySendError::Disconnected(_)) => return Err(Error::AudioStreamEnded),
190                }
191            }
192        }
193        Ok(())
194    }
195}
196
197fn build_stream<T, F>(
198    device: &Device,
199    config: &StreamConfig,
200    sender: SyncSender<[i16; AUDIO_BLOCK_SAMPLES]>,
201    status: Arc<AudioStatus>,
202    convert: F,
203) -> Result<Stream>
204where
205    T: cpal::SizedSample + Copy,
206    F: Fn(T) -> f32 + Send + 'static,
207{
208    let mut pipeline = Pipeline::new(
209        config.sample_rate.0,
210        config.channels as usize,
211        sender.clone(),
212        status.clone(),
213    )?;
214    let stream_status = status.clone();
215    device
216        .build_input_stream(
217            config,
218            move |data: &[T], _| {
219                if let Err(error) = pipeline.push(data, &convert) {
220                    pipeline.status.report_error(error.to_string());
221                }
222            },
223            move |error| {
224                stream_status.report_error(error.to_string());
225            },
226            None,
227        )
228        .map_err(|e| Error::Audio(e.to_string()))
229}
230
231#[cfg(test)]
232mod tests {
233    use super::*;
234
235    #[test]
236    fn downmixes_stereo_into_one_ten_millisecond_block() {
237        let (sender, receiver) = std::sync::mpsc::sync_channel(1);
238        let status = Arc::new(AudioStatus::default());
239        let mut pipeline = Pipeline::new(SAMPLE_RATE, 2, sender, status).unwrap();
240        pipeline
241            .push(&[0.5_f32; AUDIO_BLOCK_SAMPLES * 2], |sample| sample)
242            .unwrap();
243        let samples = receiver.try_recv().unwrap();
244        assert!(samples.iter().all(|sample| *sample == 16_383));
245    }
246
247    #[test]
248    fn reports_queue_overflow() {
249        let (sender, _receiver) = std::sync::mpsc::sync_channel(0);
250        let status = Arc::new(AudioStatus::default());
251        let mut pipeline = Pipeline::new(SAMPLE_RATE, 1, sender, status.clone()).unwrap();
252        pipeline
253            .push(&[0.0_f32; AUDIO_BLOCK_SAMPLES], |sample| sample)
254            .unwrap();
255        assert_eq!(status.dropped(), 1);
256        assert_eq!(status.dropped(), 0);
257    }
258
259    #[test]
260    fn microphone_errors_are_not_lost_when_the_sample_queue_is_full() {
261        let (sender, _receiver) = std::sync::mpsc::sync_channel(1);
262        let status = Arc::new(AudioStatus::default());
263        let mut pipeline = Pipeline::new(SAMPLE_RATE, 1, sender, status.clone()).unwrap();
264        pipeline
265            .push(&[0.0_f32; AUDIO_BLOCK_SAMPLES], |sample| sample)
266            .unwrap();
267
268        status.report_error("device disconnected".into());
269
270        assert_eq!(status.take_error().as_deref(), Some("device disconnected"));
271        assert!(status.take_error().is_none());
272    }
273}