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#[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}