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