diff --git a/README.md b/README.md index 9f891d5..f51fc22 100644 --- a/README.md +++ b/README.md @@ -86,7 +86,7 @@ Examples: | Processors | `GainProcessor` | | Sinks | `AsrSink`, `WavSink` | | ASR delivery | sync iteration and `asyncio` iteration | -| Output formats | `AsrSink`: `f32` / `i16`, mono or stereo | +| Output formats | `AsrSink` and `WavSink`: `f32` / `i16`, mono or stereo | | Metrics | `engine.stats()`, `asr_sink.stats()`, `wav_sink.stats()` | --- @@ -150,7 +150,13 @@ with macloop.AudioEngine() as engine: mic_for_asr = engine.route("mic_for_asr", stream=mic) mic_for_wav = engine.route("mic_for_wav", stream=mic) - wav_sink = macloop.WavSink(route=mic_for_wav, file="out/mic.wav") + wav_sink = macloop.WavSink( + route=mic_for_wav, + file="out/mic.wav", + sample_rate=16_000, + channels=1, + sample_format="i16", + ) asr_sink = macloop.AsrSink( routes=[mic_for_asr], chunk_frames=320, @@ -204,6 +210,9 @@ with macloop.AudioEngine() as engine: wav_sink = macloop.WavSink( routes=[mic_for_wav, zoom_for_wav], file="out/meeting.wav", + sample_rate=16_000, + channels=1, + sample_format="i16", ) asr_sink = macloop.AsrSink( @@ -224,7 +233,9 @@ Notes: - `AsrSink` emits **independent chunks per route**. - `WavSink` can mix several routes into one file. -- If `mix_gain` is not provided, `WavSink` uses `1 / N` by default. +- Set `sample_rate`, `channels`, and `sample_format` together to choose the WAV format. `channels` supports `1` or `2`; `sample_format` supports `"f32"` or `"i16"`. +- If all three format fields are omitted, `WavSink` writes the existing 48 kHz stereo float32 format. +- If `mix_gain` is omitted, `WavSink` uses `1 / N`. Set `mix_gain=1.0` to sum routes before format conversion; it does not normalize loudness. --- diff --git a/core_engine/src/converter.rs b/core_engine/src/converter.rs index c6c41bc..a9c0062 100644 --- a/core_engine/src/converter.rs +++ b/core_engine/src/converter.rs @@ -70,6 +70,9 @@ pub struct MasterFormatConverter { pending_interleaved: Vec, deinterleaved_in: Vec>, deinterleaved_out: Vec>, + output_delay_frames_remaining: usize, + total_input_frames: u64, + total_output_frames: u64, } impl MasterFormatConverter { @@ -111,6 +114,9 @@ impl MasterFormatConverter { None }; + let output_delay_frames_remaining = + resampler.as_ref().map(Resampler::output_delay).unwrap_or(0); + Ok(Self { input_format, output_format, @@ -119,9 +125,96 @@ impl MasterFormatConverter { pending_interleaved: Vec::new(), deinterleaved_in: Vec::new(), deinterleaved_out: Vec::new(), + output_delay_frames_remaining, + total_input_frames: 0, + total_output_frames: 0, }) } + pub fn finish(&mut self, output: &mut Vec) -> Result<(), InputConversionError> { + output.clear(); + if self.resampler.is_none() { + return Ok(()); + } + + let target_output_frames = self.expected_output_frames(); + let channels = self.output_format.channels as usize; + while self.total_output_frames < target_output_frames { + let input_frames_next = self + .resampler + .as_ref() + .expect("resampler presence checked above") + .input_frames_next(); + let pending_frames = self.pending_interleaved.len() / channels; + let partial_frames = pending_frames.min(input_frames_next); + let partial_samples = partial_frames * channels; + Self::deinterleave_padded_into( + &mut self.deinterleaved_in, + &self.pending_interleaved[..partial_samples], + channels, + input_frames_next, + ); + + let (used_in, used_out) = { + let resampler = self + .resampler + .as_mut() + .expect("resampler presence checked above"); + let adapter_in = + SequentialSliceOfVecs::new(&self.deinterleaved_in, channels, input_frames_next) + .map_err(|e| InputConversionError::ResamplerProcess(e.to_string()))?; + + let output_frames_cap = resampler.output_frames_max().max(1); + Self::ensure_output_storage( + &mut self.deinterleaved_out, + channels, + output_frames_cap, + ); + let mut adapter_out = SequentialSliceOfVecs::new_mut( + &mut self.deinterleaved_out, + channels, + output_frames_cap, + ) + .map_err(|e| InputConversionError::ResamplerProcess(e.to_string()))?; + + resampler + .process_into_buffer( + &adapter_in, + &mut adapter_out, + Some(&Indexing { + input_offset: 0, + output_offset: 0, + partial_len: Some(partial_frames), + active_channels_mask: None, + }), + ) + .map_err(|e| InputConversionError::ResamplerProcess(e.to_string()))? + }; + + if partial_frames > 0 { + let consumed_samples = used_in.min(partial_frames) * channels; + self.pending_interleaved.drain(..consumed_samples); + } + if used_out == 0 { + return Err(InputConversionError::ResamplerProcess( + "resampler produced no output while finishing".to_string(), + )); + } + + self.append_trimmed_resampler_output(used_out, target_output_frames, output); + } + + self.pending_interleaved.clear(); + Ok(()) + } + + fn expected_output_frames(&self) -> u64 { + let numerator = self.total_input_frames as u128 * self.output_format.sample_rate as u128; + let denominator = self.input_format.sample_rate as u128; + let frames = numerator.div_ceil(denominator); + frames.min(u64::MAX as u128) as u64 + } + fn convert_channels( input: &[f32], input_channels: u16, @@ -189,6 +282,23 @@ impl MasterFormatConverter { } } + fn deinterleave_padded_into( + storage: &mut Vec>, + input: &[f32], + channels: usize, + padded_frames: usize, + ) { + Self::ensure_channels_storage(storage, channels, padded_frames); + for channel in storage.iter_mut() { + channel.resize(padded_frames, 0.0); + } + for (frame_index, frame) in input.chunks_exact(channels).enumerate() { + for (channel_index, sample) in frame.iter().enumerate() { + storage[channel_index][frame_index] = *sample; + } + } + } + fn ensure_output_storage(storage: &mut Vec>, channels: usize, frames_cap: usize) { if storage.len() != channels { storage.clear(); @@ -202,15 +312,45 @@ impl MasterFormatConverter { } } - fn interleave_append_to(channels_data: &[Vec], frames: usize, out: &mut Vec) { + fn interleave_range_append_to( + channels_data: &[Vec], + first_frame: usize, + frames: usize, + out: &mut Vec, + ) { let channels = channels_data.len(); out.reserve(frames * channels); - for frame in 0..frames { + for frame in first_frame..first_frame + frames { for channel in channels_data.iter().take(channels) { out.push(channel[frame]); } } } + + fn append_trimmed_resampler_output( + &mut self, + produced_frames: usize, + target_output_frames: u64, + output: &mut Vec, + ) { + let delay_frames = self.output_delay_frames_remaining.min(produced_frames); + self.output_delay_frames_remaining -= delay_frames; + + let available_frames = produced_frames - delay_frames; + let remaining_target_frames = target_output_frames + .saturating_sub(self.total_output_frames) + .min(usize::MAX as u64) as usize; + let emitted_frames = available_frames.min(remaining_target_frames); + Self::interleave_range_append_to( + &self.deinterleaved_out, + delay_frames, + emitted_frames, + output, + ); + self.total_output_frames = self + .total_output_frames + .saturating_add(emitted_frames as u64); + } } impl InputConverter for MasterFormatConverter { @@ -233,25 +373,35 @@ impl InputConverter for MasterFormatConverter { self.output_format.channels, &mut self.channels_buffer, )?; + let converted_input_frames = + self.channels_buffer.len() / self.output_format.channels as usize; + self.total_input_frames = self + .total_input_frames + .saturating_add(converted_input_frames as u64); if self.resampler.is_none() { output.clear(); output.extend_from_slice(&self.channels_buffer); + self.total_output_frames = self + .total_output_frames + .saturating_add(converted_input_frames as u64); return Ok(()); } output.clear(); let channels = self.output_format.channels as usize; - let Some(resampler) = self.resampler.as_mut() else { - return Ok(()); - }; self.pending_interleaved .extend_from_slice(&self.channels_buffer); + let target_output_frames = self.expected_output_frames(); let mut consumed_frames = 0usize; loop { let pending_frames = self.pending_interleaved.len() / channels; - let input_frames_next = resampler.input_frames_next(); + let input_frames_next = self + .resampler + .as_ref() + .expect("resampler presence checked above") + .input_frames_next(); if pending_frames.saturating_sub(consumed_frames) < input_frames_next { break; } @@ -263,33 +413,45 @@ impl InputConverter for MasterFormatConverter { &self.pending_interleaved[start..end], channels, ); - let adapter_in = - SequentialSliceOfVecs::new(&self.deinterleaved_in, channels, input_frames_next) - .map_err(|e| InputConversionError::ResamplerProcess(e.to_string()))?; - let output_frames_cap = resampler.output_frames_max().max(1); - Self::ensure_output_storage(&mut self.deinterleaved_out, channels, output_frames_cap); - let mut adapter_out = SequentialSliceOfVecs::new_mut( - &mut self.deinterleaved_out, - channels, - output_frames_cap, - ) - .map_err(|e| InputConversionError::ResamplerProcess(e.to_string()))?; - - let (_used_in, used_out) = resampler - .process_into_buffer( - &adapter_in, - &mut adapter_out, - Some(&Indexing { - input_offset: 0, - output_offset: 0, - partial_len: None, - active_channels_mask: None, - }), + let used_out = { + let resampler = self + .resampler + .as_mut() + .expect("resampler presence checked above"); + let adapter_in = + SequentialSliceOfVecs::new(&self.deinterleaved_in, channels, input_frames_next) + .map_err(|e| InputConversionError::ResamplerProcess(e.to_string()))?; + + let output_frames_cap = resampler.output_frames_max().max(1); + Self::ensure_output_storage( + &mut self.deinterleaved_out, + channels, + output_frames_cap, + ); + let mut adapter_out = SequentialSliceOfVecs::new_mut( + &mut self.deinterleaved_out, + channels, + output_frames_cap, ) .map_err(|e| InputConversionError::ResamplerProcess(e.to_string()))?; - Self::interleave_append_to(&self.deinterleaved_out, used_out, output); + let (_used_in, used_out) = resampler + .process_into_buffer( + &adapter_in, + &mut adapter_out, + Some(&Indexing { + input_offset: 0, + output_offset: 0, + partial_len: None, + active_channels_mask: None, + }), + ) + .map_err(|e| InputConversionError::ResamplerProcess(e.to_string()))?; + used_out + }; + + self.append_trimmed_resampler_output(used_out, target_output_frames, output); consumed_frames += input_frames_next; } @@ -403,6 +565,105 @@ mod tests { assert_eq!(produced % MASTER_FORMAT.channels as usize, 0); } + fn convert_and_finish(converter: &mut MasterFormatConverter, input: &[f32]) -> Vec { + let mut output = Vec::new(); + converter.convert(input, &mut output).expect("convert"); + let mut all_output = output.clone(); + converter.finish(&mut output).expect("finish"); + all_output.extend_from_slice(&output); + all_output + } + + fn peak_index(samples: &[f32]) -> usize { + samples + .iter() + .enumerate() + .max_by(|(_, left), (_, right)| left.abs().total_cmp(&right.abs())) + .map(|(index, _)| index) + .expect("non-empty samples") + } + + #[test] + fn finish_preserves_resampled_duration_for_partial_input() { + let mut converter = MasterFormatConverter::new( + MASTER_FORMAT, + StreamFormat::with_sample_format(16_000, 1, SampleFormat::F32), + ) + .expect("converter"); + let input = vec![0.25_f32; 4_800 * MASTER_FORMAT.channels as usize]; + let mut output = Vec::new(); + + converter.convert(&input, &mut output).expect("convert"); + let mut all_output = output.clone(); + converter.finish(&mut output).expect("finish"); + all_output.extend_from_slice(&output); + + assert_eq!(all_output.len(), 1_600); + assert!(all_output[200..1_400] + .iter() + .all(|sample| (*sample - 0.25).abs() < 1e-3)); + + converter.finish(&mut output).expect("finish again"); + assert!(output.is_empty()); + } + + #[test] + fn resampler_discards_startup_delay_for_beginning_impulse() { + let mut converter = MasterFormatConverter::new( + MASTER_FORMAT, + StreamFormat::with_sample_format(16_000, 1, SampleFormat::F32), + ) + .expect("converter"); + let mut input = vec![0.0_f32; 4_800 * MASTER_FORMAT.channels as usize]; + input[..MASTER_FORMAT.channels as usize].fill(1.0); + + let output = convert_and_finish(&mut converter, &input); + + assert_eq!(output.len(), 1_600); + assert!(peak_index(&output) <= 1, "beginning impulse was delayed"); + assert!(output[..4].iter().any(|sample| sample.abs() > 0.1)); + } + + #[test] + fn finish_preserves_end_impulse_in_zero_tail() { + let mut converter = MasterFormatConverter::new( + MASTER_FORMAT, + StreamFormat::with_sample_format(16_000, 1, SampleFormat::F32), + ) + .expect("converter"); + let mut input = vec![0.0_f32; 4_800 * MASTER_FORMAT.channels as usize]; + let last_frame = input.len() - MASTER_FORMAT.channels as usize; + input[last_frame..].fill(1.0); + + let output = convert_and_finish(&mut converter, &input); + + assert_eq!(output.len(), 1_600); + assert!( + peak_index(&output) >= output.len() - 2, + "end impulse was truncated" + ); + assert!(output[output.len() - 4..] + .iter() + .any(|sample| sample.abs() > 0.1)); + } + + #[test] + fn finish_emits_short_resampled_input() { + let mut converter = MasterFormatConverter::new( + MASTER_FORMAT, + StreamFormat::with_sample_format(16_000, 1, SampleFormat::F32), + ) + .expect("converter"); + let input = vec![0.25_f32; 100 * MASTER_FORMAT.channels as usize]; + let mut output = Vec::new(); + + converter.convert(&input, &mut output).expect("convert"); + assert!(output.is_empty()); + converter.finish(&mut output).expect("finish"); + + assert_eq!(output.len(), 34); + } + #[test] fn rejects_non_f32_formats_for_now() { match MasterFormatConverter::new( diff --git a/core_engine/src/lib.rs b/core_engine/src/lib.rs index f1bc6e9..c31e1e1 100644 --- a/core_engine/src/lib.rs +++ b/core_engine/src/lib.rs @@ -26,7 +26,7 @@ pub use outputs::asr_sink::{ AsrChunkView, AsrInputId, AsrInputMetricsSnapshot, AsrSampleSlice, AsrSink, AsrSinkCallback, AsrSinkConfig, AsrSinkError, AsrSinkInput, AsrSinkMetricsSnapshot, }; -pub use outputs::wav_file::{WavFileOutput, WavOutputError, WavSinkMetricsSnapshot}; +pub use outputs::wav_file::{WavFileOutput, WavOutputError, WavSinkConfig, WavSinkMetricsSnapshot}; pub use permissions::{microphone_access, screen_capture_access}; pub use processor::{AudioProcessor, NodeId, OutputId, StreamId}; pub use sources::app_audio::{ diff --git a/core_engine/src/outputs/asr_sink.rs b/core_engine/src/outputs/asr_sink.rs index b6e4d49..3bf7f89 100644 --- a/core_engine/src/outputs/asr_sink.rs +++ b/core_engine/src/outputs/asr_sink.rs @@ -214,8 +214,10 @@ impl InputState { let input_channels = MASTER_FORMAT.channels.max(1) as usize; let ready_input_samples = self.drained_master.len() / input_channels * input_channels; if ready_input_samples > 0 { - self.converter - .convert(&self.drained_master[..ready_input_samples], &mut self.converted_output)?; + self.converter.convert( + &self.drained_master[..ready_input_samples], + &mut self.converted_output, + )?; if !self.converted_output.is_empty() { self.pending_output .extend_from_slice(&self.converted_output); diff --git a/core_engine/src/outputs/wav_file.rs b/core_engine/src/outputs/wav_file.rs index 3959506..e2397d0 100644 --- a/core_engine/src/outputs/wav_file.rs +++ b/core_engine/src/outputs/wav_file.rs @@ -1,10 +1,14 @@ +use crate::converter::{ + convert_f32_to_i16, InputConversionError, InputConverter, MasterFormatConverter, +}; use crate::engine::RouteConsumer; -use crate::format::{SampleFormat, StreamFormat}; +use crate::format::{SampleFormat, StreamFormat, MASTER_FORMAT}; use crate::metrics::{LatencyHistogram, LatencyHistogramSnapshot}; use ringbuf::traits::{Consumer, Observer}; use std::collections::VecDeque; use std::fs::File; use std::io::{BufWriter, Seek, Write}; +use std::panic::{catch_unwind, AssertUnwindSafe}; use std::path::Path; use std::sync::atomic::{AtomicBool, AtomicU64, Ordering}; use std::sync::Arc; @@ -14,6 +18,8 @@ use std::time::{Duration, Instant}; #[derive(Debug)] pub enum WavOutputError { UnsupportedSampleFormat(SampleFormat), + UnsupportedOutputChannels(u16), + Converter(InputConversionError), Io(String), Hound(String), ThreadPanic, @@ -26,6 +32,10 @@ impl std::fmt::Display for WavOutputError { Self::UnsupportedSampleFormat(fmt) => { write!(f, "unsupported WAV sample format: {:?}", fmt) } + Self::UnsupportedOutputChannels(channels) => { + write!(f, "unsupported output channels for wav sink: {channels}") + } + Self::Converter(err) => write!(f, "converter error: {err}"), Self::Io(err) => write!(f, "wav io error: {err}"), Self::Hound(err) => write!(f, "wav writer error: {err}"), Self::ThreadPanic => write!(f, "wav writer thread panicked"), @@ -48,6 +58,27 @@ impl From for WavOutputError { } } +impl From for WavOutputError { + fn from(value: InputConversionError) -> Self { + Self::Converter(value) + } +} + +#[derive(Debug, Clone, Copy)] +pub struct WavSinkConfig { + pub format: StreamFormat, + pub mix_gain: f32, +} + +impl Default for WavSinkConfig { + fn default() -> Self { + Self { + format: MASTER_FORMAT, + mix_gain: 1.0, + } + } +} + pub struct WavSinkMetrics { write_calls: AtomicU64, samples_written: AtomicU64, @@ -89,18 +120,39 @@ pub struct WavSinkMetricsSnapshot { pub finalize: LatencyHistogramSnapshot, } +struct WavThreadResult { + consumers: Vec, + result: Result<(), WavOutputError>, +} + pub struct WavFileOutput { stop: Arc, - handle: Option, WavOutputError>>>, + handle: Option>, metrics: Arc, } impl WavFileOutput { - pub fn try_spawn_mix( + pub fn validate_config(config: WavSinkConfig) -> Result<(), WavOutputError> { + if !(1..=2).contains(&config.format.channels) { + return Err(WavOutputError::UnsupportedOutputChannels( + config.format.channels, + )); + } + + let f32_format = StreamFormat::with_sample_format( + config.format.sample_rate, + config.format.channels, + SampleFormat::F32, + ); + MasterFormatConverter::new(MASTER_FORMAT, f32_format)?; + Self::wav_spec(config.format)?; + Ok(()) + } + + pub fn try_spawn_mix_with_config( writer: W, - format: StreamFormat, consumers: Vec, - mix_gain: f32, + config: WavSinkConfig, ) -> Result)> where W: Write + Seek + Send + 'static, @@ -112,98 +164,113 @@ impl WavFileOutput { )); } - let spec = match Self::wav_spec(format) { + if let Err(err) = Self::validate_config(config) { + return Err((err, consumers)); + } + + let spec = match Self::wav_spec(config.format) { Ok(spec) => spec, Err(err) => return Err((err, consumers)), }; + let f32_format = StreamFormat::with_sample_format( + config.format.sample_rate, + config.format.channels, + SampleFormat::F32, + ); + let mut converter = match MasterFormatConverter::new(MASTER_FORMAT, f32_format) { + Ok(converter) => converter, + Err(err) => return Err((WavOutputError::Converter(err), consumers)), + }; let stop = Arc::new(AtomicBool::new(false)); let stop_thread = stop.clone(); let metrics = Arc::new(WavSinkMetrics::default()); let metrics_thread = metrics.clone(); - let channels = format.channels.max(1) as u64; - let frame_channels = format.channels.max(1) as usize; + let frame_channels = MASTER_FORMAT.channels as usize; - let handle = thread::spawn(move || -> Result, WavOutputError> { - let mut writer = hound::WavWriter::new(writer, spec)?; - let idle_sleep = Duration::from_micros(200); + let handle = thread::spawn(move || { let mut consumers = consumers; - let mut input_buffers = vec![VecDeque::::new(); consumers.len()]; - let mut mixed_buffer = Vec::::new(); - - loop { - let stopping = stop_thread.load(Ordering::Acquire); - let mut drained_any = false; - for (consumer, buffer) in consumers.iter_mut().zip(input_buffers.iter_mut()) { - let drain_limit = consumer.occupied_len() / frame_channels * frame_channels; - let mut drained = 0_usize; - while drained < drain_limit { - let Some(sample) = consumer.try_pop() else { - break; - }; - buffer.push_back(sample); - drained_any = true; - drained += 1; + let result = catch_unwind(AssertUnwindSafe(|| -> Result<(), WavOutputError> { + let mut writer = hound::WavWriter::new(writer, spec)?; + let idle_sleep = Duration::from_micros(200); + let mut input_buffers = vec![VecDeque::::new(); consumers.len()]; + let mut mixed_buffer = Vec::::new(); + let mut converted_buffer = Vec::::new(); + let mut quantized_buffer = Vec::::new(); + + loop { + let stopping = stop_thread.load(Ordering::Acquire); + let mut drained_any = false; + for (consumer, buffer) in consumers.iter_mut().zip(input_buffers.iter_mut()) { + let drain_limit = consumer.occupied_len() / frame_channels * frame_channels; + let mut drained = 0_usize; + while drained < drain_limit { + let Some(sample) = consumer.try_pop() else { + break; + }; + buffer.push_back(sample); + drained_any = true; + drained += 1; + } } - } - let ready_samples = input_buffers - .iter() - .map(VecDeque::len) - .min() - .unwrap_or(0) - / frame_channels - * frame_channels; + let ready_samples = input_buffers.iter().map(VecDeque::len).min().unwrap_or(0) + / frame_channels + * frame_channels; - if ready_samples > 0 { - mixed_buffer.clear(); - mixed_buffer.reserve(ready_samples); + if ready_samples > 0 { + mixed_buffer.clear(); + mixed_buffer.reserve(ready_samples); - for _ in 0..ready_samples { - let mut mixed_sample = 0.0_f32; - for input in &mut input_buffers { - if let Some(sample) = input.pop_front() { - mixed_sample += sample; + for _ in 0..ready_samples { + let mut mixed_sample = 0.0_f32; + for input in &mut input_buffers { + if let Some(sample) = input.pop_front() { + mixed_sample += sample; + } } + mixed_buffer.push(mixed_sample * config.mix_gain); } - mixed_buffer.push(mixed_sample * mix_gain); - } - let write_start = Instant::now(); - for sample in &mixed_buffer { - writer.write_sample(*sample)?; + converter.convert(&mixed_buffer, &mut converted_buffer)?; + write_output_samples( + &mut writer, + &converted_buffer, + config.format, + &mut quantized_buffer, + &metrics_thread, + )?; } - let samples_written = mixed_buffer.len() as u64; - metrics_thread.write_calls.fetch_add(1, Ordering::Relaxed); - metrics_thread - .samples_written - .fetch_add(samples_written, Ordering::Relaxed); - metrics_thread - .frames_written - .fetch_add(samples_written / channels, Ordering::Relaxed); - metrics_thread - .write - .record(duration_to_u32_us(write_start.elapsed())); - } + if stopping { + converter.finish(&mut converted_buffer)?; + write_output_samples( + &mut writer, + &converted_buffer, + config.format, + &mut quantized_buffer, + &metrics_thread, + )?; + for input in &mut input_buffers { + input.clear(); + } + break; + } - if stopping { - for input in &mut input_buffers { - input.clear(); + if !drained_any && ready_samples == 0 { + thread::sleep(idle_sleep); } - break; } - if !drained_any && ready_samples == 0 { - thread::sleep(idle_sleep); - } - } + let finalize_start = Instant::now(); + let finalize_result = writer.finalize().map_err(WavOutputError::from); + metrics_thread + .finalize + .record(duration_to_u32_us(finalize_start.elapsed())); + finalize_result + })) + .unwrap_or(Err(WavOutputError::ThreadPanic)); - let finalize_start = Instant::now(); - writer.finalize()?; - metrics_thread - .finalize - .record(duration_to_u32_us(finalize_start.elapsed())); - Ok(consumers) + WavThreadResult { consumers, result } }); Ok(Self { @@ -213,6 +280,29 @@ impl WavFileOutput { }) } + pub fn try_spawn_mix( + writer: W, + format: StreamFormat, + consumers: Vec, + mix_gain: f32, + ) -> Result)> + where + W: Write + Seek + Send + 'static, + { + Self::try_spawn_mix_with_config(writer, consumers, WavSinkConfig { format, mix_gain }) + } + + pub fn spawn_mix_with_config( + writer: W, + consumers: Vec, + config: WavSinkConfig, + ) -> Result + where + W: Write + Seek + Send + 'static, + { + Self::try_spawn_mix_with_config(writer, consumers, config).map_err(|(err, _consumers)| err) + } + pub fn spawn_mix( writer: W, format: StreamFormat, @@ -244,6 +334,14 @@ impl WavFileOutput { Self::spawn(BufWriter::new(file), format, consumer) } + pub fn try_spawn_file_mix_with_config( + file: File, + consumers: Vec, + config: WavSinkConfig, + ) -> Result)> { + Self::try_spawn_mix_with_config(BufWriter::new(file), consumers, config) + } + pub fn try_spawn_file_mix( file: File, format: StreamFormat, @@ -253,6 +351,14 @@ impl WavFileOutput { Self::try_spawn_mix(BufWriter::new(file), format, consumers, mix_gain) } + pub fn spawn_file_mix_with_config( + file: File, + consumers: Vec, + config: WavSinkConfig, + ) -> Result { + Self::spawn_mix_with_config(BufWriter::new(file), consumers, config) + } + pub fn spawn_file_mix( file: File, format: StreamFormat, @@ -271,6 +377,15 @@ impl WavFileOutput { Self::spawn_file(file, format, consumer) } + pub fn spawn_path_mix_with_config>( + path: P, + consumers: Vec, + config: WavSinkConfig, + ) -> Result { + let file = File::create(path)?; + Self::spawn_file_mix_with_config(file, consumers, config) + } + pub fn spawn_path_mix>( path: P, format: StreamFormat, @@ -282,17 +397,15 @@ impl WavFileOutput { } fn wav_spec(format: StreamFormat) -> Result { - let sample_format = match format.sample_format { - SampleFormat::F32 => hound::SampleFormat::Float, - SampleFormat::I16 => { - return Err(WavOutputError::UnsupportedSampleFormat(SampleFormat::I16)) - } + let (sample_format, bits_per_sample) = match format.sample_format { + SampleFormat::F32 => (hound::SampleFormat::Float, 32), + SampleFormat::I16 => (hound::SampleFormat::Int, 16), }; Ok(hound::WavSpec { channels: format.channels, sample_rate: format.sample_rate, - bits_per_sample: 32, + bits_per_sample, sample_format, }) } @@ -301,17 +414,73 @@ impl WavFileOutput { self.metrics.snapshot() } - pub fn stop(&mut self) -> Result, WavOutputError> { + pub fn stop_with_consumers( + &mut self, + ) -> Result, (WavOutputError, Vec)> { self.stop.store(true, Ordering::Release); let Some(handle) = self.handle.take() else { - return Err(WavOutputError::AlreadyStopped); + return Err((WavOutputError::AlreadyStopped, Vec::new())); }; match handle.join() { - Ok(res) => res, - Err(_) => Err(WavOutputError::ThreadPanic), + Ok(WavThreadResult { + consumers, + result: Ok(()), + }) => Ok(consumers), + Ok(WavThreadResult { + consumers, + result: Err(err), + }) => Err((err, consumers)), + Err(_) => Err((WavOutputError::ThreadPanic, Vec::new())), + } + } + + pub fn stop(&mut self) -> Result, WavOutputError> { + self.stop_with_consumers().map_err(|(err, _consumers)| err) + } +} + +fn write_output_samples( + writer: &mut hound::WavWriter, + samples: &[f32], + format: StreamFormat, + quantized: &mut Vec, + metrics: &WavSinkMetrics, +) -> Result<(), WavOutputError> +where + W: Write + Seek, +{ + if samples.is_empty() { + return Ok(()); + } + + let write_start = Instant::now(); + match format.sample_format { + SampleFormat::F32 => { + for sample in samples { + writer.write_sample(*sample)?; + } + } + SampleFormat::I16 => { + convert_f32_to_i16(samples, quantized); + for sample in quantized { + writer.write_sample(*sample)?; + } } } + + let samples_written = samples.len() as u64; + metrics.write_calls.fetch_add(1, Ordering::Relaxed); + metrics + .samples_written + .fetch_add(samples_written, Ordering::Relaxed); + metrics + .frames_written + .fetch_add(samples_written / format.channels as u64, Ordering::Relaxed); + metrics + .write + .record(duration_to_u32_us(write_start.elapsed())); + Ok(()) } fn duration_to_u32_us(duration: Duration) -> u32 { @@ -332,19 +501,101 @@ mod tests { use ringbuf::traits::{Producer, Split}; use ringbuf::HeapRb; use std::fs::{self, File}; + use std::io::{self, Cursor, SeekFrom}; use std::path::PathBuf; - use std::sync::atomic::{AtomicBool, Ordering as AtomicOrdering}; + use std::sync::atomic::{AtomicBool, AtomicU64, Ordering as AtomicOrdering}; use std::sync::Arc; use std::thread; use std::time::Duration; use std::time::{SystemTime, UNIX_EPOCH}; fn test_wav_path() -> PathBuf { - let suffix = SystemTime::now() + static NEXT_PATH_ID: AtomicU64 = AtomicU64::new(0); + + let timestamp = SystemTime::now() .duration_since(UNIX_EPOCH) .expect("time") .as_nanos(); - std::env::temp_dir().join(format!("core_engine_wav_test_{suffix}.wav")) + let path_id = NEXT_PATH_ID.fetch_add(1, AtomicOrdering::Relaxed); + std::env::temp_dir().join(format!("core_engine_wav_test_{timestamp}_{path_id}.wav")) + } + + struct WriteFailingWriter; + + impl Write for WriteFailingWriter { + fn write(&mut self, _buffer: &[u8]) -> io::Result { + Err(io::Error::new( + io::ErrorKind::PermissionDenied, + "test writer rejected write", + )) + } + + fn flush(&mut self) -> io::Result<()> { + Ok(()) + } + } + + impl Seek for WriteFailingWriter { + fn seek(&mut self, _position: SeekFrom) -> io::Result { + Ok(0) + } + } + + struct PanickingWriter { + inner: Cursor>, + panic_next_write: bool, + } + + impl PanickingWriter { + fn new() -> Self { + Self { + inner: Cursor::new(Vec::new()), + panic_next_write: true, + } + } + } + + impl Write for PanickingWriter { + fn write(&mut self, buffer: &[u8]) -> io::Result { + if std::mem::take(&mut self.panic_next_write) { + panic!("test writer panicked"); + } + self.inner.write(buffer) + } + + fn flush(&mut self) -> io::Result<()> { + self.inner.flush() + } + } + + impl Seek for PanickingWriter { + fn seek(&mut self, position: SeekFrom) -> io::Result { + self.inner.seek(position) + } + } + + #[derive(Default)] + struct FinalizeFailingWriter { + inner: Cursor>, + } + + impl Write for FinalizeFailingWriter { + fn write(&mut self, buffer: &[u8]) -> io::Result { + self.inner.write(buffer) + } + + fn flush(&mut self) -> io::Result<()> { + self.inner.flush() + } + } + + impl Seek for FinalizeFailingWriter { + fn seek(&mut self, _position: SeekFrom) -> io::Result { + Err(io::Error::new( + io::ErrorKind::PermissionDenied, + "test writer rejected finalize seek", + )) + } } #[test] @@ -364,18 +615,102 @@ mod tests { } #[test] - fn wav_spec_rejects_i16_output() { - let err = WavFileOutput::wav_spec(StreamFormat::with_sample_format( - 48_000, - 2, + fn write_error_returns_consumers_for_reuse() { + let ring = HeapRb::::new(32); + let (mut producer, consumer) = ring.split(); + let mut failing = + WavFileOutput::spawn_mix(WriteFailingWriter, MASTER_FORMAT, vec![consumer], 1.0) + .expect("spawn failing wav"); + + let (err, consumers) = match failing.stop_with_consumers() { + Err(result) => result, + Ok(_) => panic!("writer should fail"), + }; + assert!(matches!(err, WavOutputError::Hound(_))); + assert_eq!(consumers.len(), 1); + + assert_eq!(producer.push_slice(&[0.25_f32; 8]), 8); + let mut reused = + WavFileOutput::spawn_mix(Cursor::new(Vec::::new()), MASTER_FORMAT, consumers, 1.0) + .expect("reuse returned consumer"); + reused.stop().expect("stop reused wav"); + } + + #[test] + fn writer_panic_returns_all_consumers_for_reuse() { + let ring_a = HeapRb::::new(32); + let ring_b = HeapRb::::new(32); + let (mut producer_a, consumer_a) = ring_a.split(); + let (mut producer_b, consumer_b) = ring_b.split(); + let mut panicking = WavFileOutput::spawn_mix( + PanickingWriter::new(), + MASTER_FORMAT, + vec![consumer_a, consumer_b], + 1.0, + ) + .expect("spawn panicking wav"); + + let (err, consumers) = match panicking.stop_with_consumers() { + Err(result) => result, + Ok(_) => panic!("writer should panic"), + }; + assert!(matches!(err, WavOutputError::ThreadPanic)); + assert_eq!(consumers.len(), 2); + + assert_eq!(producer_a.push_slice(&[0.25_f32; 8]), 8); + assert_eq!(producer_b.push_slice(&[0.75_f32; 8]), 8); + let mut reused = + WavFileOutput::spawn_mix(Cursor::new(Vec::::new()), MASTER_FORMAT, consumers, 1.0) + .expect("reuse returned consumers"); + assert_eq!(reused.stop().expect("stop reused wav").len(), 2); + } + + #[test] + fn finalize_error_returns_consumers_and_records_stats() { + let ring = HeapRb::::new(32); + let (_producer, consumer) = ring.split(); + let mut wav = WavFileOutput::spawn_mix( + FinalizeFailingWriter::default(), + MASTER_FORMAT, + vec![consumer], + 1.0, + ) + .expect("spawn finalize-failing wav"); + + let (err, consumers) = match wav.stop_with_consumers() { + Err(result) => result, + Ok(_) => panic!("finalize should fail"), + }; + + assert!(matches!(err, WavOutputError::Hound(_))); + assert_eq!(consumers.len(), 1); + assert_eq!(wav.stats().finalize.count, 1); + } + + #[test] + fn wav_spec_supports_i16_output() { + let spec = WavFileOutput::wav_spec(StreamFormat::with_sample_format( + 16_000, + 1, SampleFormat::I16, )) - .expect_err("i16 unsupported"); + .expect("i16 spec"); - assert!(matches!( - err, - WavOutputError::UnsupportedSampleFormat(SampleFormat::I16) - )); + assert_eq!(spec.channels, 1); + assert_eq!(spec.sample_rate, 16_000); + assert_eq!(spec.bits_per_sample, 16); + assert_eq!(spec.sample_format, hound::SampleFormat::Int); + } + + #[test] + fn validate_config_rejects_unsupported_channels() { + let err = WavFileOutput::validate_config(WavSinkConfig { + format: StreamFormat::new(48_000, 3), + mix_gain: 1.0, + }) + .expect_err("unsupported channels"); + + assert!(matches!(err, WavOutputError::UnsupportedOutputChannels(3))); } #[test] @@ -384,7 +719,7 @@ mod tests { } #[test] - fn writes_wav_from_routed_stream_to_file() { + fn writes_default_master_f32_wav_from_routed_stream_to_file() { let path = test_wav_path(); let mut engine = AudioEngineController::new(32, 32, 4096); @@ -430,6 +765,194 @@ mod tests { let _ = fs::remove_file(path); } + #[test] + fn converts_master_stereo_to_16khz_mono_i16() { + let path = test_wav_path(); + + let mut engine = AudioEngineController::new(32, 32, 32_768); + let stream = "capture".to_string(); + let output = "wav".to_string(); + let mut pipeline = engine + .create_stream(stream.clone(), SourceType::SystemAudio, 8, 4) + .expect("create stream"); + engine.route(&stream, &output).expect("route output"); + + let consumer = engine + .take_output_consumer(&output) + .expect("output consumer present"); + let file = File::create(&path).expect("create output file"); + let mut wav = WavFileOutput::spawn_file_mix_with_config( + file, + vec![consumer], + WavSinkConfig { + format: StreamFormat::with_sample_format(16_000, 1, SampleFormat::I16), + mix_gain: 1.0, + }, + ) + .expect("spawn converted wav output"); + + let mut frame = [0.25_f32; 320]; + for _ in 0..30 { + pipeline.process_callback(&mut frame); + } + wav.stop().expect("stop wav output"); + let stats = wav.stats(); + + let mut reader = hound::WavReader::open(&path).expect("open wav"); + let spec = reader.spec(); + let samples: Vec = reader + .samples::() + .map(|sample| sample.expect("sample")) + .collect(); + assert_eq!(spec.channels, 1); + assert_eq!(spec.sample_rate, 16_000); + assert_eq!(spec.bits_per_sample, 16); + assert_eq!(spec.sample_format, hound::SampleFormat::Int); + assert_eq!(samples.len(), 1_600); + assert!(samples[500..samples.len() - 200] + .iter() + .all(|sample| (*sample - 8192).abs() <= 2)); + assert_eq!(stats.samples_written, samples.len() as u64); + assert_eq!(stats.frames_written, samples.len() as u64); + + let _ = fs::remove_file(path); + } + + #[test] + fn converted_wav_preserves_final_impulse_tail() { + let path = test_wav_path(); + let ring = HeapRb::::new(10_000); + let (mut producer, consumer) = ring.split(); + let file = File::create(&path).expect("create output file"); + let mut wav = WavFileOutput::spawn_file_mix_with_config( + file, + vec![consumer], + WavSinkConfig { + format: StreamFormat::with_sample_format(16_000, 1, SampleFormat::F32), + mix_gain: 1.0, + }, + ) + .expect("spawn converted wav output"); + let mut input = vec![0.0_f32; 4_800 * MASTER_FORMAT.channels as usize]; + let last_frame = input.len() - MASTER_FORMAT.channels as usize; + input[last_frame..].fill(1.0); + assert_eq!(producer.push_slice(&input), input.len()); + + wav.stop().expect("stop wav output"); + + let mut reader = hound::WavReader::open(&path).expect("open wav"); + let samples: Vec = reader + .samples::() + .map(|sample| sample.expect("sample")) + .collect(); + let peak_index = samples + .iter() + .enumerate() + .max_by(|(_, left), (_, right)| left.abs().total_cmp(&right.abs())) + .map(|(index, _)| index) + .expect("wav samples"); + assert_eq!(samples.len(), 1_600); + assert!( + peak_index >= samples.len() - 2, + "final impulse was truncated" + ); + assert!(samples[samples.len() - 4..] + .iter() + .any(|sample| sample.abs() > 0.1)); + + let _ = fs::remove_file(path); + } + + #[test] + fn explicit_mix_gain_one_sums_routes_before_conversion() { + let path = test_wav_path(); + + let mut engine = AudioEngineController::new(32, 32, 4096); + let stream = "capture".to_string(); + let output_a = "wav_a".to_string(); + let output_b = "wav_b".to_string(); + let mut pipeline = engine + .create_stream(stream.clone(), SourceType::SystemAudio, 8, 4) + .expect("create stream"); + engine.route(&stream, &output_a).expect("route output a"); + engine.route(&stream, &output_b).expect("route output b"); + + let consumer_a = engine + .take_output_consumer(&output_a) + .expect("output consumer a present"); + let consumer_b = engine + .take_output_consumer(&output_b) + .expect("output consumer b present"); + let file = File::create(&path).expect("create output file"); + let mut wav = WavFileOutput::spawn_file_mix_with_config( + file, + vec![consumer_a, consumer_b], + WavSinkConfig { + format: StreamFormat::with_sample_format(48_000, 1, SampleFormat::I16), + mix_gain: 1.0, + }, + ) + .expect("spawn mixed wav output"); + + let mut frame = [0.25_f32; 8]; + pipeline.process_callback(&mut frame); + wav.stop().expect("stop wav output"); + + let mut reader = hound::WavReader::open(&path).expect("open wav"); + let samples: Vec = reader + .samples::() + .map(|sample| sample.expect("sample")) + .collect(); + assert_eq!(samples, vec![16384; 4]); + + let _ = fs::remove_file(path); + } + + #[test] + fn pcm16_output_clips_mixed_samples_safely() { + let path = test_wav_path(); + + let mut engine = AudioEngineController::new(32, 32, 4096); + let stream = "capture".to_string(); + let output_a = "wav_a".to_string(); + let output_b = "wav_b".to_string(); + let mut pipeline = engine + .create_stream(stream.clone(), SourceType::SystemAudio, 8, 4) + .expect("create stream"); + engine.route(&stream, &output_a).expect("route output a"); + engine.route(&stream, &output_b).expect("route output b"); + + let consumer_a = engine + .take_output_consumer(&output_a) + .expect("output consumer a present"); + let consumer_b = engine + .take_output_consumer(&output_b) + .expect("output consumer b present"); + let file = File::create(&path).expect("create output file"); + let mut wav = WavFileOutput::spawn_file_mix_with_config( + file, + vec![consumer_a, consumer_b], + WavSinkConfig { + format: StreamFormat::with_sample_format(48_000, 1, SampleFormat::I16), + mix_gain: 1.0, + }, + ) + .expect("spawn mixed wav output"); + + let mut frame = [0.75_f32, 0.75, -0.75, -0.75]; + pipeline.process_callback(&mut frame); + wav.stop().expect("stop wav output"); + + let mut reader = hound::WavReader::open(&path).expect("open wav"); + let samples: Vec = reader + .samples::() + .map(|sample| sample.expect("sample")) + .collect(); + assert_eq!(samples, vec![i16::MAX, i16::MIN + 1]); + + let _ = fs::remove_file(path); + } + #[test] fn mixes_multiple_routes_with_mix_gain() { let path = test_wav_path(); diff --git a/examples/write_to_wav.py b/examples/write_to_wav.py index cad4ca6..365f783 100644 --- a/examples/write_to_wav.py +++ b/examples/write_to_wav.py @@ -18,6 +18,14 @@ def main() -> None: parser.add_argument("--seconds", type=float, default=5.0, help="How long to record.") parser.add_argument("--device-id", type=int, default=None, help="Optional microphone device id.") parser.add_argument("--output", default="out/mic.wav", help="Output WAV path.") + parser.add_argument("--sample-rate", type=int, default=16_000, help="Output sample rate.") + parser.add_argument("--channels", type=int, choices=(1, 2), default=1, help="Output channels.") + parser.add_argument( + "--sample-format", + choices=("f32", "i16"), + default="i16", + help="Output sample format.", + ) parser.add_argument("--list-mics", action="store_true", help="List available microphones and exit.") args = parser.parse_args() @@ -37,7 +45,13 @@ def main() -> None: vpio_enabled=vpio_enabled, ) mic_for_wav = engine.route(stream=mic) - wav_sink = macloop.WavSink(route=mic_for_wav, file=output) + wav_sink = macloop.WavSink( + route=mic_for_wav, + file=output, + sample_rate=args.sample_rate, + channels=args.channels, + sample_format=args.sample_format, + ) print(f"Recording microphone to {output.resolve()} for {args.seconds:.1f}s...") time.sleep(args.seconds) diff --git a/macloop/__init__.py b/macloop/__init__.py index d382835..b574727 100644 --- a/macloop/__init__.py +++ b/macloop/__init__.py @@ -24,7 +24,7 @@ _AudioEngineBackend, _WavSinkBackend, _create_asr_sink, - _create_wav_sink, + _create_wav_sink_with_config, ) from ._macloop import list_applications as _list_applications from ._macloop import list_displays as _list_displays @@ -35,6 +35,9 @@ AudioSamples = Union[npt.NDArray[np.int16], npt.NDArray[np.float32]] _STOP = object() +_WAV_DEFAULT_SAMPLE_RATE = 48_000 +_WAV_DEFAULT_CHANNELS = 2 +_WAV_DEFAULT_SAMPLE_FORMAT = "f32" def _generate_id(prefix: str) -> str: @@ -632,8 +635,16 @@ def __init__( route: Optional[RouteHandle] = None, routes: Optional[Sequence[RouteHandle]] = None, file: Any, + sample_rate: Optional[int] = None, + channels: Optional[int] = None, + sample_format: Optional[str] = None, mix_gain: Optional[float] = None, ) -> None: + output_sample_rate, output_channels, output_sample_format = _resolve_wav_format( + sample_rate=sample_rate, + channels=channels, + sample_format=sample_format, + ) route_list = _resolve_sink_routes(route=route, routes=routes) engine = _engine_from_routes(route_list) engine._ensure_open() @@ -644,11 +655,14 @@ def __init__( effective_mix_gain = mix_gain if mix_gain is not None else (1.0 / len(route_ids)) fd, should_close = _resolve_wav_fd(file) try: - backend = _create_wav_sink( + backend = _create_wav_sink_with_config( engine._backend, sink_id, route_ids, fd, + output_sample_rate, + output_channels, + output_sample_format, float(effective_mix_gain), ) finally: @@ -733,6 +747,30 @@ def _resolve_sink_routes( return routes +def _resolve_wav_format( + *, + sample_rate: Optional[int], + channels: Optional[int], + sample_format: Optional[str], +) -> Tuple[int, int, str]: + values = (sample_rate, channels, sample_format) + if all(value is None for value in values): + return ( + _WAV_DEFAULT_SAMPLE_RATE, + _WAV_DEFAULT_CHANNELS, + _WAV_DEFAULT_SAMPLE_FORMAT, + ) + if any(value is None for value in values): + raise ValueError( + "sample_rate, channels, and sample_format must be provided together" + ) + + assert sample_rate is not None + assert channels is not None + assert sample_format is not None + return sample_rate, channels, sample_format + + def _resolve_wav_fd(file: Any) -> Tuple[int, bool]: if isinstance(file, int): if file < 0: diff --git a/macloop/__init__.pyi b/macloop/__init__.pyi index 36c01c2..e763c3f 100644 --- a/macloop/__init__.pyi +++ b/macloop/__init__.pyi @@ -1,6 +1,6 @@ from __future__ import annotations -from typing import Any, AsyncIterator, Iterator, Optional, Sequence, Union +from typing import Any, AsyncIterator, Iterator, Literal, Optional, Sequence, Union, overload import numpy as np import numpy.typing as npt @@ -154,6 +154,7 @@ class AsrSink: class WavSink: id: str + @overload def __init__( self, id: Optional[str] = None, @@ -161,6 +162,22 @@ class WavSink: route: Optional[RouteHandle] = None, routes: Optional[Sequence[RouteHandle]] = None, file: Any, + sample_rate: None = None, + channels: None = None, + sample_format: None = None, + mix_gain: Optional[float] = None, + ) -> None: ... + @overload + def __init__( + self, + id: Optional[str] = None, + *, + route: Optional[RouteHandle] = None, + routes: Optional[Sequence[RouteHandle]] = None, + file: Any, + sample_rate: int, + channels: Literal[1, 2], + sample_format: Literal["f32", "i16"], mix_gain: Optional[float] = None, ) -> None: ... def stats(self) -> WavSinkStats: ... diff --git a/macloop/_macloop.pyi b/macloop/_macloop.pyi index 28a000d..9e7511a 100644 --- a/macloop/_macloop.pyi +++ b/macloop/_macloop.pyi @@ -1,6 +1,6 @@ from __future__ import annotations -from typing import Any, Callable, Optional, Union +from typing import Any, Callable, Literal, Optional, Union import numpy as np import numpy.typing as npt @@ -45,7 +45,7 @@ class _AsrSinkBackend: class _WavSinkBackend: def stats(self) -> WavSinkStats: ... - def close(self) -> None: ... + def close(self, engine: Optional[_AudioEngineBackend] = None) -> None: ... class AsrInputStats: @@ -106,3 +106,13 @@ def _create_wav_sink( fd: int, mix_gain: float, ) -> _WavSinkBackend: ... +def _create_wav_sink_with_config( + engine: _AudioEngineBackend, + sink_id: str, + route_ids: list[str], + fd: int, + sample_rate: int, + channels: Literal[1, 2], + sample_format: Literal["f32", "i16"], + mix_gain: float, +) -> _WavSinkBackend: ... diff --git a/python_ffi/src/lib.rs b/python_ffi/src/lib.rs index b018844..a0bcd5a 100644 --- a/python_ffi/src/lib.rs +++ b/python_ffi/src/lib.rs @@ -1,12 +1,13 @@ mod stats; use core_engine::{ - microphone_access as core_microphone_access, screen_capture_access as core_screen_capture_access, - AppAudioSource, AppAudioSourceConfig, ApplicationInfo, AsrChunkView, AsrSampleSlice, AsrSink, - AsrSinkCallback, AsrSinkConfig, AsrSinkInput, AsrSinkMetricsSnapshot, AudioEngineController, - AudioProcessor, DisplayInfo, EngineError, MicInfo, MicrophoneSource, MicrophoneSourceConfig, - NodeMetrics, RouteConsumer, SampleFormat, SourceType, StreamFormat, SyntheticSource, - SyntheticSourceConfig, SystemAudioSource, SystemAudioSourceConfig, WavFileOutput, + microphone_access as core_microphone_access, + screen_capture_access as core_screen_capture_access, AppAudioSource, AppAudioSourceConfig, + ApplicationInfo, AsrChunkView, AsrSampleSlice, AsrSink, AsrSinkCallback, AsrSinkConfig, + AsrSinkInput, AsrSinkMetricsSnapshot, AudioEngineController, AudioProcessor, DisplayInfo, + EngineError, MicInfo, MicrophoneSource, MicrophoneSourceConfig, NodeMetrics, RouteConsumer, + SampleFormat, SourceType, StreamFormat, SyntheticSource, SyntheticSourceConfig, + SystemAudioSource, SystemAudioSourceConfig, WavFileOutput, WavSinkConfig, WavSinkMetricsSnapshot, }; use numpy::ToPyArray; @@ -951,20 +952,12 @@ impl PythonAsrCallback { input_id, frames, samples, - } => ( - input_id, - frames, - samples.to_pyarray(py).into_any().unbind(), - ), + } => (input_id, frames, samples.to_pyarray(py).into_any().unbind()), AsrWorkerPayload::I16 { input_id, frames, samples, - } => ( - input_id, - frames, - samples.to_pyarray(py).into_any().unbind(), - ), + } => (input_id, frames, samples.to_pyarray(py).into_any().unbind()), }; if let Err(err) = callback.call1(py, (input_id, frames, samples_obj)) { @@ -1097,7 +1090,9 @@ impl PyAsrSinkBackend { .map(|input| (input.input_id, input.consumer)) .collect(), ) - .map_err(|e| PyRuntimeError::new_err(format!("failed to restore asr sink routes: {e}")))?; + .map_err(|e| { + PyRuntimeError::new_err(format!("failed to restore asr sink routes: {e}")) + })?; } Ok(()) } @@ -1160,21 +1155,36 @@ impl PyWavSinkBackend { }; let route_ids = self.route_ids.clone(); - let (stop_result, final_stats) = py.detach(move || { - let stop_result = sink.stop().map_err(|e| e.to_string()); + let (stop_result, consumers, final_stats) = py.detach(move || { + let (stop_result, consumers) = match sink.stop_with_consumers() { + Ok(consumers) => (Ok(()), consumers), + Err((err, consumers)) => (Err(err.to_string()), consumers), + }; let final_stats = sink.stats(); - (stop_result, final_stats) + (stop_result, consumers, final_stats) }); self.final_stats = Some(final_stats); - let consumers = stop_result - .map_err(|e| PyRuntimeError::new_err(format!("failed to stop wav sink: {e}")))?; - if let Some(engine) = engine.as_mut() { + let restore_result = if let Some(engine) = engine.as_mut() { engine .restore_route_consumers(route_ids.into_iter().zip(consumers).collect()) - .map_err(|e| PyRuntimeError::new_err(format!("failed to restore wav sink routes: {e}")))?; + .map_err(|err| err.to_string()) + } else { + Ok(()) + }; + + match (stop_result, restore_result) { + (Ok(()), Ok(())) => Ok(()), + (Err(stop_err), Ok(())) => Err(PyRuntimeError::new_err(format!( + "failed to stop wav sink: {stop_err}" + ))), + (Ok(()), Err(restore_err)) => Err(PyRuntimeError::new_err(format!( + "failed to restore wav sink routes: {restore_err}" + ))), + (Err(stop_err), Err(restore_err)) => Err(PyRuntimeError::new_err(format!( + "failed to stop wav sink: {stop_err}; additionally failed to restore wav sink routes: {restore_err}" + ))), } - Ok(()) } fn close_no_restore(&mut self, py: Python<'_>) -> PyResult<()> { @@ -1182,13 +1192,11 @@ impl PyWavSinkBackend { return Ok(()); }; - let (stop_result, final_stats) = py.detach(move || { - let stop_result = sink.stop().map_err(|e| e.to_string()); - let final_stats = sink.stats(); - (stop_result, final_stats) + let final_stats = py.detach(move || { + let _ = sink.stop_with_consumers(); + sink.stats() }); self.final_stats = Some(final_stats); - let _ = stop_result; Ok(()) } } @@ -1760,11 +1768,25 @@ impl PyAudioEngineBackend { py: Python<'_>, route_ids: Vec, fd: i32, + sample_rate: u32, + channels: u16, + sample_format: String, mix_gain: f32, ) -> PyResult { self.ensure_open()?; self.ensure_route_consumers_available(&route_ids)?; + let config = WavSinkConfig { + format: StreamFormat::with_sample_format( + sample_rate, + channels, + parse_sample_format(&sample_format)?, + ), + mix_gain, + }; + WavFileOutput::validate_config(config) + .map_err(|err| PyValueError::new_err(err.to_string()))?; + let file = duplicate_file_descriptor(fd) .map_err(|e| PyOSError::new_err(format!("failed to duplicate file descriptor: {e}")))?; let stream_ids = self.stream_ids_for_routes(&route_ids)?; @@ -1789,14 +1811,13 @@ impl PyAudioEngineBackend { let route_consumers = self.take_route_consumers(&route_ids)?; let route_ids_for_sink = route_ids.clone(); - let master_format = self.controller.master_format(); let detached_result: DetachedWavStartResult = py.detach(move || { let consumers = route_consumers .into_iter() .map(|(_, consumer)| consumer) .collect(); - match WavFileOutput::try_spawn_file_mix(file, master_format, consumers, mix_gain) { + match WavFileOutput::try_spawn_file_mix_with_config(file, consumers, config) { Ok(sink) => Ok(sink), Err((err, consumers)) => Err(( format!("failed to create wav sink: {err}"), @@ -1909,7 +1930,31 @@ fn create_wav_sink( mix_gain: f32, ) -> PyResult { let _ = sink_id; - engine.build_wav_sink(py, route_ids, fd, mix_gain) + engine.build_wav_sink(py, route_ids, fd, 48_000, 2, "f32".to_string(), mix_gain) +} + +#[pyfunction(name = "_create_wav_sink_with_config")] +fn create_wav_sink_with_config( + py: Python<'_>, + mut engine: PyRefMut<'_, PyAudioEngineBackend>, + sink_id: String, + route_ids: Vec, + fd: i32, + sample_rate: u32, + channels: u16, + sample_format: String, + mix_gain: f32, +) -> PyResult { + let _ = sink_id; + engine.build_wav_sink( + py, + route_ids, + fd, + sample_rate, + channels, + sample_format, + mix_gain, + ) } fn append_microphone_info(list: &Bound<'_, PyList>, mic: MicInfo) -> PyResult<()> { @@ -2054,6 +2099,7 @@ fn _macloop(m: &Bound<'_, PyModule>) -> PyResult<()> { m.add_function(wrap_pyfunction!(microphone_access, m)?)?; m.add_function(wrap_pyfunction!(create_asr_sink, m)?)?; m.add_function(wrap_pyfunction!(create_wav_sink, m)?)?; + m.add_function(wrap_pyfunction!(create_wav_sink_with_config, m)?)?; Ok(()) } diff --git a/tests/conftest.py b/tests/conftest.py index 1aeca22..eba5dcf 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -277,7 +277,34 @@ def _fake_create_asr_sink( def _fake_create_wav_sink(engine, sink_id, route_ids, fd, mix_gain): - engine.calls.append(("create_wav_sink", sink_id, tuple(route_ids), fd, mix_gain)) + engine.calls.append( + ("create_wav_sink", sink_id, tuple(route_ids), fd, mix_gain) + ) + return _FakeWavSinkBackend() + + +def _fake_create_wav_sink_with_config( + engine, + sink_id, + route_ids, + fd, + sample_rate, + channels, + sample_format, + mix_gain, +): + engine.calls.append( + ( + "create_wav_sink_with_config", + sink_id, + tuple(route_ids), + fd, + sample_rate, + channels, + sample_format, + mix_gain, + ) + ) return _FakeWavSinkBackend() @@ -295,6 +322,7 @@ def macloop_module(monkeypatch): fake_ext.WavSinkStats = _FakeWavSinkStats fake_ext._create_asr_sink = _fake_create_asr_sink fake_ext._create_wav_sink = _fake_create_wav_sink + fake_ext._create_wav_sink_with_config = _fake_create_wav_sink_with_config fake_ext.list_microphones = lambda: [ {"id": 11, "name": "Built-in Mic", "is_default": True}, {"id": 22, "name": "USB Mic", "is_default": False}, diff --git a/tests/test_e2e_synthetic.py b/tests/test_e2e_synthetic.py index fa56345..8757b45 100644 --- a/tests/test_e2e_synthetic.py +++ b/tests/test_e2e_synthetic.py @@ -7,6 +7,7 @@ import sys import textwrap import threading +import time from pathlib import Path import pytest @@ -49,6 +50,52 @@ def _read_float_wav(path: Path) -> tuple[tuple[int, int, int], list[float]]: return (fmt_info[0], fmt_info[1], fmt_info[2]), samples +def _read_pcm16_wav(path: Path) -> tuple[tuple[int, int, int, int], list[int]]: + data = path.read_bytes() + assert data[:4] == b"RIFF" + assert data[8:12] == b"WAVE" + + fmt_info = None + samples = None + offset = 12 + while offset + 8 <= len(data): + chunk_id = data[offset : offset + 4] + size = struct.unpack_from(" None: + deadline = time.monotonic() + timeout + while True: + stats = engine.stats() + if all(stats[stream_id].pipeline.latency.count >= count for stream_id, count in expected.items()): + return + if time.monotonic() >= deadline: + pytest.fail(f"Synthetic streams did not produce every configured callback: {expected}") + time.sleep(0.01) + + def _collect_chunks(sink: macloop.AsrSink, count: int) -> list[macloop.AudioChunk]: collected: "queue.Queue[macloop.AudioChunk]" = queue.Queue() @@ -140,6 +187,60 @@ def test_synthetic_source_reaches_asr_and_wav(tmp_path: Path) -> None: ] +def test_explicit_pcm16_wav_and_asr_include_processed_gain(tmp_path: Path) -> None: + output_path = tmp_path / "synthetic_16khz_mono_pcm16.wav" + input_frames = 160 * 151 + expected_output_frames = (input_frames * 16_000 + 47_999) // 48_000 + expected_sample = 4096 + + with macloop.AudioEngine() as engine: + stream = engine.create_stream( + macloop.SyntheticSource, + frames_per_callback=160, + callback_count=151, + start_value=0.25, + step_value=0.0, + interval_ms=2, + start_delay_ms=200, + ) + engine.add_processor(stream=stream, processor=macloop.GainProcessor(gain=0.5)) + + asr_route = engine.route("explicit_asr", stream=stream) + wav_route = engine.route("explicit_wav", stream=stream) + asr_sink = macloop.AsrSink( + routes=[asr_route], + chunk_frames=320, + sample_rate=16_000, + channels=1, + sample_format="i16", + ) + wav_sink = macloop.WavSink( + route=wav_route, + file=output_path, + sample_rate=16_000, + channels=1, + sample_format="i16", + ) + + chunks = _collect_chunks(asr_sink, 3) + _wait_for_stream_callbacks(engine, {stream.id: 151}) + asr_sink.close() + wav_sink.close() + wav_stats = wav_sink.stats() + + assert chunks[-1].samples.dtype == np.int16 + assert np.all(np.abs(chunks[-1].samples.astype(np.int32) - expected_sample) <= 2) + + fmt_info, wav_samples = _read_pcm16_wav(output_path) + assert fmt_info == (1, 1, 16_000, 16) + assert len(wav_samples) == expected_output_frames + assert all( + abs(sample - expected_sample) <= 2 for sample in wav_samples[1000:-1000] + ) + assert wav_stats.samples_written == len(wav_samples) + assert wav_stats.frames_written == len(wav_samples) + + def test_synthetic_source_reaches_i16_asr_output() -> None: with macloop.AudioEngine() as engine: stream = engine.create_stream( @@ -234,8 +335,7 @@ def test_two_synthetic_sources_mix_into_aligned_wav(tmp_path: Path) -> None: wav_sink = macloop.WavSink(routes=[route_a, route_b], file=output_path) - # Let both synthetic sources finish and give the writer thread time to flush. - threading.Event().wait(0.5) + _wait_for_stream_callbacks(engine, {stream_a.id: 6, stream_b.id: 6}) wav_sink.close() fmt_info, wav_samples = _read_float_wav(output_path) diff --git a/tests/test_ffi_backend.py b/tests/test_ffi_backend.py index 13ceda3..2d9dd76 100644 --- a/tests/test_ffi_backend.py +++ b/tests/test_ffi_backend.py @@ -1,8 +1,10 @@ from __future__ import annotations +import os import queue import struct import threading +import time from pathlib import Path import pytest @@ -50,6 +52,18 @@ def _wait_for_chunks( return [collected.get(timeout=2.0) for _ in range(count)] +def _open_file_descriptors() -> set[int]: + descriptors = set() + for name in os.listdir("/dev/fd"): + try: + fd = int(name) + os.fstat(fd) + except (OSError, ValueError): + continue + descriptors.add(fd) + return descriptors + + def test_low_level_ffi_asr_sink_round_trip() -> None: engine = ffi._AudioEngineBackend() engine.create_stream( @@ -104,7 +118,9 @@ def callback(route_id: str, frames: int, samples) -> None: assert sink_stats.callback.count == 2 -def test_low_level_ffi_wav_sink_writes_output(tmp_path: Path) -> None: +def test_low_level_ffi_wav_sink_old_five_arg_api_writes_default_format( + tmp_path: Path, +) -> None: output_path = tmp_path / "ffi_synthetic.wav" engine = ffi._AudioEngineBackend() @@ -122,7 +138,13 @@ def test_low_level_ffi_wav_sink_writes_output(tmp_path: Path) -> None: engine.route("wav_route", "synthetic_stream") with output_path.open("w+b") as fileobj: - sink = ffi._create_wav_sink(engine, "ffi_wav_sink", ["wav_route"], fileobj.fileno(), 1.0) + sink = ffi._create_wav_sink( + engine, + "ffi_wav_sink", + ["wav_route"], + fileobj.fileno(), + 1.0, + ) try: threading.Event().wait(0.5) stats = sink.stats() @@ -139,7 +161,150 @@ def test_low_level_ffi_wav_sink_writes_output(tmp_path: Path) -> None: assert stats.frames_written == 12 -def test_low_level_ffi_validates_invalid_configs() -> None: +def test_low_level_ffi_wav_sink_converts_to_mono_pcm16(tmp_path: Path) -> None: + output_path = tmp_path / "ffi_synthetic_pcm16.wav" + + engine = ffi._AudioEngineBackend() + engine.create_stream( + "synthetic_stream", + "synthetic", + { + "frames_per_callback": 160, + "callback_count": 60, + "start_value": 0.25, + "step_value": 0.0, + "interval_ms": 3, + "start_delay_ms": 100, + }, + ) + engine.route("wav_route", "synthetic_stream") + + with output_path.open("w+b") as fileobj: + sink = ffi._create_wav_sink_with_config( + engine, + "ffi_wav_sink_pcm16", + ["wav_route"], + fileobj.fileno(), + 16_000, + 1, + "i16", + 1.0, + ) + try: + deadline = time.monotonic() + 5.0 + while engine.get_stats()["synthetic_stream"].pipeline.latency.count < 60: + if time.monotonic() >= deadline: + pytest.fail("Synthetic source did not produce every configured callback") + time.sleep(0.01) + finally: + sink.close() + stats = sink.stats() + engine.close() + + data = output_path.read_bytes() + fmt_offset = data.index(b"fmt ") + 8 + audio_format, channels, sample_rate, _, _, bits_per_sample = struct.unpack_from( + " None: + read_only_path = tmp_path / "read_only.wav" + read_only_path.touch() + recovered_path = tmp_path / "recovered.wav" + engine = ffi._AudioEngineBackend() + engine.create_stream( + "synthetic_stream", + "synthetic", + { + "frames_per_callback": 4, + "callback_count": 1, + "start_value": 0.25, + "start_delay_ms": 100, + }, + ) + engine.route("wav_route", "synthetic_stream") + + try: + with read_only_path.open("rb") as fileobj: + descriptors_before = _open_file_descriptors() + sink = ffi._create_wav_sink( + engine, + "failing_wav_sink", + ["wav_route"], + fileobj.fileno(), + 1.0, + ) + with pytest.raises( + RuntimeError, + match="failed to stop wav sink: wav writer error", + ): + sink.close(engine) + + stats = sink.stats() + assert stats.finalize.count == 1 + assert _open_file_descriptors() == descriptors_before + + with recovered_path.open("w+b") as fileobj: + replacement = ffi._create_wav_sink( + engine, + "replacement_wav_sink", + ["wav_route"], + fileobj.fileno(), + 1.0, + ) + replacement.close(engine) + finally: + engine.close() + + +def test_wav_close_combines_write_and_restoration_errors(tmp_path: Path) -> None: + read_only_path = tmp_path / "read_only_combined.wav" + read_only_path.touch() + engine = ffi._AudioEngineBackend() + wrong_engine = ffi._AudioEngineBackend() + engine.create_stream("synthetic_stream", "synthetic", {"start_delay_ms": 100}) + engine.route("wav_route", "synthetic_stream") + wrong_engine.create_stream("other_stream", "synthetic", {}) + wrong_engine.route("wav_route", "other_stream") + + try: + with read_only_path.open("rb") as fileobj: + sink = ffi._create_wav_sink( + engine, + "failing_wav_sink", + ["wav_route"], + fileobj.fileno(), + 1.0, + ) + with pytest.raises(RuntimeError) as exc_info: + sink.close(wrong_engine) + + message = str(exc_info.value) + assert "failed to stop wav sink: wav writer error" in message + assert "additionally failed to restore wav sink routes" in message + assert sink.stats().finalize.count == 1 + finally: + engine.close() + wrong_engine.close() + + +def test_low_level_ffi_validates_invalid_configs(tmp_path: Path) -> None: engine = ffi._AudioEngineBackend() with pytest.raises(ValueError, match="unsupported source_kind"): @@ -167,4 +332,42 @@ def test_low_level_ffi_validates_invalid_configs() -> None: lambda *_args: None, ) + output_path = tmp_path / "invalid.wav" + with output_path.open("w+b") as fileobj: + with pytest.raises(ValueError, match="unsupported sample_format"): + ffi._create_wav_sink_with_config( + engine, + "bad_wav_format", + ["synthetic_route"], + fileobj.fileno(), + 16_000, + 1, + "u8", + 1.0, + ) + + with pytest.raises(ValueError, match="unsupported output channels"): + ffi._create_wav_sink_with_config( + engine, + "bad_wav_channels", + ["synthetic_route"], + fileobj.fileno(), + 16_000, + 3, + "i16", + 1.0, + ) + + sink = ffi._create_wav_sink_with_config( + engine, + "valid_wav_after_errors", + ["synthetic_route"], + fileobj.fileno(), + 48_000, + 2, + "f32", + 1.0, + ) + sink.close(engine) + engine.close() diff --git a/tests/test_public_api.py b/tests/test_public_api.py index 32b0c9e..bf7ac58 100644 --- a/tests/test_public_api.py +++ b/tests/test_public_api.py @@ -284,7 +284,16 @@ def test_wav_sink_accepts_file_objects_and_engine_closes_sinks(macloop_module) - with tempfile.TemporaryFile() as fileobj: sink = macloop_module.WavSink(route=route, file=fileobj) assert sink.id.startswith("wav_sink_") - assert engine._backend.calls[-1] == ("create_wav_sink", sink.id, (route.id,), fileobj.fileno(), 1.0) + assert engine._backend.calls[-1] == ( + "create_wav_sink_with_config", + sink.id, + (route.id,), + fileobj.fileno(), + 48_000, + 2, + "f32", + 1.0, + ) assert sink._backend.closed is False stats = sink.stats() assert stats.write_calls == 3 @@ -310,10 +319,13 @@ def test_wav_sink_accepts_multiple_routes(macloop_module) -> None: sink = macloop_module.WavSink(routes=[route_a, route_b], file=fileobj) assert engine._backend.calls[-1] == ( - "create_wav_sink", + "create_wav_sink_with_config", sink.id, (route_a.id, route_b.id), fileobj.fileno(), + 48_000, + 2, + "f32", 0.5, ) @@ -325,10 +337,64 @@ def test_wav_sink_accepts_multiple_routes(macloop_module) -> None: with tempfile.TemporaryFile() as fileobj: second = macloop_module.WavSink(routes=[route_a, route_b], file=fileobj, mix_gain=0.25) assert engine._backend.calls[-1] == ( - "create_wav_sink", + "create_wav_sink_with_config", second.id, (route_a.id, route_b.id), fileobj.fileno(), + 48_000, + 2, + "f32", 0.25, ) second.close() + + +def test_wav_sink_passes_explicit_output_format(macloop_module) -> None: + with macloop_module.AudioEngine() as engine: + stream = engine.create_stream(macloop_module.MicrophoneSource, "mic") + route = engine.route("wav_route", stream=stream) + + with tempfile.TemporaryFile() as fileobj: + sink = macloop_module.WavSink( + route=route, + file=fileobj, + sample_rate=16_000, + channels=1, + sample_format="i16", + mix_gain=1.0, + ) + + assert engine._backend.calls[-1] == ( + "create_wav_sink_with_config", + sink.id, + (route.id,), + fileobj.fileno(), + 16_000, + 1, + "i16", + 1.0, + ) + sink.close() + + +@pytest.mark.parametrize( + "format_kwargs", + [ + {"sample_rate": 16_000}, + {"channels": 1, "sample_format": "i16"}, + {"sample_rate": 16_000, "sample_format": "i16"}, + ], +) +def test_wav_sink_requires_complete_explicit_output_format( + macloop_module, format_kwargs +) -> None: + with macloop_module.AudioEngine() as engine: + stream = engine.create_stream(macloop_module.MicrophoneSource, "mic") + route = engine.route("wav_route", stream=stream) + + with tempfile.TemporaryFile() as fileobj: + with pytest.raises( + ValueError, + match="sample_rate, channels, and sample_format must be provided together", + ): + macloop_module.WavSink(route=route, file=fileobj, **format_kwargs)