use core::f64::consts::PI;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ImdctError {
ZeroLength,
OutputLenMismatch {
got: usize,
want: usize,
},
}
impl core::fmt::Display for ImdctError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match *self {
ImdctError::ZeroLength => {
write!(f, "oxideav-opus: CELT §4.3.7 inverse MDCT requires N >= 1")
}
ImdctError::OutputLenMismatch { got, want } => write!(
f,
"oxideav-opus: CELT §4.3.7 inverse MDCT output length {got} \
!= required 2*N = {want}"
),
}
}
}
impl std::error::Error for ImdctError {}
pub fn imdct_into(spectrum: &[f64], out: &mut [f64]) -> Result<(), ImdctError> {
let n = spectrum.len();
if n == 0 {
return Err(ImdctError::ZeroLength);
}
let want = 2 * n;
if out.len() != want {
return Err(ImdctError::OutputLenMismatch {
got: out.len(),
want,
});
}
let nf = n as f64;
let scale = 1.0 / nf;
let half_n = nf / 2.0;
for (idx, slot) in out.iter_mut().enumerate() {
let n_term = idx as f64 + 0.5 + half_n;
let mut acc = 0.0_f64;
for (k, &xk) in spectrum.iter().enumerate() {
let phase = (PI / nf) * n_term * (k as f64 + 0.5);
acc += xk * phase.cos();
}
*slot = scale * acc;
}
Ok(())
}
pub fn imdct(spectrum: &[f64]) -> Result<Vec<f64>, ImdctError> {
let mut out = vec![0.0_f64; 2 * spectrum.len()];
imdct_into(spectrum, &mut out)?;
Ok(out)
}
pub fn mdct_forward(time: &[f64]) -> Result<Vec<f64>, ImdctError> {
if time.is_empty() {
return Err(ImdctError::ZeroLength);
}
if time.len() % 2 != 0 {
return Err(ImdctError::OutputLenMismatch {
got: time.len(),
want: time.len() + 1,
});
}
let n = time.len() / 2;
let nf = n as f64;
let half_n = nf / 2.0;
let mut out = vec![0.0_f64; n];
for (k, coeff) in out.iter_mut().enumerate() {
let mut acc = 0.0_f64;
for (idx, &xn) in time.iter().enumerate() {
let n_term = idx as f64 + 0.5 + half_n;
let phase = (PI / nf) * n_term * (k as f64 + 0.5);
acc += xn * phase.cos();
}
*coeff = acc;
}
Ok(out)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::celt_mdct_window::window_tap;
const EPS: f64 = 1e-9;
#[test]
fn empty_spectrum_rejected() {
assert_eq!(imdct(&[]), Err(ImdctError::ZeroLength));
let mut out = [];
assert_eq!(imdct_into(&[], &mut out), Err(ImdctError::ZeroLength));
}
#[test]
fn output_length_is_twice_n() {
let spec = [1.0, 2.0, 3.0, 4.0];
let y = imdct(&spec).unwrap();
assert_eq!(y.len(), 8);
}
#[test]
fn output_len_mismatch_rejected() {
let spec = [1.0, 2.0];
let mut out = [0.0; 3];
assert_eq!(
imdct_into(&spec, &mut out),
Err(ImdctError::OutputLenMismatch { got: 3, want: 4 })
);
}
#[test]
fn matches_direct_formula() {
let spec = [0.5, -1.0, 2.0, 0.25];
let n = spec.len();
let nf = n as f64;
let y = imdct(&spec).unwrap();
for (idx, &got) in y.iter().enumerate() {
let n_term = idx as f64 + 0.5 + nf / 2.0;
let mut want = 0.0;
for (k, &xk) in spec.iter().enumerate() {
want += xk * ((PI / nf) * n_term * (k as f64 + 0.5)).cos();
}
want /= nf;
assert!((got - want).abs() < EPS, "n={idx}: {got} != {want}");
}
}
#[test]
fn imdct_into_matches_imdct() {
let spec = [1.0, -2.0, 3.0, -4.0, 5.0, -6.0, 7.0, -8.0];
let owned = imdct(&spec).unwrap();
let mut buf = vec![0.0; 2 * spec.len()];
imdct_into(&spec, &mut buf).unwrap();
assert_eq!(owned, buf);
}
#[test]
fn linearity() {
let x = [1.0, 2.0, -1.0, 0.5];
let yv = [0.3, -0.7, 2.1, 1.0];
let a = 1.5;
let b = -2.0;
let comb: Vec<f64> = x
.iter()
.zip(yv.iter())
.map(|(p, q)| a * p + b * q)
.collect();
let lhs = imdct(&comb).unwrap();
let ix = imdct(&x).unwrap();
let iy = imdct(&yv).unwrap();
for i in 0..lhs.len() {
let rhs = a * ix[i] + b * iy[i];
assert!((lhs[i] - rhs).abs() < EPS, "i={i}");
}
}
#[test]
fn aliasing_symmetry_pattern() {
let spec = [0.7, -1.3, 2.2, 0.1, -0.9, 1.1, 0.4, -2.5];
let n = spec.len();
let y = imdct(&spec).unwrap();
for nn in 0..n / 2 {
assert!(
(y[nn] + y[n - 1 - nn]).abs() < EPS,
"lower fold n={nn}: {} vs {}",
y[nn],
y[n - 1 - nn]
);
}
for m in 0..n / 2 {
assert!(
(y[n + m] - y[2 * n - 1 - m]).abs() < EPS,
"upper fold m={m}: {} vs {}",
y[n + m],
y[2 * n - 1 - m]
);
}
}
#[test]
fn forward_then_inverse_block_alias() {
let n = 4;
let x: Vec<f64> = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
let coeffs = mdct_forward(&x).unwrap();
assert_eq!(coeffs.len(), n);
let y = imdct(&coeffs).unwrap();
assert_eq!(y.len(), 2 * n);
for nn in 0..n {
let want = (x[nn] - x[n - 1 - nn]) / 2.0;
assert!(
(y[nn] - want).abs() < EPS,
"lower n={nn}: {} != {}",
y[nn],
want
);
}
for m in 0..n {
let want = (x[n + m] + x[2 * n - 1 - m]) / 2.0;
assert!(
(y[n + m] - want).abs() < EPS,
"upper m={m}: {} != {}",
y[n + m],
want
);
}
}
fn pb_window(two_n: usize) -> Vec<f64> {
(0..two_n)
.map(|i| (PI / (two_n as f64) * (i as f64 + 0.5)).sin())
.collect()
}
#[test]
fn windowed_overlap_add_reconstructs_at_half_gain() {
let n = 8;
let two_n = 2 * n;
let win = pb_window(two_n);
let total = 3 * 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 analyse_synthesise = |start: usize| -> Vec<f64> {
let block: Vec<f64> = (0..two_n).map(|i| signal[start + i] * win[i]).collect();
let coeffs = mdct_forward(&block).unwrap();
let rec = imdct(&coeffs).unwrap();
rec.iter().zip(win.iter()).map(|(r, w)| r * w).collect()
};
let f0 = analyse_synthesise(0);
let f1 = analyse_synthesise(n);
for j in 0..n {
let global = n + j;
let recon = f0[n + j] + f1[j];
let want = 0.5 * signal[global];
assert!(
(recon - want).abs() < 1e-9,
"overlap j={j}: recon={recon} != 0.5*signal={want}"
);
}
}
#[test]
fn celt_low_overlap_window_is_power_complementary() {
let len = 16;
for nn in 0..len {
let a = window_tap(nn, len).unwrap();
let b = window_tap(len - 1 - nn, len).unwrap();
assert!((a * a + b * b - 1.0).abs() < EPS, "nn={nn}");
}
}
#[test]
fn dc_spectrum_shape() {
let mut spec = vec![0.0; 16];
spec[0] = 4.0;
let y = imdct(&spec).unwrap();
assert!(y.iter().all(|v| v.is_finite()));
let n = spec.len();
for nn in 0..n / 2 {
assert!((y[nn] + y[n - 1 - nn]).abs() < EPS);
}
}
#[test]
fn error_display_messages() {
assert!(ImdctError::ZeroLength.to_string().contains("N >= 1"));
assert!(ImdctError::OutputLenMismatch { got: 3, want: 4 }
.to_string()
.contains("!= required 2*N"));
}
#[test]
fn mdct_forward_rejects_odd_and_empty() {
assert_eq!(mdct_forward(&[]), Err(ImdctError::ZeroLength));
assert_eq!(
mdct_forward(&[1.0, 2.0, 3.0]),
Err(ImdctError::OutputLenMismatch { got: 3, want: 4 })
);
}
}