use crate::celt_mdct_window::{mdct_window, MdctWindowError};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum OverlapAddError {
ZeroLength,
BlockLenNotEven {
got: usize,
},
BlockLenMismatch {
got: usize,
want: usize,
},
BadOverlap {
overlap: usize,
n: usize,
},
OutputTooSmall {
want: usize,
got: usize,
},
}
impl core::fmt::Display for OverlapAddError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match *self {
OverlapAddError::ZeroLength => {
write!(f, "oxideav-opus: CELT §4.3.7 overlap-add requires N >= 1")
}
OverlapAddError::BlockLenNotEven { got } => write!(
f,
"oxideav-opus: CELT §4.3.7 overlap-add block length {got} is not 2*N (odd)"
),
OverlapAddError::BlockLenMismatch { got, want } => write!(
f,
"oxideav-opus: CELT §4.3.7 overlap-add block length {got} != required 2*N = {want}"
),
OverlapAddError::BadOverlap { overlap, n } => write!(
f,
"oxideav-opus: CELT §4.3.7 overlap-add overlap {overlap} invalid for N={n} \
(require 0 < overlap <= N and overlap even)"
),
OverlapAddError::OutputTooSmall { want, got } => write!(
f,
"oxideav-opus: CELT §4.3.7 overlap-add output length {got} < required N = {want}"
),
}
}
}
impl std::error::Error for OverlapAddError {}
pub fn apply_synthesis_window(
block: &mut [f64],
ramp: &[f64],
n: usize,
) -> Result<(), OverlapAddError> {
if n == 0 {
return Err(OverlapAddError::ZeroLength);
}
if block.len() != 2 * n {
if block.len() % 2 != 0 {
return Err(OverlapAddError::BlockLenNotEven { got: block.len() });
}
return Err(OverlapAddError::BlockLenMismatch {
got: block.len(),
want: 2 * n,
});
}
let overlap = ramp.len();
if overlap == 0 || overlap > n || overlap % 2 != 0 {
return Err(OverlapAddError::BadOverlap { overlap, n });
}
let two_n = 2 * n;
for i in 0..overlap {
block[i] *= ramp[i];
block[two_n - 1 - i] *= ramp[i];
}
Ok(())
}
#[derive(Debug, Clone, PartialEq)]
pub struct WeightedOverlapAdd {
n: usize,
ramp: Vec<f64>,
history: Vec<f64>,
}
impl WeightedOverlapAdd {
pub fn new(n: usize, overlap: usize) -> Result<Self, OverlapAddError> {
if n == 0 {
return Err(OverlapAddError::ZeroLength);
}
if overlap == 0 || overlap > n || overlap % 2 != 0 {
return Err(OverlapAddError::BadOverlap { overlap, n });
}
let ramp = mdct_window(overlap).map_err(|e| match e {
MdctWindowError::ZeroLength => OverlapAddError::BadOverlap { overlap, n },
MdctWindowError::OddOverlap { overlap } => OverlapAddError::BadOverlap { overlap, n },
MdctWindowError::PositionOutOfRange { .. } => {
OverlapAddError::BadOverlap { overlap, n }
}
})?;
Ok(Self {
n,
ramp,
history: vec![0.0_f64; n],
})
}
#[must_use]
pub fn frame_len(&self) -> usize {
self.n
}
#[must_use]
pub fn overlap(&self) -> usize {
self.ramp.len()
}
#[must_use]
pub fn history(&self) -> &[f64] {
&self.history
}
pub fn reset(&mut self) {
for h in self.history.iter_mut() {
*h = 0.0;
}
}
pub fn process(&mut self, block: &[f64]) -> Result<Vec<f64>, OverlapAddError> {
let mut out = vec![0.0_f64; self.n];
self.process_into(block, &mut out)?;
Ok(out)
}
pub fn process_into(&mut self, block: &[f64], out: &mut [f64]) -> Result<(), OverlapAddError> {
let two_n = 2 * self.n;
if block.len() != two_n {
if block.len() % 2 != 0 {
return Err(OverlapAddError::BlockLenNotEven { got: block.len() });
}
return Err(OverlapAddError::BlockLenMismatch {
got: block.len(),
want: two_n,
});
}
if out.len() < self.n {
return Err(OverlapAddError::OutputTooSmall {
want: self.n,
got: out.len(),
});
}
let mut windowed = block.to_vec();
apply_synthesis_window(&mut windowed, &self.ramp, self.n)?;
for i in 0..self.n {
out[i] = windowed[i] + self.history[i];
}
self.history.copy_from_slice(&windowed[self.n..two_n]);
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::celt_imdct::{imdct, mdct_forward};
use crate::celt_mdct_window::window_tap;
const EPS: f64 = 1e-9;
#[test]
fn rejects_zero_n() {
assert_eq!(
WeightedOverlapAdd::new(0, 2),
Err(OverlapAddError::ZeroLength)
);
}
#[test]
fn rejects_bad_overlap() {
assert_eq!(
WeightedOverlapAdd::new(4, 6),
Err(OverlapAddError::BadOverlap { overlap: 6, n: 4 })
);
assert_eq!(
WeightedOverlapAdd::new(4, 0),
Err(OverlapAddError::BadOverlap { overlap: 0, n: 4 })
);
assert_eq!(
WeightedOverlapAdd::new(8, 3),
Err(OverlapAddError::BadOverlap { overlap: 3, n: 8 })
);
}
#[test]
fn overlap_equal_to_n_is_allowed() {
let ola = WeightedOverlapAdd::new(4, 4).unwrap();
assert_eq!(ola.frame_len(), 4);
assert_eq!(ola.overlap(), 4);
}
#[test]
fn rejects_wrong_block_length() {
let mut ola = WeightedOverlapAdd::new(4, 2).unwrap();
assert_eq!(
ola.process(&[0.0; 7]),
Err(OverlapAddError::BlockLenNotEven { got: 7 })
);
assert_eq!(
ola.process(&[0.0; 6]),
Err(OverlapAddError::BlockLenMismatch { got: 6, want: 8 })
);
}
#[test]
fn rejects_output_too_small() {
let mut ola = WeightedOverlapAdd::new(4, 2).unwrap();
let block = [0.0; 8];
let mut out = [0.0; 3];
assert_eq!(
ola.process_into(&block, &mut out),
Err(OverlapAddError::OutputTooSmall { want: 4, got: 3 })
);
}
#[test]
fn silence_stays_silent() {
let mut ola = WeightedOverlapAdd::new(8, 4).unwrap();
for _ in 0..3 {
let out = ola.process(&[0.0; 16]).unwrap();
assert_eq!(out, vec![0.0; 8]);
}
}
#[test]
fn first_frame_emits_windowed_leading_half() {
let n = 4;
let overlap = 2;
let mut ola = WeightedOverlapAdd::new(n, overlap).unwrap();
let block: Vec<f64> = (0..2 * n).map(|i| (i as f64 + 1.0) * 0.5).collect();
let out = ola.process(&block).unwrap();
let ramp = mdct_window(overlap).unwrap();
for i in 0..overlap {
assert!((out[i] - block[i] * ramp[i]).abs() < EPS, "i={i}");
}
for i in overlap..n {
assert!((out[i] - block[i]).abs() < EPS, "i={i}");
}
}
#[test]
fn history_holds_windowed_trailing_half() {
let n = 4;
let overlap = 2;
let mut ola = WeightedOverlapAdd::new(n, overlap).unwrap();
assert_eq!(ola.history(), &[0.0; 4]);
let block: Vec<f64> = (0..2 * n).map(|i| i as f64 + 1.0).collect();
ola.process(&block).unwrap();
let ramp = mdct_window(overlap).unwrap();
let mut windowed = block.clone();
apply_synthesis_window(&mut windowed, &ramp, n).unwrap();
assert_eq!(ola.history(), &windowed[n..2 * n]);
}
#[test]
fn reset_zeroes_history() {
let mut ola = WeightedOverlapAdd::new(4, 2).unwrap();
ola.process(&[1.0; 8]).unwrap();
assert!(ola.history().iter().any(|&h| h != 0.0));
ola.reset();
assert_eq!(ola.history(), &[0.0; 4]);
}
#[test]
fn apply_synthesis_window_layout() {
let n = 8;
let overlap = 4;
let ramp = mdct_window(overlap).unwrap();
let block: Vec<f64> = (0..2 * n).map(|i| i as f64 + 1.0).collect();
let mut w = block.clone();
apply_synthesis_window(&mut w, &ramp, n).unwrap();
let two_n = 2 * n;
for i in 0..overlap {
assert!((w[i] - block[i] * ramp[i]).abs() < EPS, "lead i={i}");
assert!(
(w[two_n - 1 - i] - block[two_n - 1 - i] * ramp[i]).abs() < EPS,
"trail i={i}"
);
}
for i in overlap..(two_n - overlap) {
assert!((w[i] - block[i]).abs() < EPS, "middle i={i}");
}
}
#[test]
fn apply_synthesis_window_rejects_bad_args() {
let ramp = mdct_window(2).unwrap();
assert_eq!(
apply_synthesis_window(&mut [0.0; 8], &ramp, 0),
Err(OverlapAddError::ZeroLength)
);
assert_eq!(
apply_synthesis_window(&mut [0.0; 6], &ramp, 4),
Err(OverlapAddError::BlockLenMismatch { got: 6, want: 8 })
);
assert!(matches!(
apply_synthesis_window(&mut [0.0; 7], &ramp, 4),
Err(OverlapAddError::BlockLenNotEven { got: 7 })
));
let big = mdct_window(6).unwrap();
assert_eq!(
apply_synthesis_window(&mut [0.0; 4], &big, 2),
Err(OverlapAddError::BadOverlap { overlap: 6, n: 2 })
);
}
fn pb_window(two_n: usize) -> Vec<f64> {
(0..two_n)
.map(|i| (core::f64::consts::PI / (two_n as f64) * (i as f64 + 0.5)).sin())
.collect()
}
#[test]
fn multi_frame_arithmetic_matches_reference() {
let n = 6;
let overlap = 4;
let mut ola = WeightedOverlapAdd::new(n, overlap).unwrap();
let ramp = mdct_window(overlap).unwrap();
assert!((ramp[0] - window_tap(0, overlap).unwrap()).abs() < EPS);
let blocks: Vec<Vec<f64>> = (0..3)
.map(|b| {
(0..2 * n)
.map(|i| ((b * n + i) as f64 * 0.21).sin() - 0.3)
.collect()
})
.collect();
let mut prev_tail = vec![0.0_f64; n];
for block in &blocks {
let mut w = block.clone();
apply_synthesis_window(&mut w, &ramp, n).unwrap();
let want: Vec<f64> = (0..n).map(|i| w[i] + prev_tail[i]).collect();
let got = ola.process(block).unwrap();
for i in 0..n {
assert!(
(got[i] - want[i]).abs() < EPS,
"i={i}: {} != {}",
got[i],
want[i]
);
}
prev_tail.copy_from_slice(&w[n..2 * n]);
}
}
#[test]
fn end_to_end_tdac_via_imdct_full_overlap() {
let n = 8;
let two_n = 2 * n;
let win = pb_window(two_n);
let total = 4 * n;
let signal: Vec<f64> = (0..total)
.map(|i| (i as f64 * 0.37).sin() * 1.5 - (i as f64 * 0.11).cos())
.collect();
let synth = |start: usize| -> Vec<f64> {
let blk: Vec<f64> = (0..two_n).map(|i| signal[start + i] * win[i]).collect();
let coeffs = mdct_forward(&blk).unwrap();
let rec = imdct(&coeffs).unwrap();
rec.iter().zip(win.iter()).map(|(r, w)| r * w).collect()
};
let f0 = synth(0);
let f1 = synth(n);
for j in 0..n {
let recon = f0[n + j] + f1[j];
let want = 0.5 * signal[n + j];
assert!(
(recon - want).abs() < 1e-9,
"overlap j={j}: {recon} != 0.5*signal {want}"
);
}
}
#[test]
fn error_display_messages() {
assert!(OverlapAddError::ZeroLength.to_string().contains("N >= 1"));
assert!(OverlapAddError::BlockLenNotEven { got: 7 }
.to_string()
.contains("not 2*N"));
assert!(OverlapAddError::BlockLenMismatch { got: 6, want: 8 }
.to_string()
.contains("!= required 2*N"));
assert!(OverlapAddError::BadOverlap { overlap: 6, n: 4 }
.to_string()
.contains("invalid"));
assert!(OverlapAddError::OutputTooSmall { want: 4, got: 3 }
.to_string()
.contains("< required N"));
}
}