use crate::{Result, error::invalid};
use super::bits::BitReader;
pub(super) struct Huffman {
symbols: Vec<u16>,
lengths: Vec<u8>,
max_bits: u8,
}
impl Huffman {
pub(super) fn new(lengths: &[u8]) -> Result<Self> {
let max_bits = lengths.iter().copied().max().unwrap_or(0);
if max_bits == 0 || max_bits > 15 {
return Err(invalid("invalid Huffman code lengths"));
}
let mut counts = [0_u16; 16];
for length in lengths {
if *length > 15 {
return Err(invalid("Huffman code is longer than 15 bits"));
}
if *length != 0 {
counts[usize::from(*length)] += 1;
}
}
let mut left = 1_i32;
for count in counts.iter().skip(1) {
left = (left << 1) - i32::from(*count);
if left < 0 {
return Err(invalid("oversubscribed Huffman tree"));
}
}
let mut next = [0_u16; 16];
let mut code = 0_u16;
for bits in 1..=15 {
code = (code + counts[bits - 1]) << 1;
next[bits] = code;
}
let size = 1_usize << max_bits;
let mut symbols = vec![u16::MAX; size];
let mut table_lengths = vec![0_u8; size];
for (symbol, length) in lengths.iter().copied().enumerate() {
if length == 0 {
continue;
}
let canonical = next[usize::from(length)];
next[usize::from(length)] += 1;
let reversed = reverse_bits(canonical, length);
let repetitions = 1_usize << (max_bits - length);
for suffix in 0..repetitions {
let index = usize::from(reversed) | (suffix << length);
symbols[index] =
u16::try_from(symbol).map_err(|_| invalid("too many Huffman symbols"))?;
table_lengths[index] = length;
}
}
Ok(Self {
symbols,
lengths: table_lengths,
max_bits,
})
}
pub(super) fn decode(&self, reader: &mut BitReader<'_>) -> Result<u16> {
let index = usize::try_from(reader.peek(self.max_bits)?)
.map_err(|_| invalid("Huffman table index overflow"))?;
let length = self.lengths[index];
if length == 0 {
return Err(invalid("invalid Huffman code"));
}
reader.consume(length)?;
Ok(self.symbols[index])
}
}
fn reverse_bits(value: u16, length: u8) -> u16 {
value.reverse_bits() >> (u16::BITS - u32::from(length))
}
#[cfg(test)]
mod tests {
use super::Huffman;
use crate::inflate::bits::BitReader;
#[test]
fn decodes_lsb_first_codes() {
let table = Huffman::new(&[1, 2, 2]).unwrap();
let mut reader = BitReader::new(&[0b0000_0001, 0]);
assert_eq!(table.decode(&mut reader).unwrap(), 1);
}
#[test]
fn rejects_empty_long_and_oversubscribed_trees() {
assert!(Huffman::new(&[]).is_err());
assert!(Huffman::new(&[16]).is_err());
assert!(Huffman::new(&[1, 1, 1]).is_err());
}
}