use std::panic::{AssertUnwindSafe, catch_unwind};
use mediaway_common::{AudioFrame, SampleFormat};
use sonora::{AudioProcessing, StreamConfig};
use crate::apm::ApmConfig;
use crate::apm::error::ApmError;
use crate::apm::pcm::{bytes_to_f32, f32_to_bytes};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct AudioStreamFormat {
pub sample_rate: u32,
pub channels: u16,
pub sample_format: SampleFormat,
}
#[derive(Debug)]
pub struct AudioProcessor {
inner: Option<AudioProcessing>,
capture_format: AudioStreamFormat,
render_format: AudioStreamFormat,
capture_block_frames: usize,
render_block_frames: usize,
capture_channels: usize,
render_channels: usize,
capture_accum: Vec<f32>,
render_accum: Vec<f32>,
capture_in: Vec<Vec<f32>>,
capture_out: Vec<Vec<f32>>,
render_in: Vec<Vec<f32>>,
render_out: Vec<Vec<f32>>,
next_capture_pts: Option<i64>,
}
fn deinterleave(interleaved: &[f32], channels: usize, out: &mut [Vec<f32>]) {
if channels == 0 {
return;
}
let frames = interleaved.len() / channels;
for (ch, out_ch) in out.iter_mut().enumerate().take(channels) {
for (f, sample) in out_ch.iter_mut().enumerate().take(frames) {
*sample = interleaved[f * channels + ch];
}
}
}
fn interleave(channel_data: &[Vec<f32>], channels: usize, frames: usize) -> Vec<f32> {
let mut out = vec![0.0_f32; channels * frames];
for (ch, buf) in channel_data.iter().enumerate() {
for (f, &sample) in buf.iter().enumerate().take(frames) {
out[f * channels + ch] = sample;
}
}
out
}
fn validate_frame(frame: &AudioFrame, expected: AudioStreamFormat) -> Result<(), ApmError> {
if frame.format != SampleFormat::F32 {
return Err(ApmError::UnsupportedSampleFormat(frame.format));
}
if frame.sample_rate != expected.sample_rate || frame.channels != expected.channels {
return Err(ApmError::StreamFormatMismatch {
expected_sample_rate: expected.sample_rate,
expected_channels: expected.channels,
actual_sample_rate: frame.sample_rate,
actual_channels: frame.channels,
});
}
Ok(())
}
impl AudioProcessor {
pub fn open(
config: ApmConfig,
capture_format: AudioStreamFormat,
render_format: AudioStreamFormat,
) -> Result<Self, ApmError> {
if capture_format.sample_format != SampleFormat::F32 {
return Err(ApmError::UnsupportedSampleFormat(
capture_format.sample_format,
));
}
if render_format.sample_format != SampleFormat::F32 {
return Err(ApmError::UnsupportedSampleFormat(
render_format.sample_format,
));
}
let capture_stream = StreamConfig::new(capture_format.sample_rate, capture_format.channels);
let render_stream = StreamConfig::new(render_format.sample_rate, render_format.channels);
let build_result = catch_unwind(AssertUnwindSafe(|| {
AudioProcessing::builder()
.config(config)
.capture_config(capture_stream)
.render_config(render_stream)
.build()
}));
let inner = match build_result {
Ok(apm) => Some(apm),
Err(_) => return Err(ApmError::BackendPanicked),
};
let capture_channels = usize::from(capture_format.channels);
let render_channels = usize::from(render_format.channels);
let capture_block_frames = capture_stream.num_frames();
let render_block_frames = render_stream.num_frames();
Ok(Self {
inner,
capture_format,
render_format,
capture_block_frames,
render_block_frames,
capture_channels,
render_channels,
capture_accum: Vec::new(),
render_accum: Vec::new(),
capture_in: vec![vec![0.0; capture_block_frames]; capture_channels],
capture_out: vec![vec![0.0; capture_block_frames]; capture_channels],
render_in: vec![vec![0.0; render_block_frames]; render_channels],
render_out: vec![vec![0.0; render_block_frames]; render_channels],
next_capture_pts: None,
})
}
pub fn push_render_frame(&mut self, frame: &AudioFrame) -> Result<(), ApmError> {
if self.inner.is_none() {
return Ok(());
}
validate_frame(frame, self.render_format)?;
self.render_accum.extend(bytes_to_f32(&frame.data));
let block_len = self.render_block_frames * self.render_channels;
while block_len > 0 && self.render_accum.len() >= block_len {
deinterleave(
&self.render_accum[..block_len],
self.render_channels,
&mut self.render_in,
);
self.render_accum.drain(..block_len);
let Some(inner) = self.inner.as_mut() else {
return Ok(());
};
let src: Vec<&[f32]> = self.render_in.iter().map(Vec::as_slice).collect();
let mut dst: Vec<&mut [f32]> =
self.render_out.iter_mut().map(Vec::as_mut_slice).collect();
let result = catch_unwind(AssertUnwindSafe(|| {
inner.process_render_f32(&src, &mut dst)
}));
match result {
Ok(Ok(())) => {}
Ok(Err(err)) => return Err(ApmError::Backend(err)),
Err(_) => {
self.inner = None;
return Err(ApmError::BackendPanicked);
}
}
}
Ok(())
}
pub fn push_capture_frame(&mut self, frame: &AudioFrame) -> Result<(), ApmError> {
validate_frame(frame, self.capture_format)?;
if self.next_capture_pts.is_none() {
self.next_capture_pts = Some(frame.pts);
}
self.capture_accum.extend(bytes_to_f32(&frame.data));
Ok(())
}
pub fn poll_processed_frame(&mut self) -> Result<Option<AudioFrame>, ApmError> {
let block_len = self.capture_block_frames * self.capture_channels;
if block_len == 0 || self.capture_accum.len() < block_len {
return Ok(None);
}
let base_pts = self.next_capture_pts.unwrap_or(0);
let block_frames_i64 = i64::try_from(self.capture_block_frames).unwrap_or(i64::MAX);
let Some(inner) = self.inner.as_mut() else {
let block: Vec<f32> = self.capture_accum.drain(..block_len).collect();
self.next_capture_pts = Some(base_pts.saturating_add(block_frames_i64));
return Ok(Some(self.raw_capture_frame(base_pts, &block)));
};
deinterleave(
&self.capture_accum[..block_len],
self.capture_channels,
&mut self.capture_in,
);
let src: Vec<&[f32]> = self.capture_in.iter().map(Vec::as_slice).collect();
let mut dst: Vec<&mut [f32]> = self.capture_out.iter_mut().map(Vec::as_mut_slice).collect();
let result = catch_unwind(AssertUnwindSafe(|| {
inner.process_capture_f32(&src, &mut dst)
}));
match result {
Ok(Ok(())) => {
self.capture_accum.drain(..block_len);
let interleaved = interleave(
&self.capture_out,
self.capture_channels,
self.capture_block_frames,
);
self.next_capture_pts = Some(base_pts.saturating_add(block_frames_i64));
Ok(Some(self.raw_capture_frame(base_pts, &interleaved)))
}
Ok(Err(err)) => Err(ApmError::Backend(err)),
Err(_) => {
self.inner = None;
Err(ApmError::BackendPanicked)
}
}
}
fn raw_capture_frame(&self, pts: i64, interleaved: &[f32]) -> AudioFrame {
AudioFrame {
pts,
duration: u64::try_from(self.capture_block_frames).unwrap_or(u64::MAX),
sample_rate: self.capture_format.sample_rate,
channels: self.capture_format.channels,
format: SampleFormat::F32,
data: f32_to_bytes(interleaved),
}
}
pub fn set_stream_delay_ms(&mut self, ms: i32) {
if let Some(inner) = self.inner.as_mut() {
let _ = inner.set_stream_delay_ms(ms);
}
}
#[must_use]
pub const fn is_disabled(&self) -> bool {
self.inner.is_none()
}
}
#[cfg(test)]
#[path = "processor_tests.rs"]
mod tests;