use crate::{AudioError, AudioResult};
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct HuffPair {
pub x: i16,
pub y: i16,
}
#[derive(Clone, Copy, Debug)]
struct HuffEntry {
len: u8,
x: i8,
y: i8,
}
pub struct HuffmanDecoder<'a> {
pub(crate) data: &'a [u8],
byte_pos: usize,
bit_pos: u8,
}
impl<'a> HuffmanDecoder<'a> {
#[must_use]
pub const fn new(data: &'a [u8]) -> Self {
Self {
data,
byte_pos: 0,
bit_pos: 0,
}
}
#[must_use]
pub const fn bit_position(&self) -> usize {
self.byte_pos * 8 + self.bit_pos as usize
}
pub fn seek(&mut self, bit_pos: usize) {
self.byte_pos = bit_pos / 8;
self.bit_pos = (bit_pos % 8) as u8;
}
fn peek_bits(&self, n: u8) -> AudioResult<u32> {
if n == 0 || n > 32 {
return Err(AudioError::InvalidData("Invalid bit count".into()));
}
let mut result = 0u32;
let mut bits_read = 0u8;
let mut byte_pos = self.byte_pos;
let mut bit_pos = self.bit_pos;
while bits_read < n {
if byte_pos >= self.data.len() {
return Err(AudioError::NeedMoreData);
}
let bits_available = 8 - bit_pos;
let bits_to_read = (n - bits_read).min(bits_available);
let byte = self.data[byte_pos];
let mask = ((1u8 << bits_to_read) - 1) << (bits_available - bits_to_read);
let bits = (byte & mask) >> (bits_available - bits_to_read);
result = (result << bits_to_read) | u32::from(bits);
bits_read += bits_to_read;
bit_pos += bits_to_read;
if bit_pos >= 8 {
bit_pos = 0;
byte_pos += 1;
}
}
Ok(result)
}
pub fn skip_bits(&mut self, n: usize) -> AudioResult<()> {
let total_bits = n + self.bit_pos as usize;
self.byte_pos += total_bits / 8;
self.bit_pos = (total_bits % 8) as u8;
if self.byte_pos > self.data.len() {
return Err(AudioError::NeedMoreData);
}
Ok(())
}
pub fn read_bits(&mut self, n: u8) -> AudioResult<u32> {
let result = self.peek_bits(n)?;
self.skip_bits(n as usize)?;
Ok(result)
}
pub fn decode(&mut self, table: u8, linbits: u8) -> AudioResult<HuffPair> {
let table_data = get_huffman_table(table)?;
let mut code = 0u32;
for len in 1..=16 {
let bit = self.read_bits(1)? as u32;
code = (code << 1) | bit;
for entry in table_data {
if entry.len == len && self.matches_code(code, entry, len) {
let mut x = i16::from(entry.x);
let mut y = i16::from(entry.y);
if linbits > 0 {
if x == 15 {
let extra = self.read_bits(linbits)? as i16;
x += extra;
}
if y == 15 {
let extra = self.read_bits(linbits)? as i16;
y += extra;
}
}
if x != 0 {
let sign = self.read_bits(1)?;
if sign != 0 {
x = -x;
}
}
if y != 0 {
let sign = self.read_bits(1)?;
if sign != 0 {
y = -y;
}
}
return Ok(HuffPair { x, y });
}
}
}
Err(AudioError::InvalidData("Invalid Huffman code".into()))
}
fn matches_code(&self, _code: u32, _entry: &HuffEntry, _len: u8) -> bool {
true
}
pub fn decode_quad(&mut self) -> AudioResult<[i16; 4]> {
let code = self.read_bits(4)?;
let mut values = [0i16; 4];
for (i, value) in values.iter_mut().enumerate() {
if (code & (1 << (3 - i))) != 0 {
let sign = self.read_bits(1)?;
*value = if sign != 0 { -1 } else { 1 };
}
}
Ok(values)
}
pub fn byte_align(&mut self) {
if self.bit_pos != 0 {
self.bit_pos = 0;
self.byte_pos += 1;
}
}
}
fn get_huffman_table(table: u8) -> AudioResult<&'static [HuffEntry]> {
match table {
0 => Ok(&TABLE_0),
1 => Ok(&TABLE_1),
_ => {
Ok(&TABLE_0)
}
}
}
const TABLE_0: [HuffEntry; 1] = [HuffEntry { len: 0, x: 0, y: 0 }];
const TABLE_1: [HuffEntry; 4] = [
HuffEntry { len: 1, x: 0, y: 0 },
HuffEntry { len: 3, x: 1, y: 1 },
HuffEntry { len: 3, x: 0, y: 1 },
HuffEntry { len: 3, x: 1, y: 0 },
];
#[must_use]
pub const fn get_linbits(table: u8) -> u8 {
match table {
16..=18 => 1,
19..=21 => 2,
22 | 23 => 3,
24 => 4,
25 | 26 => 6,
27 | 28 => 8,
29 | 30 => 10,
31 => 13,
_ => 0,
}
}
#[must_use]
pub const fn uses_linbits(table: u8) -> bool {
table >= 16 && table <= 31
}
#[must_use]
pub const fn get_max_value(table: u8) -> u8 {
match table {
0 => 0,
1 => 1,
2 | 3 => 2,
4..=6 => 3,
7..=9 => 5,
10..=12 => 7,
13..=15 => 15,
_ => 15,
}
}