use std::f32::consts::FRAC_1_SQRT_2;
use rubato::audioadapter_buffers::direct::SequentialSliceOfVecs;
use rubato::{
Async, FixedAsync, Resampler as RubatoTrait, SincInterpolationParameters, SincInterpolationType, WindowFunction,
};
use crate::layout::Speaker;
use crate::{Error, Layout};
#[derive(Debug, thiserror::Error)]
enum BackendError {
#[error(transparent)]
Construction(#[from] rubato::ResamplerConstructionError),
#[error(transparent)]
Process(#[from] rubato::ResampleError),
}
impl BackendError {
fn into_public(self) -> Error {
match self {
Self::Construction(err) => Error::ResamplerConstruction(err.to_string()),
Self::Process(err) => Error::Resample(err.to_string()),
}
}
}
pub struct Resampler {
resampler: Async<f32>,
chunk_frames: usize,
input_rate: u32,
ratio: f64,
delay: usize,
started: bool,
skip: usize,
channels: usize,
input_planar: Vec<Vec<f32>>,
output_planar: Vec<Vec<f32>>,
output_frames_max: usize,
pending: Vec<f32>,
held: Option<moq_net::Timestamp>,
}
impl Resampler {
pub fn new(input_rate: u32, output_rate: u32, channels: u32, chunk_frames: usize) -> Result<Self, Error> {
if chunk_frames == 0 {
return Err(Error::Unsupported("chunk_frames must be > 0".into()));
}
Self::new_inner(input_rate, output_rate, channels, chunk_frames).map_err(BackendError::into_public)
}
fn new_inner(input_rate: u32, output_rate: u32, channels: u32, chunk_frames: usize) -> Result<Self, BackendError> {
let params = SincInterpolationParameters {
sinc_len: 128,
f_cutoff: Some(0.95),
interpolation: SincInterpolationType::Linear,
oversampling_factor: 128,
window: WindowFunction::BlackmanHarris2,
};
let ratio = output_rate as f64 / input_rate as f64;
let resampler =
Async::<f32>::new_sinc(ratio, 1.0, ¶ms, chunk_frames, channels as usize, FixedAsync::Input)?;
let delay = resampler.output_delay();
let input_planar = (0..channels as usize).map(|_| vec![0.0f32; chunk_frames]).collect();
let output_frames_max = resampler.output_frames_max();
let output_planar = vec![vec![0.0f32; output_frames_max]; channels as usize];
Ok(Self {
resampler,
chunk_frames,
input_rate,
ratio,
delay,
started: false,
skip: delay,
channels: channels as usize,
input_planar,
output_planar,
output_frames_max,
pending: Vec::new(),
held: None,
})
}
pub fn skipped(&self) -> usize {
self.delay - self.skip
}
pub fn pending_frames(&self) -> usize {
self.pending.len() / self.channels
}
pub(crate) fn held_at(&self) -> Option<moq_net::Timestamp> {
self.held
}
pub fn reset(&mut self) {
self.resampler.reset();
self.pending.clear();
self.skip = self.delay;
self.started = false;
self.held = None;
}
pub fn flush(mut self) -> Result<Vec<f32>, Error> {
self.drain()
}
pub fn drain(&mut self) -> Result<Vec<f32>, Error> {
let out = self.drained()?;
self.reset();
Ok(out)
}
fn drained(&mut self) -> Result<Vec<f32>, Error> {
if !self.started {
return Ok(Vec::new());
}
let pending = self.pending_frames();
let repaid = self.delay - self.skip;
let wanted = ((pending as f64 * self.ratio).round() as usize + repaid) * self.channels;
let mut out = Vec::new();
while out.len() < wanted {
let skip_before = self.skip;
self.pending.resize(self.chunk_frames * self.channels, 0.0);
let produced = self.convert().map_err(BackendError::into_public)?;
if produced.is_empty() && self.skip == skip_before {
break;
}
out.extend_from_slice(&produced);
}
out.truncate(wanted);
Ok(out)
}
pub fn process(&mut self, samples: &[f32], at: moq_net::Timestamp) -> Result<Vec<f32>, Error> {
if !samples.len().is_multiple_of(self.channels) {
return Err(Error::Misaligned {
got: samples.len(),
expected: samples.len().next_multiple_of(self.channels),
});
}
if self.pending.is_empty() {
self.held = Some(at);
}
self.started |= !samples.is_empty();
self.pending.extend_from_slice(samples);
let buffered = self.pending.len();
let out = self.convert().map_err(BackendError::into_public)?;
if self.pending.len() < buffered {
let consumed = (samples.len() - self.pending.len()) / self.channels;
let elapsed =
moq_net::Timestamp::from_scale(consumed as u64, self.input_rate as u64)?.convert(at.scale())?;
self.held = Some(at.checked_add(elapsed)?);
}
Ok(out)
}
fn convert(&mut self) -> Result<Vec<f32>, BackendError> {
let chunk_samples = self.chunk_frames * self.channels;
let mut out = Vec::new();
while self.pending.len() >= chunk_samples {
for (frame_idx, frame) in self.pending[..chunk_samples].chunks_exact(self.channels).enumerate() {
for (ch, &sample) in frame.iter().enumerate() {
self.input_planar[ch][frame_idx] = sample;
}
}
let input = SequentialSliceOfVecs::new(&self.input_planar, self.channels, self.chunk_frames)
.expect("resampler input buffer dimensions");
let mut output =
SequentialSliceOfVecs::new_mut(&mut self.output_planar, self.channels, self.output_frames_max)
.expect("resampler output buffer dimensions");
let (_, produced) = self.resampler.process_into_buffer(&input, &mut output, None)?;
let prev_len = out.len();
out.resize(prev_len + produced * self.channels, 0.0);
for frame_idx in 0..produced {
for ch in 0..self.channels {
out[prev_len + frame_idx * self.channels + ch] = self.output_planar[ch][frame_idx];
}
}
self.pending.drain(..chunk_samples);
}
if self.skip > 0 {
let drop = self.skip.min(out.len() / self.channels) * self.channels;
out.drain(..drop);
self.skip -= drop / self.channels;
}
Ok(out)
}
}
pub(crate) struct Remix {
inputs: usize,
outputs: usize,
weights: Vec<f32>,
}
impl Remix {
pub(crate) fn new(input: Layout, output: Layout) -> Result<Self, Error> {
input.validate()?;
output.validate()?;
let (inputs, outputs) = (input.channels() as usize, output.channels() as usize);
let unchanged = input == output || output == Layout::Discrete(inputs as u32);
let weights = match (input.speakers(), output.speakers()) {
_ if unchanged => (0..outputs)
.flat_map(|o| (0..inputs).map(move |i| if i == o { 1.0 } else { 0.0 }))
.collect(),
(Some(from), Some(to)) => weights(from, to),
_ => {
return Err(Error::Unsupported(format!(
"cannot convert audio layout {input:?} to {output:?} without speaker positions"
)));
}
};
Ok(Self {
inputs,
outputs,
weights,
})
}
pub(crate) fn apply(&self, input: &[f32], output: &mut [f32]) {
for (frame, out) in input
.chunks_exact(self.inputs)
.zip(output.chunks_exact_mut(self.outputs))
{
for (sample, row) in out.iter_mut().zip(self.weights.chunks_exact(self.inputs)) {
*sample = row.iter().zip(frame).map(|(weight, input)| weight * input).sum();
}
}
}
pub(crate) fn process(&self, input: &[f32]) -> Vec<f32> {
let mut output = vec![0.0; input.len() / self.inputs * self.outputs];
self.apply(input, &mut output);
output
}
}
fn weights(input: &[Speaker], output: &[Speaker]) -> Vec<f32> {
use Speaker::*;
if output == [FrontCenter] && input != [FrontCenter] {
let stereo = weights(input, &[FrontLeft, FrontRight]);
let (left, right) = stereo.split_at(input.len());
return left.iter().zip(right).map(|(l, r)| (l + r) * 0.5).collect();
}
let mut weights = vec![0.0; output.len() * input.len()];
let has = |speaker| output.contains(&speaker);
for (i, &speaker) in input.iter().enumerate() {
let mut feed = |to: Speaker, weight: f32| {
if let Some(o) = output.iter().position(|s| *s == to) {
weights[o * input.len() + i] += weight;
}
};
if has(speaker) {
feed(speaker, 1.0);
continue;
}
match speaker {
FrontCenter => {
let weight = if input == [FrontCenter] { 1.0 } else { FRAC_1_SQRT_2 };
feed(FrontLeft, weight);
feed(FrontRight, weight);
}
Lfe => {}
SideLeft | BackLeft | SideRight | BackRight => {
let (front, side, back) = match speaker {
SideLeft | BackLeft => (FrontLeft, SideLeft, BackLeft),
_ => (FrontRight, SideRight, BackRight),
};
let other = if speaker == side { back } else { side };
if has(other) {
feed(other, 1.0);
} else {
feed(front, FRAC_1_SQRT_2);
}
}
BackCenter => {
if has(BackLeft) {
feed(BackLeft, FRAC_1_SQRT_2);
feed(BackRight, FRAC_1_SQRT_2);
} else if has(SideLeft) {
feed(SideLeft, FRAC_1_SQRT_2);
feed(SideRight, FRAC_1_SQRT_2);
} else {
feed(FrontLeft, 0.5);
feed(FrontRight, 0.5);
}
}
FrontLeft | FrontRight => unreachable!("{output:?} has no front pair"),
}
}
weights
}
#[cfg(test)]
mod tests {
use super::*;
fn remix(samples: &[f32], input: Layout, output: Layout) -> Result<Vec<f32>, Error> {
Ok(Remix::new(input, output)?.process(samples))
}
fn at(frames: u64, rate: u64) -> moq_net::Timestamp {
moq_net::Timestamp::from_scale(frames, rate).unwrap()
}
#[test]
fn rejects_zero_chunk_frames() {
let r = Resampler::new(48_000, 48_000, 2, 0);
assert!(matches!(r, Err(Error::Unsupported(_))));
}
#[test]
fn upsample_44100_to_48000_preserves_energy_roughly() {
let mut r = Resampler::new(44_100, 48_000, 1, 1024).unwrap();
let input: Vec<f32> = (0..44_100)
.map(|i| (2.0 * std::f32::consts::PI * 440.0 * i as f32 / 44_100.0).sin() * 0.5)
.collect();
let mut out = r.process(&input, at(0, 44_100)).unwrap();
out.extend(r.process(&vec![0.0; 1024], at(44_100, 44_100)).unwrap());
assert!(
(47_000..50_000).contains(&out.len()),
"expected ~48k samples, got {}",
out.len()
);
}
#[test]
fn flush_drains_the_delayed_tail() {
let mut r = Resampler::new(44_100, 48_000, 1, 882).unwrap();
let mut input = vec![0.0f32; 1024];
input[1000] = 1.0;
let body = r.process(&input, at(0, 44_100)).unwrap();
let tail = r.flush().unwrap();
let peak = |samples: &[f32]| samples.iter().fold(0.0f32, |max, s| max.max(s.abs()));
assert!(peak(&body) < 0.01, "the sample emerged early: peak {}", peak(&body));
assert!(peak(&tail) > 0.5, "the tail lost the sample: peak {}", peak(&tail));
}
#[test]
fn flush_drains_on_an_exact_chunk_boundary() {
let mut r = Resampler::new(44_100, 48_000, 1, 882).unwrap();
let mut input = vec![0.0f32; 1764];
input[1750] = 1.0;
let body = r.process(&input, at(0, 44_100)).unwrap();
assert_eq!(r.pending_frames(), 0, "the input should divide evenly");
let tail = r.flush().unwrap();
let peak = |samples: &[f32]| samples.iter().fold(0.0f32, |max, s| max.max(s.abs()));
assert!(peak(&body) < 0.01, "the sample emerged early: peak {}", peak(&body));
assert!(peak(&tail) > 0.5, "the tail lost the sample: peak {}", peak(&tail));
}
#[test]
fn flush_sizes_a_stream_shorter_than_a_chunk() {
let mut r = Resampler::new(44_100, 48_000, 1, 882).unwrap();
let body = r.process(&[0.25f32; 441], at(0, 44_100)).unwrap();
let tail = r.flush().unwrap();
let total = body.len() + tail.len();
assert!((475..=485).contains(&total), "unexpected total: {total}");
}
#[test]
fn flush_survives_a_chunk_smaller_than_the_delay() {
let mut r = Resampler::new(44_100, 48_000, 1, 32).unwrap();
let body = r.process(&[0.5f32; 20], at(0, 44_100)).unwrap();
let tail = r.flush().unwrap();
let total = body.len() + tail.len();
assert!((18..=26).contains(&total), "unexpected total: {total}");
assert!(
tail.iter().any(|s| s.abs() > 0.25),
"the stream came back silent: peak {}",
tail.iter().fold(0.0f32, |m, s| m.max(s.abs()))
);
}
#[test]
fn drain_ends_the_stream_and_starts_a_new_one() {
let mut r = Resampler::new(44_100, 48_000, 1, 882).unwrap();
let mut input = vec![0.0f32; 1024];
input[1000] = 1.0;
let body = r.process(&input, at(0, 44_100)).unwrap();
let tail = r.drain().unwrap();
let peak = |samples: &[f32]| samples.iter().fold(0.0f32, |max, s| max.max(s.abs()));
assert!(peak(&tail) > 0.5, "the tail lost the sample: peak {}", peak(&tail));
assert!(
(1105..=1120).contains(&(body.len() + tail.len())),
"unexpected total: {}",
body.len() + tail.len()
);
assert_eq!(r.pending_frames(), 0);
assert_eq!(r.held_at(), None, "the drain should forget where the old stream was");
let after = r.process(&vec![0.0f32; 1024], at(2048, 44_100)).unwrap();
assert!(peak(&after) < 0.01, "audio crossed the gap: peak {}", peak(&after));
}
#[test]
fn held_frames_keep_the_stamp_they_arrived_under() {
let mut r = Resampler::new(44_100, 48_000, 1, 882).unwrap();
assert!(r.process(&[0.25f32; 441], at(0, 44_100)).unwrap().is_empty());
assert_eq!(r.held_at(), Some(at(0, 44_100)));
assert!(!r.process(&[0.25f32; 441], at(44_100, 44_100)).unwrap().is_empty());
assert_eq!(
r.held_at(),
Some(at(44_541, 44_100)),
"the tail starts at the end of the last packet consumed"
);
}
#[test]
fn an_emptied_buffer_re_anchors_on_the_next_input() {
let mut r = Resampler::new(44_100, 48_000, 1, 882).unwrap();
r.process(&[0.25f32; 882], at(0, 44_100)).unwrap();
assert_eq!(r.pending_frames(), 0);
r.process(&[0.25f32; 441], at(44_100, 44_100)).unwrap();
assert_eq!(r.held_at(), Some(at(44_100, 44_100)));
}
#[test]
fn leftover_frames_keep_the_new_packet_timestamp() {
let mut r = Resampler::new(44_100, 48_000, 1, 882).unwrap();
r.process(&[0.25; 441], at(0, 44_100)).unwrap();
r.process(&[0.25; 882], at(44_100, 44_100)).unwrap();
assert_eq!(r.pending_frames(), 441);
assert_eq!(r.held_at(), Some(at(44_541, 44_100)));
}
#[test]
fn held_timestamp_preserves_fractional_chunk_progress() {
let mut r = Resampler::new(11_025, 48_000, 1, 220).unwrap();
r.process(&vec![0.25; 11_025], at(0, 1000)).unwrap();
assert_eq!(r.pending_frames(), 25);
assert_eq!(r.held_at(), Some(at(997, 1000)));
}
#[test]
fn remix_mono_to_stereo_duplicates_samples() {
assert_eq!(
remix(&[1.0, 2.0], Layout::Mono, Layout::Stereo).unwrap(),
[1.0, 1.0, 2.0, 2.0]
);
}
#[test]
fn remix_stereo_to_mono_averages_channels() {
assert_eq!(
remix(&[1.0, 3.0, 2.0, 4.0], Layout::Stereo, Layout::Mono).unwrap(),
[2.0, 3.0]
);
}
const FIVE_ONE: [f32; 6] = [0.1, 0.2, 0.3, 0.4, 0.5, 0.6];
fn close(got: &[f32], want: &[f32]) {
assert_eq!(got.len(), want.len(), "{got:?} vs {want:?}");
for (g, w) in got.iter().zip(want) {
assert!((g - w).abs() < 1e-6, "{got:?} vs {want:?}");
}
}
#[test]
fn remix_downmixes_five_one_to_stereo_by_bs775() {
let h = FRAC_1_SQRT_2;
let [l, r, c, _lfe, ls, rs] = FIVE_ONE;
close(
&remix(&FIVE_ONE, Layout::FivePointOne, Layout::Stereo).unwrap(),
&[l + h * c + h * ls, r + h * c + h * rs],
);
}
#[test]
fn remix_downmixes_five_one_to_mono_through_stereo() {
let stereo = remix(&FIVE_ONE, Layout::FivePointOne, Layout::Stereo).unwrap();
close(
&remix(&FIVE_ONE, Layout::FivePointOne, Layout::Mono).unwrap(),
&[(stereo[0] + stereo[1]) * 0.5],
);
}
#[test]
fn remix_upmixes_stereo_into_the_front_pair() {
close(
&remix(&[0.25, 0.75], Layout::Stereo, Layout::FivePointOne).unwrap(),
&[0.25, 0.75, 0.0, 0.0, 0.0, 0.0],
);
}
#[test]
fn remix_upmixes_mono_into_the_center() {
close(
&remix(&[0.5], Layout::Mono, Layout::FivePointOne).unwrap(),
&[0.0, 0.0, 0.5, 0.0, 0.0, 0.0],
);
close(
&remix(&[0.5], Layout::Mono, Layout::Quad).unwrap(),
&[0.5, 0.5, 0.0, 0.0],
);
}
#[test]
fn remix_folds_back_into_side_surrounds() {
let seven = [0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8];
close(
&remix(&seven, Layout::SevenPointOne, Layout::FivePointOne).unwrap(),
&[0.1, 0.2, 0.3, 0.4, 0.5 + 0.7, 0.6 + 0.8],
);
}
#[test]
fn remix_refuses_positions_it_would_invent() {
for (input, output) in [
(Layout::Discrete(6), Layout::Stereo),
(Layout::Mono, Layout::Discrete(2)),
(Layout::Discrete(0), Layout::Discrete(0)),
] {
assert!(
matches!(Remix::new(input, output), Err(Error::Unsupported(_))),
"{input:?} -> {output:?}"
);
}
assert_eq!(
remix(&[1.0, 2.0, 3.0], Layout::Discrete(3), Layout::Discrete(3)).unwrap(),
[1.0, 2.0, 3.0]
);
assert_eq!(
remix(&[1.0, 2.0, 3.0], Layout::TwoPointOne, Layout::Discrete(3)).unwrap(),
[1.0, 2.0, 3.0]
);
assert!(Remix::new(Layout::Discrete(3), Layout::TwoPointOne).is_err());
}
#[test]
fn remix_converts_between_every_named_layout() {
let layouts: Vec<Layout> = (1..=8)
.map(|n| Layout::from_channels(n).unwrap())
.chain([Layout::ThreePointZero, Layout::FourPointZero])
.collect();
for &input in &layouts {
for &output in &layouts {
let frame = vec![0.5; input.channels() as usize * 2];
let mixed = remix(&frame, input, output).unwrap();
assert_eq!(mixed.len(), output.channels() as usize * 2, "{input:?} -> {output:?}");
}
}
}
}