#![forbid(unsafe_code)]
use crate::AudioError;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum CeltMode {
#[default]
Normal,
Custom,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
#[allow(dead_code)]
pub enum CeltFrameSize {
Ms2_5,
Ms5,
#[default]
Ms10,
Ms20,
}
impl CeltFrameSize {
#[must_use]
pub fn samples_48khz(self) -> usize {
match self {
CeltFrameSize::Ms2_5 => 120,
CeltFrameSize::Ms5 => 240,
CeltFrameSize::Ms10 => 480,
CeltFrameSize::Ms20 => 960,
}
}
#[must_use]
pub fn short_blocks(self) -> usize {
1
}
#[must_use]
pub fn duration_us(self) -> u32 {
match self {
CeltFrameSize::Ms2_5 => 2500,
CeltFrameSize::Ms5 => 5000,
CeltFrameSize::Ms10 => 10000,
CeltFrameSize::Ms20 => 20000,
}
}
}
#[derive(Debug, Clone, Default)]
pub struct BandEnergy {
pub values: Vec<i16>,
pub band_count: usize,
}
impl BandEnergy {
#[must_use]
pub fn new(band_count: usize) -> Self {
Self {
values: vec![0; band_count],
band_count,
}
}
#[must_use]
pub fn get(&self, band: usize) -> Option<i16> {
self.values.get(band).copied()
}
pub fn set(&mut self, band: usize, value: i16) {
if band < self.band_count {
self.values[band] = value;
}
}
#[must_use]
pub fn total_energy(&self) -> i32 {
self.values.iter().map(|&e| i32::from(e)).sum()
}
}
#[derive(Debug, Clone, Default)]
pub struct PitchPeriod {
pub period: u16,
pub gain: i16,
pub tap_count: u8,
pub taps: [i16; 3],
}
impl PitchPeriod {
pub const MIN_PERIOD: u16 = 15;
pub const MAX_PERIOD: u16 = 1022;
#[must_use]
pub fn new(period: u16, gain: i16) -> Self {
Self {
period: period.clamp(Self::MIN_PERIOD, Self::MAX_PERIOD),
gain,
tap_count: 1,
taps: [gain, 0, 0],
}
}
#[must_use]
pub fn is_active(&self) -> bool {
self.gain != 0
}
}
#[derive(Debug, Clone)]
#[allow(dead_code)]
pub struct CeltBandConfig {
pub band_count: usize,
pub band_boundaries: Vec<usize>,
pub bits_per_band: Vec<u16>,
}
impl Default for CeltBandConfig {
fn default() -> Self {
Self::new_48khz()
}
}
impl CeltBandConfig {
#[must_use]
pub fn new_48khz() -> Self {
let band_boundaries = vec![
0, 1, 2, 3, 4, 5, 6, 7, 8, 10, 12, 14, 17, 21, 25, 30, 36, 43, 52, 63, 76, 92,
];
let band_count = band_boundaries.len() - 1;
Self {
band_count,
band_boundaries,
bits_per_band: vec![0; band_count],
}
}
#[must_use]
pub fn band_range(&self, band: usize) -> Option<(usize, usize)> {
if band < self.band_count {
Some((self.band_boundaries[band], self.band_boundaries[band + 1]))
} else {
None
}
}
#[must_use]
pub fn band_width(&self, band: usize) -> Option<usize> {
self.band_range(band).map(|(start, end)| end - start)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
#[allow(dead_code)]
pub enum TransientType {
#[default]
None,
Short,
}
#[derive(Debug, Clone, Default)]
#[allow(dead_code)]
pub struct CeltFrame {
pub mode: CeltMode,
pub frame_size: CeltFrameSize,
pub channels: u8,
pub transient: TransientType,
pub intra: bool,
pub energy: BandEnergy,
pub pitch: Option<PitchPeriod>,
pub fine_energy: Vec<i16>,
pub coefficients: Vec<Vec<i16>>,
pub spread: u8,
pub dual_stereo: bool,
pub intensity_stereo_band: Option<usize>,
}
impl CeltFrame {
#[must_use]
pub fn new(frame_size: CeltFrameSize, channels: u8) -> Self {
let band_config = CeltBandConfig::new_48khz();
let band_count = band_config.band_count;
Self {
frame_size,
channels,
energy: BandEnergy::new(band_count),
fine_energy: vec![0; band_count],
coefficients: vec![Vec::new(); band_count],
..Default::default()
}
}
#[must_use]
pub fn band_count(&self) -> usize {
self.energy.band_count
}
#[must_use]
pub fn has_transient(&self) -> bool {
self.transient == TransientType::Short
}
#[must_use]
pub fn is_stereo(&self) -> bool {
self.channels == 2
}
#[must_use]
pub fn sample_count(&self) -> usize {
self.frame_size.samples_48khz() * usize::from(self.channels)
}
}
#[derive(Debug, Clone, Default)]
#[allow(dead_code)]
pub struct CeltDecoderState {
pub prev_energy: BandEnergy,
pub prev_fine_energy: Vec<i16>,
pub overlap_buffer: Vec<f32>,
pub pitch_buffer: Vec<f32>,
pub synthesis_state: Vec<f32>,
pub channels: u8,
pub frame_size: CeltFrameSize,
pub postfilter_enabled: bool,
}
impl CeltDecoderState {
#[must_use]
pub fn new(channels: u8, frame_size: CeltFrameSize) -> Self {
let band_config = CeltBandConfig::new_48khz();
let samples = frame_size.samples_48khz() * usize::from(channels);
Self {
prev_energy: BandEnergy::new(band_config.band_count),
prev_fine_energy: vec![0; band_config.band_count],
overlap_buffer: vec![0.0; samples / 2],
pitch_buffer: vec![0.0; PitchPeriod::MAX_PERIOD as usize + samples],
synthesis_state: vec![0.0; samples],
channels,
frame_size,
postfilter_enabled: true,
}
}
pub fn reset(&mut self) {
self.prev_energy.values.fill(0);
self.prev_fine_energy.fill(0);
self.overlap_buffer.fill(0.0);
self.pitch_buffer.fill(0.0);
self.synthesis_state.fill(0.0);
}
#[allow(dead_code)]
pub fn conceal_frame(&mut self) {
for energy in &mut self.prev_energy.values {
*energy = (*energy).saturating_sub(64); }
}
}
#[derive(Debug, Clone, Default)]
#[allow(dead_code)]
pub struct PvqAllocation {
pub pulses_per_band: Vec<u16>,
pub total_bits: u32,
pub remaining_bits: u32,
}
impl PvqAllocation {
#[must_use]
pub fn new(band_count: usize) -> Self {
Self {
pulses_per_band: vec![0; band_count],
total_bits: 0,
remaining_bits: 0,
}
}
#[must_use]
pub fn pulses(&self, band: usize) -> u16 {
self.pulses_per_band.get(band).copied().unwrap_or(0)
}
}
#[must_use]
fn make_sine_window(size: usize) -> Vec<f32> {
(0..size)
.map(|i| {
let x = std::f64::consts::PI * (i as f64 + 0.5) / size as f64;
x.sin() as f32
})
.collect()
}
#[must_use]
pub fn imdct(coeffs: &[f32]) -> Vec<f32> {
let n = coeffs.len();
if n == 0 {
return Vec::new();
}
let two_n = 2 * n;
let mut output = vec![0.0f32; two_n];
let scale = 2.0_f32 / n as f32;
for i in 0..two_n {
let mut sum = 0.0f32;
for (k, &c) in coeffs.iter().enumerate() {
let angle = std::f64::consts::PI / n as f64
* (k as f64 + 0.5)
* (i as f64 - n as f64 + 0.5 + n as f64 * 2.0);
let angle2 = std::f64::consts::PI / n as f64 * (k as f64 + 0.5) * (i as f64 + 0.5);
let _ = angle; sum += c * (angle2.cos() as f32);
}
output[i] = sum * scale;
}
output
}
#[must_use]
pub fn imdct_fast(coeffs: &[f32]) -> Vec<f32> {
let n = coeffs.len();
if n == 0 {
return Vec::new();
}
let two_n = 2 * n;
let mut output = vec![0.0f32; two_n];
let scale = 2.0_f64 / n as f64;
for i in 0..two_n {
let mut sum = 0.0_f64;
for (k, &c) in coeffs.iter().enumerate() {
let angle =
std::f64::consts::PI * (2 * i + n + 1) as f64 * (2 * k + 1) as f64 / (4 * n) as f64;
sum += c as f64 * angle.cos();
}
output[i] = (sum * scale) as f32;
}
output
}
pub fn decode_frame(
frame: &CeltFrame,
state: &mut CeltDecoderState,
) -> Result<Vec<f32>, AudioError> {
let band_config = CeltBandConfig::new_48khz();
let n = frame.frame_size.samples_48khz(); let channels = usize::from(frame.channels);
let mut interleaved = vec![0.0f32; n * channels];
for ch in 0..channels {
let mut spectrum = vec![0.0f32; n];
for band in 0..band_config.band_count {
let (start, end) = match band_config.band_range(band) {
Some(r) => r,
None => break,
};
if start >= n {
break;
}
let end = end.min(n);
let width = end - start;
if width == 0 {
continue;
}
let energy_q8 = frame.energy.get(band).unwrap_or(0);
let energy_db = f64::from(energy_q8) / 256.0;
let amplitude = (10.0_f64.powf(energy_db / 20.0)) as f32;
let band_coeffs = frame.coefficients.get(band);
if let Some(coeffs) = band_coeffs {
if !coeffs.is_empty() {
let sq_sum: f64 = coeffs.iter().map(|&c| (c as f64) * (c as f64)).sum();
let norm = if sq_sum > 0.0 {
sq_sum.sqrt() as f32
} else {
1.0
};
for (j, &c) in coeffs.iter().enumerate().take(width) {
spectrum[start + j] = (c as f32 / norm) * amplitude;
}
for j in coeffs.len()..width {
spectrum[start + j] = 0.0;
}
} else {
let per_bin = amplitude / (width as f32).sqrt();
for j in 0..width {
spectrum[start + j] = if j % 2 == 0 { per_bin } else { -per_bin };
}
}
} else {
let per_bin = amplitude / (width as f32).sqrt();
for j in 0..width {
spectrum[start + j] = if j % 2 == 0 { per_bin } else { -per_bin };
}
}
}
let time_domain = imdct_fast(&spectrum);
let window = make_sine_window(2 * n);
let windowed: Vec<f32> = time_domain
.iter()
.zip(window.iter())
.map(|(&s, &w)| s * w)
.collect();
let overlap_len = state.overlap_buffer.len().min(n);
while state.overlap_buffer.len() < n {
state.overlap_buffer.push(0.0);
}
let mut channel_out = vec![0.0f32; n];
for i in 0..n {
channel_out[i] = windowed[i]
+ if i < overlap_len {
state.overlap_buffer[i]
} else {
0.0
};
}
let second_half_start = n;
state.overlap_buffer.resize(n, 0.0);
for i in 0..n {
state.overlap_buffer[i] = windowed[second_half_start + i];
}
if state.postfilter_enabled {
if let Some(ref pitch) = frame.pitch {
if pitch.is_active() {
let period = usize::from(pitch.period);
let gain = pitch.gain as f64 / 32768.0;
for i in period..channel_out.len() {
channel_out[i] += (channel_out[i - period] as f64 * gain) as f32;
}
}
}
}
for s in &mut channel_out {
*s = s.clamp(-1.0, 1.0);
}
for band in 0..band_config.band_count {
if let Some(e) = frame.energy.get(band) {
state.prev_energy.set(band, e);
}
}
for (i, &s) in channel_out.iter().enumerate() {
interleaved[i * channels + ch] = s;
}
}
Ok(interleaved)
}
#[must_use]
pub fn mdct(input: &[f32]) -> Vec<f32> {
let two_n = input.len();
if two_n == 0 {
return Vec::new();
}
let n = two_n / 2;
let mut output = vec![0.0f32; n];
let scale = 2.0_f64 / two_n as f64;
for k in 0..n {
let mut sum = 0.0_f64;
for (i, &x) in input.iter().enumerate() {
let angle =
std::f64::consts::PI * (2 * i + n + 1) as f64 * (2 * k + 1) as f64 / (4 * n) as f64;
sum += x as f64 * angle.cos();
}
output[k] = (sum * scale) as f32;
}
output
}
#[derive(Debug, Clone, Default)]
pub struct CeltEncoderState {
pub prev_tail: Vec<f32>,
pub channels: u8,
pub frame_size: CeltFrameSize,
}
impl CeltEncoderState {
#[must_use]
pub fn new(channels: u8, frame_size: CeltFrameSize) -> Self {
let n = frame_size.samples_48khz();
Self {
prev_tail: vec![0.0; n],
channels,
frame_size,
}
}
pub fn reset(&mut self) {
self.prev_tail.fill(0.0);
}
}
pub fn encode_frame(
samples: &[f32],
state: &mut CeltEncoderState,
) -> Result<CeltFrame, AudioError> {
let n = state.frame_size.samples_48khz();
let channels = usize::from(state.channels);
if samples.len() != n * channels {
return Err(AudioError::InvalidParameter(format!(
"Expected {} samples, got {}",
n * channels,
samples.len()
)));
}
let band_config = CeltBandConfig::new_48khz();
let mut frame = CeltFrame::new(state.frame_size, state.channels);
for ch in 0..channels {
let channel_samples: Vec<f32> = (0..n).map(|i| samples[i * channels + ch]).collect();
let window = make_sine_window(2 * n);
state.prev_tail.resize(n, 0.0);
let windowed_input: Vec<f32> = (0..2 * n)
.map(|i| {
let s = if i < n {
state.prev_tail[i]
} else {
channel_samples[i - n]
};
s * window[i]
})
.collect();
state.prev_tail = channel_samples.clone();
let spectrum = mdct(&windowed_input);
for band in 0..band_config.band_count {
let (start, end) = match band_config.band_range(band) {
Some(r) => r,
None => break,
};
if start >= n {
break;
}
let end = end.min(n);
let width = end - start;
if width == 0 {
continue;
}
let band_slice = &spectrum[start..end];
let sq_sum: f64 = band_slice.iter().map(|&c| (c as f64) * (c as f64)).sum();
let rms = (sq_sum / width as f64).sqrt() as f32;
let energy_db = if rms > 1e-10_f32 {
20.0_f32 * rms.log10()
} else {
-60.0_f32
};
let energy_q8 = ((energy_db * 256.0) as i32).clamp(-32768, 32767) as i16;
if ch == 0 {
frame.energy.set(band, energy_q8);
}
let norm = if rms > 1e-10_f32 { rms } else { 1.0 };
let quantized: Vec<i16> = band_slice
.iter()
.map(|&c| ((c / norm * 1024.0) as i32).clamp(-32767, 32767) as i16)
.collect();
if ch == 0 {
frame.coefficients[band] = quantized;
}
}
}
Ok(frame)
}
pub struct BitPacker {
buffer: Vec<u8>,
current_byte: u8,
bits_used: u8,
}
impl BitPacker {
#[must_use]
pub fn new() -> Self {
Self {
buffer: Vec::new(),
current_byte: 0,
bits_used: 0,
}
}
pub fn write_bits(&mut self, mut value: u32, mut n: u8) {
while n > 0 {
let space = 8 - self.bits_used;
let take = n.min(space);
let mask = (1u32 << take) - 1;
self.current_byte |= ((value & mask) as u8) << self.bits_used;
value >>= take;
n -= take;
self.bits_used += take;
if self.bits_used == 8 {
self.buffer.push(self.current_byte);
self.current_byte = 0;
self.bits_used = 0;
}
}
}
pub fn flush(&mut self) {
if self.bits_used > 0 {
self.buffer.push(self.current_byte);
self.current_byte = 0;
self.bits_used = 0;
}
}
#[must_use]
pub fn finish(mut self) -> Vec<u8> {
self.flush();
self.buffer
}
}
impl Default for BitPacker {
fn default() -> Self {
Self::new()
}
}
pub struct BitUnpacker {
data: Vec<u8>,
byte_pos: usize,
bits_used: u8,
}
impl BitUnpacker {
#[must_use]
pub fn new(data: &[u8]) -> Self {
Self {
data: data.to_vec(),
byte_pos: 0,
bits_used: 0,
}
}
pub fn read_bits(&mut self, mut n: u8) -> u32 {
let mut result = 0u32;
let mut shift = 0u8;
while n > 0 {
if self.byte_pos >= self.data.len() {
break;
}
let space = 8 - self.bits_used;
let take = n.min(space);
let mask: u8 = if take >= 8 {
0xFF
} else {
(1u8 << take).wrapping_sub(1)
};
let bits = (self.data[self.byte_pos] >> self.bits_used) & mask;
result |= (bits as u32) << shift;
shift += take;
self.bits_used += take;
n -= take;
if self.bits_used == 8 {
self.bits_used = 0;
self.byte_pos += 1;
}
}
result
}
}
pub fn serialize_celt_frame(frame: &CeltFrame) -> Result<Vec<u8>, AudioError> {
let mut packer = BitPacker::new();
let band_count = frame.energy.band_count.min(255) as u8;
packer.write_bits(u32::from(band_count), 8);
for band in 0..band_count as usize {
let energy = frame.energy.get(band).unwrap_or(0);
let eu = energy as u16;
packer.write_bits(u32::from(eu & 0xFF), 8);
packer.write_bits(u32::from((eu >> 8) & 0xFF), 8);
let coeffs = frame
.coefficients
.get(band)
.map_or(&[][..], |v| v.as_slice());
let coeff_count = coeffs.len().min(255) as u8;
packer.write_bits(u32::from(coeff_count), 8);
for &c in &coeffs[..coeff_count as usize] {
let cu = c as u16;
packer.write_bits(u32::from(cu & 0xFF), 8);
packer.write_bits(u32::from((cu >> 8) & 0xFF), 8);
}
}
Ok(packer.finish())
}
pub fn deserialize_celt_frame(
data: &[u8],
frame_size: CeltFrameSize,
channels: u8,
) -> Result<CeltFrame, AudioError> {
if data.is_empty() {
return Err(AudioError::InvalidData("Empty CELT frame data".into()));
}
let mut unpacker = BitUnpacker::new(data);
let band_count = unpacker.read_bits(8) as usize;
let band_config = CeltBandConfig::new_48khz();
let mut frame = CeltFrame::new(frame_size, channels);
for band in 0..band_count.min(band_config.band_count) {
let lo = unpacker.read_bits(8) as u16;
let hi = unpacker.read_bits(8) as u16;
let energy = ((hi << 8) | lo) as i16;
frame.energy.set(band, energy);
let coeff_count = unpacker.read_bits(8) as usize;
let mut coeffs = Vec::with_capacity(coeff_count);
for _ in 0..coeff_count {
let lo2 = unpacker.read_bits(8) as u16;
let hi2 = unpacker.read_bits(8) as u16;
let c = ((hi2 << 8) | lo2) as i16;
coeffs.push(c);
}
if band < frame.coefficients.len() {
frame.coefficients[band] = coeffs;
}
}
Ok(frame)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_celt_frame_size_samples() {
assert_eq!(CeltFrameSize::Ms2_5.samples_48khz(), 120);
assert_eq!(CeltFrameSize::Ms5.samples_48khz(), 240);
assert_eq!(CeltFrameSize::Ms10.samples_48khz(), 480);
assert_eq!(CeltFrameSize::Ms20.samples_48khz(), 960);
}
#[test]
fn test_celt_frame_size_duration() {
assert_eq!(CeltFrameSize::Ms2_5.duration_us(), 2500);
assert_eq!(CeltFrameSize::Ms5.duration_us(), 5000);
assert_eq!(CeltFrameSize::Ms10.duration_us(), 10000);
assert_eq!(CeltFrameSize::Ms20.duration_us(), 20000);
}
#[test]
fn test_band_energy() {
let mut energy = BandEnergy::new(21);
assert_eq!(energy.band_count, 21);
assert_eq!(energy.get(0), Some(0));
energy.set(5, 100);
assert_eq!(energy.get(5), Some(100));
}
#[test]
fn test_band_energy_total() {
let mut energy = BandEnergy::new(4);
energy.values = vec![10, 20, 30, 40];
assert_eq!(energy.total_energy(), 100);
}
#[test]
fn test_pitch_period() {
let pitch = PitchPeriod::new(100, 1000);
assert_eq!(pitch.period, 100);
assert_eq!(pitch.gain, 1000);
assert!(pitch.is_active());
let no_pitch = PitchPeriod::default();
assert!(!no_pitch.is_active());
}
#[test]
fn test_pitch_period_clamping() {
let pitch = PitchPeriod::new(5, 100); assert_eq!(pitch.period, PitchPeriod::MIN_PERIOD);
let pitch = PitchPeriod::new(2000, 100); assert_eq!(pitch.period, PitchPeriod::MAX_PERIOD);
}
#[test]
fn test_celt_band_config() {
let config = CeltBandConfig::new_48khz();
assert_eq!(config.band_count, 21);
let range = config.band_range(0);
assert!(range.is_some());
assert_eq!(range.expect("should succeed"), (0, 1));
}
#[test]
fn test_celt_band_width() {
let config = CeltBandConfig::new_48khz();
assert_eq!(config.band_width(0), Some(1));
assert_eq!(config.band_width(20), Some(16)); }
#[test]
fn test_celt_frame() {
let frame = CeltFrame::new(CeltFrameSize::Ms20, 2);
assert_eq!(frame.frame_size, CeltFrameSize::Ms20);
assert_eq!(frame.channels, 2);
assert_eq!(frame.band_count(), 21);
assert!(frame.is_stereo());
}
#[test]
fn test_celt_frame_sample_count() {
let mono = CeltFrame::new(CeltFrameSize::Ms10, 1);
assert_eq!(mono.sample_count(), 480);
let stereo = CeltFrame::new(CeltFrameSize::Ms10, 2);
assert_eq!(stereo.sample_count(), 960);
}
#[test]
fn test_celt_decoder_state() {
let state = CeltDecoderState::new(2, CeltFrameSize::Ms20);
assert_eq!(state.channels, 2);
assert_eq!(state.frame_size, CeltFrameSize::Ms20);
assert!(!state.overlap_buffer.is_empty());
}
#[test]
fn test_celt_decoder_state_reset() {
let mut state = CeltDecoderState::new(1, CeltFrameSize::Ms10);
state.prev_energy.set(0, 100);
state.reset();
assert_eq!(state.prev_energy.get(0), Some(0));
}
#[test]
fn test_pvq_allocation() {
let alloc = PvqAllocation::new(21);
assert_eq!(alloc.pulses_per_band.len(), 21);
assert_eq!(alloc.pulses(0), 0);
assert_eq!(alloc.pulses(100), 0); }
#[test]
fn test_transient_type() {
let frame = CeltFrame::new(CeltFrameSize::Ms20, 1);
assert!(!frame.has_transient());
}
}