use crate::celt_pvq_decode::{decode_pvq_shape_into, PvqShapeError};
use crate::celt_spreading::{apply_spreading, SpreadingError};
use crate::celt_tf_adjust::{TfAdjustment, TfDirection};
use crate::celt_tf_hadamard::{apply_tf_hadamard, TfHadamardError};
pub const PVQ_MAX_CODEBOOK_BITS: u32 = 32;
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum BandShapeError {
Pvq(PvqShapeError),
Spreading(SpreadingError),
TfHadamard(TfHadamardError),
OutputBufferTooSmall {
required: usize,
provided: usize,
},
}
impl core::fmt::Display for BandShapeError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
BandShapeError::Pvq(e) => write!(f, "oxideav-opus: §4.3.4 band shape PVQ: {e}"),
BandShapeError::Spreading(e) => {
write!(f, "oxideav-opus: §4.3.4 band shape spreading: {e}")
}
BandShapeError::TfHadamard(e) => {
write!(f, "oxideav-opus: §4.3.4 band shape TF transform: {e}")
}
BandShapeError::OutputBufferTooSmall { required, provided } => write!(
f,
"oxideav-opus: §4.3.4 band shape output buffer too small: \
required={required}, provided={provided}"
),
}
}
}
impl std::error::Error for BandShapeError {}
impl From<PvqShapeError> for BandShapeError {
fn from(e: PvqShapeError) -> Self {
BandShapeError::Pvq(e)
}
}
impl From<SpreadingError> for BandShapeError {
fn from(e: SpreadingError) -> Self {
BandShapeError::Spreading(e)
}
}
impl From<TfHadamardError> for BandShapeError {
fn from(e: TfHadamardError) -> Self {
BandShapeError::TfHadamard(e)
}
}
#[allow(clippy::too_many_arguments)]
pub fn decode_band_shape_into(
rd: &mut crate::RangeDecoder<'_>,
n: u32,
k: u32,
spread: u8,
tf_adjust: TfAdjustment,
nb_blocks: usize,
out: &mut [f64],
) -> Result<usize, BandShapeError> {
let n_usize = n as usize;
if out.len() < n_usize {
return Err(BandShapeError::OutputBufferTooSmall {
required: n_usize,
provided: out.len(),
});
}
let band = &mut out[..n_usize];
decode_pvq_shape_into(rd, n, k, band)?;
apply_spreading(band, k, spread, nb_blocks)?;
let direction = TfDirection::from_adjustment(tf_adjust);
apply_tf_hadamard(band, nb_blocks, direction)?;
Ok(n_usize)
}
#[allow(clippy::too_many_arguments)]
pub fn decode_band_shape(
rd: &mut crate::RangeDecoder<'_>,
n: u32,
k: u32,
spread: u8,
tf_adjust: TfAdjustment,
nb_blocks: usize,
) -> Result<Vec<f64>, BandShapeError> {
let mut out = vec![0.0f64; n as usize];
decode_band_shape_into(rd, n, k, spread, tf_adjust, nb_blocks, &mut out)?;
Ok(out)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::RangeDecoder;
fn l2(x: &[f64]) -> f64 {
x.iter().map(|v| v * v).sum::<f64>().sqrt()
}
#[test]
fn k_zero_yields_all_zero_shape() {
let buf = [0xAB, 0xCD, 0xEF, 0x12];
let mut rd = RangeDecoder::new(&buf);
let shape = decode_band_shape(&mut rd, 6, 0, 2, 0, 1).unwrap();
assert_eq!(shape, vec![0.0; 6]);
}
#[test]
fn nonzero_k_no_spread_no_tf_is_unit_norm() {
let buf = [0x12, 0x34, 0x56, 0x78, 0x9A, 0xBC];
let mut rd = RangeDecoder::new(&buf);
let shape = decode_band_shape(&mut rd, 8, 3, 0, 0, 1).unwrap();
assert!((l2(&shape) - 1.0).abs() < 1e-9, "norm = {}", l2(&shape));
}
#[test]
fn spreading_preserves_unit_norm() {
for spread in 1u8..=3 {
let buf = [0x42, 0x18, 0x7F, 0x03, 0xC1, 0x5E];
let mut rd = RangeDecoder::new(&buf);
let shape = decode_band_shape(&mut rd, 8, 4, spread, 0, 1).unwrap();
assert!(
(l2(&shape) - 1.0).abs() < 1e-9,
"spread {spread}: norm = {}",
l2(&shape)
);
}
}
#[test]
fn tf_transform_preserves_unit_norm() {
let buf = [0x71, 0x2C, 0x90, 0xEE, 0x44, 0x08];
let mut rd = RangeDecoder::new(&buf);
let shape = decode_band_shape(&mut rd, 8, 3, 0, -1, 2).unwrap();
assert!((l2(&shape) - 1.0).abs() < 1e-9, "norm = {}", l2(&shape));
}
#[test]
fn output_buffer_too_small_errors() {
let buf = [0x00, 0x00, 0x00, 0x00];
let mut rd = RangeDecoder::new(&buf);
let mut out = [0.0f64; 3];
let r = decode_band_shape_into(&mut rd, 8, 2, 1, 0, 1, &mut out);
assert_eq!(
r,
Err(BandShapeError::OutputBufferTooSmall {
required: 8,
provided: 3
})
);
}
#[test]
fn bad_block_geometry_surfaces_error() {
let buf = [0x33, 0x77, 0xAA, 0xDD];
let mut rd = RangeDecoder::new(&buf);
let r = decode_band_shape(&mut rd, 5, 1, 0, -1, 2);
assert!(matches!(r, Err(BandShapeError::Spreading(_))), "{r:?}");
}
#[test]
fn tf_levels_exceed_blocks_surfaces_tf_error() {
let buf = [0x33, 0x77, 0xAA, 0xDD];
let mut rd = RangeDecoder::new(&buf);
let r = decode_band_shape(&mut rd, 6, 1, 0, -1, 1);
assert!(matches!(r, Err(BandShapeError::TfHadamard(_))), "{r:?}");
}
#[test]
fn into_and_allocating_agree() {
let buf = [0x21, 0x43, 0x65, 0x87, 0xA9, 0xCB];
let mut rd_a = RangeDecoder::new(&buf);
let mut rd_b = RangeDecoder::new(&buf);
let alloc = decode_band_shape(&mut rd_a, 8, 5, 2, 0, 1).unwrap();
let mut into = vec![0.0f64; 8];
let written = decode_band_shape_into(&mut rd_b, 8, 5, 2, 0, 1, &mut into).unwrap();
assert_eq!(written, 8);
assert_eq!(alloc, into);
}
#[test]
fn deterministic_across_calls() {
let buf = [0x5A, 0xA5, 0x33, 0xCC, 0x0F, 0xF0];
let mut rd1 = RangeDecoder::new(&buf);
let mut rd2 = RangeDecoder::new(&buf);
let a = decode_band_shape(&mut rd1, 6, 4, 1, 0, 1).unwrap();
let b = decode_band_shape(&mut rd2, 6, 4, 1, 0, 1).unwrap();
assert_eq!(a, b);
}
}