#![forbid(unsafe_code)]
use crate::AudioError;
const RANGE_MIN: u32 = 1 << 23;
#[derive(Debug, Clone)]
pub struct Symbol {
pub fl: u32,
pub fh: u32,
pub ft: u32,
}
impl Symbol {
#[must_use]
pub fn new(fl: u32, fh: u32, ft: u32) -> Self {
Self { fl, fh, ft }
}
#[must_use]
pub fn binary(probability: u32, total: u32) -> Self {
Self {
fl: 0,
fh: probability,
ft: total,
}
}
#[must_use]
pub fn uniform(index: u32, count: u32) -> Self {
Self {
fl: index,
fh: index + 1,
ft: count,
}
}
#[must_use]
pub fn probability(&self) -> f64 {
if self.ft == 0 {
0.0
} else {
f64::from(self.fh - self.fl) / f64::from(self.ft)
}
}
}
#[derive(Debug, Clone)]
#[allow(dead_code)]
pub struct RangeDecoder {
data: Vec<u8>,
byte_pos: usize,
bit_pos: u8,
range: u32,
value: u32,
bits_consumed: u32,
total_bits: u32,
eos: bool,
}
impl RangeDecoder {
#[allow(clippy::cast_possible_truncation)]
pub fn new(data: &[u8]) -> Result<Self, AudioError> {
if data.is_empty() {
return Err(AudioError::InvalidData("Empty range coder data".into()));
}
let mut decoder = Self {
data: data.to_vec(),
byte_pos: 0,
bit_pos: 0,
range: 128,
value: 0,
bits_consumed: 0,
total_bits: (data.len() * 8) as u32,
eos: false,
};
decoder.value = 127u32.saturating_sub(u32::from(data[0] >> 1));
decoder.byte_pos = 1;
decoder.bits_consumed = 9;
Ok(decoder)
}
#[must_use]
pub fn bits_remaining(&self) -> u32 {
self.total_bits.saturating_sub(self.bits_consumed)
}
#[must_use]
pub fn is_eos(&self) -> bool {
self.eos
}
fn read_bit(&mut self) -> u8 {
if self.byte_pos >= self.data.len() {
self.eos = true;
return 0;
}
let bit = (self.data[self.byte_pos] >> (7 - self.bit_pos)) & 1;
self.bit_pos += 1;
if self.bit_pos >= 8 {
self.bit_pos = 0;
self.byte_pos += 1;
}
self.bits_consumed += 1;
bit
}
fn normalize(&mut self) {
while self.range < RANGE_MIN {
self.range <<= 8;
let byte = if self.byte_pos < self.data.len() {
self.data[self.byte_pos]
} else {
0
};
self.byte_pos += 1;
self.value = (self.value << 8) | u32::from(byte);
self.bits_consumed += 8;
}
}
pub fn decode_symbol(&mut self, ft: u32) -> Result<u32, AudioError> {
if ft == 0 {
return Err(AudioError::InvalidData(
"Zero total frequency in range decoder".into(),
));
}
self.normalize();
let fs = self.range / ft;
if fs == 0 {
return Err(AudioError::InvalidData(
"Range too small for frequency".into(),
));
}
let k = self.value.min((fs * ft).saturating_sub(1)) / fs;
Ok(k.min(ft - 1))
}
pub fn decode_update(&mut self, sym: &Symbol) -> Result<(), AudioError> {
if sym.ft == 0 {
return Err(AudioError::InvalidData("Zero total frequency".into()));
}
let fs = self.range / sym.ft;
self.value = self.value.saturating_sub(fs * sym.fl);
self.range = if sym.fl + (sym.fh - sym.fl) == sym.ft {
self.range.saturating_sub(fs * sym.fl)
} else {
fs * (sym.fh - sym.fl)
};
Ok(())
}
#[allow(dead_code)]
pub fn decode(&mut self, sym: &Symbol) -> Result<u32, AudioError> {
let k = self.decode_symbol(sym.ft)?;
self.decode_update(sym)?;
Ok(k)
}
pub fn decode_uniform(&mut self, count: u32) -> Result<u32, AudioError> {
if count == 0 {
return Err(AudioError::InvalidData(
"Zero count in uniform decode".into(),
));
}
if count == 1 {
return Ok(0);
}
let k = self.decode_symbol(count)?;
let sym = Symbol::uniform(k, count);
self.decode_update(&sym)?;
Ok(k)
}
pub fn decode_bits(&mut self, bits: u32) -> Result<u32, AudioError> {
if bits == 0 {
return Ok(0);
}
if bits > 32 {
return Err(AudioError::InvalidData("Too many bits requested".into()));
}
let mut value = 0u32;
for _ in 0..bits {
value = (value << 1) | u32::from(self.read_bit());
}
Ok(value)
}
#[allow(dead_code)]
pub fn decode_uint(&mut self, bits: u32) -> Result<u32, AudioError> {
self.decode_bits(bits)
}
#[must_use]
pub fn tell(&self) -> u32 {
self.bits_consumed
}
#[must_use]
pub fn tell_frac(&self) -> u32 {
self.bits_consumed * 8
}
}
#[derive(Debug, Clone)]
#[allow(dead_code)]
pub struct LaplaceDecoder {
mean: i32,
decay: u32,
}
impl LaplaceDecoder {
#[must_use]
pub fn new(mean: i32, decay: u32) -> Self {
Self { mean, decay }
}
#[allow(dead_code, clippy::cast_possible_wrap)]
pub fn decode(&self, range_decoder: &mut RangeDecoder) -> Result<i32, AudioError> {
let sign = range_decoder.decode_bits(1)? != 0;
let magnitude = range_decoder.decode_uniform(self.decay)?;
let value = self.mean + magnitude as i32;
Ok(if sign { -value } else { value })
}
}
#[derive(Debug, Clone)]
pub struct IcdfTable {
values: Vec<u16>,
total: u32,
}
impl IcdfTable {
#[must_use]
pub fn from_probabilities(probs: &[u16]) -> Self {
let mut values = Vec::with_capacity(probs.len() + 1);
let mut cumsum = 0u16;
values.push(0);
for &p in probs {
cumsum = cumsum.saturating_add(p);
values.push(cumsum);
}
let total = u32::from(cumsum);
Self { values, total }
}
#[must_use]
pub fn symbol_count(&self) -> usize {
self.values.len().saturating_sub(1)
}
#[must_use]
pub fn symbol(&self, index: usize) -> Option<Symbol> {
if index + 1 < self.values.len() {
Some(Symbol {
fl: u32::from(self.values[index]),
fh: u32::from(self.values[index + 1]),
ft: self.total,
})
} else {
None
}
}
#[allow(dead_code)]
pub fn decode(&self, range_decoder: &mut RangeDecoder) -> Result<usize, AudioError> {
let k = range_decoder.decode_symbol(self.total)?;
let mut lo = 0;
let mut hi = self.values.len() - 1;
while lo < hi {
let mid = (lo + hi) / 2;
if u32::from(self.values[mid + 1]) <= k {
lo = mid + 1;
} else {
hi = mid;
}
}
let sym = self
.symbol(lo)
.ok_or_else(|| AudioError::InvalidData("Invalid symbol index".into()))?;
range_decoder.decode_update(&sym)?;
Ok(lo)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_symbol_new() {
let sym = Symbol::new(10, 20, 100);
assert_eq!(sym.fl, 10);
assert_eq!(sym.fh, 20);
assert_eq!(sym.ft, 100);
}
#[test]
fn test_symbol_binary() {
let sym = Symbol::binary(50, 100);
assert_eq!(sym.fl, 0);
assert_eq!(sym.fh, 50);
assert_eq!(sym.ft, 100);
}
#[test]
fn test_symbol_uniform() {
let sym = Symbol::uniform(5, 10);
assert_eq!(sym.fl, 5);
assert_eq!(sym.fh, 6);
assert_eq!(sym.ft, 10);
}
#[test]
fn test_symbol_probability() {
let sym = Symbol::new(0, 50, 100);
assert!((sym.probability() - 0.5).abs() < f64::EPSILON);
let zero = Symbol::new(0, 0, 0);
assert!((zero.probability() - 0.0).abs() < f64::EPSILON);
}
#[test]
fn test_range_decoder_creation() {
let data = vec![0x80, 0x00, 0x00, 0x00];
let decoder = RangeDecoder::new(&data).expect("should succeed");
assert!(!decoder.is_eos());
}
#[test]
fn test_range_decoder_empty() {
let result = RangeDecoder::new(&[]);
assert!(result.is_err());
}
#[test]
fn test_range_decoder_bits_remaining() {
let data = vec![0x80, 0x00];
let decoder = RangeDecoder::new(&data).expect("should succeed");
assert!(decoder.bits_remaining() > 0);
}
#[test]
fn test_decode_bits() {
let data = vec![0xFF, 0x00, 0xFF, 0x00];
let mut decoder = RangeDecoder::new(&data).expect("should succeed");
let _ = decoder.decode_bits(4);
}
#[test]
fn test_decode_uniform_one() {
let data = vec![0x80, 0x00];
let mut decoder = RangeDecoder::new(&data).expect("should succeed");
let result = decoder.decode_uniform(1).expect("should succeed");
assert_eq!(result, 0);
}
#[test]
fn test_tell() {
let data = vec![0x80, 0x00, 0x00, 0x00];
let decoder = RangeDecoder::new(&data).expect("should succeed");
assert!(decoder.tell() > 0);
}
#[test]
fn test_icdf_table() {
let probs = vec![10, 20, 30, 40];
let icdf = IcdfTable::from_probabilities(&probs);
assert_eq!(icdf.symbol_count(), 4);
assert_eq!(icdf.total, 100);
}
#[test]
fn test_icdf_symbol() {
let probs = vec![10, 20, 30, 40];
let icdf = IcdfTable::from_probabilities(&probs);
let sym0 = icdf.symbol(0).expect("should succeed");
assert_eq!(sym0.fl, 0);
assert_eq!(sym0.fh, 10);
let sym1 = icdf.symbol(1).expect("should succeed");
assert_eq!(sym1.fl, 10);
assert_eq!(sym1.fh, 30);
}
#[test]
fn test_laplace_decoder() {
let decoder = LaplaceDecoder::new(0, 100);
assert_eq!(decoder.mean, 0);
assert_eq!(decoder.decay, 100);
}
}