use pith_digest::{Error, Result};
pub const ZIGZAG: [u8; 64] = [
0, 1, 8, 16, 9, 2, 3, 10, 17, 24, 32, 25, 18, 11, 4, 5, 12, 19, 26, 33, 40, 48, 41, 34, 27, 20, 13, 6, 7, 14, 21, 28, 35, 42, 49, 56, 57, 50, 43, 36, 29, 22, 15, 23, 30, 37, 44, 51, 58, 59, 52, 45, 38, 31, 39, 46, 53, 60, 61, 54, 47, 55, 62, 63,
];
#[derive(Clone, Debug)]
pub struct Table {
counts: [u8; 16],
symbols: Vec<u8>,
mincode: [i32; 16],
maxcode: [i32; 16],
valptr: [i32; 16],
}
impl Table {
pub fn build(counts: [u8; 16], symbols: Vec<u8>) -> Result<Self> {
let total: usize = counts.iter().map(|&c| c as usize).sum();
if total > 256 {
return Err(Error::BadValue("DHT: more than 256 symbols"));
}
if symbols.len() != total {
return Err(Error::BadValue("DHT: symbol count mismatch"));
}
let mut mincode = [0i32; 16];
let mut maxcode = [-1i32; 16];
let mut valptr = [0i32; 16];
let mut code: i32 = 0;
let mut sym: i32 = 0;
for l in 0..16 {
let n = counts[l] as i32;
if n > 0 {
valptr[l] = sym;
mincode[l] = code;
code += n;
sym += n;
maxcode[l] = code - 1;
}
if code > (1i32 << (l + 1)) {
return Err(Error::BadValue("DHT: oversubscribed table"));
}
code <<= 1;
}
Ok(Table {
counts,
symbols,
mincode,
maxcode,
valptr,
})
}
pub fn decode(&self, bits: &mut crate::bits::Bits<'_>) -> Result<u8> {
let mut code = bits.bits(1)? as i32;
let mut l = 0usize;
while code > self.maxcode[l] {
code = (code << 1) | bits.bits(1)? as i32;
l += 1;
if l >= 16 {
return Err(Error::BadValue("Huffman code longer than 16 bits"));
}
}
let idx = self.valptr[l] + (code - self.mincode[l]);
let sym = *self
.symbols
.get(idx as usize)
.ok_or(Error::BadValue("Huffman symbol index out of range"))?;
let _ = self.counts;
Ok(sym)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::bits::Bits;
#[test]
fn single_symbol_table() {
let mut counts = [0u8; 16];
counts[0] = 1;
let t = Table::build(counts, alloc::vec![0x42]).unwrap();
let data = [0x00]; let mut b = Bits::new(&data, 0);
assert_eq!(t.decode(&mut b).unwrap(), 0x42);
assert_eq!(t.decode(&mut b).unwrap(), 0x42);
}
#[test]
fn oversubscribed_table_rejected() {
let mut counts = [0u8; 16];
counts[0] = 3; assert!(Table::build(counts, alloc::vec![1, 2, 3]).is_err());
let mut counts = [0u8; 16];
counts[0] = 2;
counts[1] = 3; assert!(Table::build(counts, alloc::vec![1, 2, 3, 4, 5]).is_err());
}
#[test]
fn standard_ac_table_first_code() {
let mut counts = [0u8; 16];
counts[1] = 2; counts[2] = 1; let syms = alloc::vec![0xaa, 0xbb, 0xcc];
let t = Table::build(counts, syms).unwrap();
let data = [0b0001_1000];
let mut b = Bits::new(&data, 0);
assert_eq!(t.decode(&mut b).unwrap(), 0xaa);
assert_eq!(t.decode(&mut b).unwrap(), 0xbb);
assert_eq!(t.decode(&mut b).unwrap(), 0xcc);
}
#[test]
fn overlong_code_errors() {
let mut counts = [0u8; 16];
counts[15] = 1; let t = Table::build(counts, alloc::vec![0x77]).unwrap();
let data = [0xff; 4]; let mut b = Bits::new(&data, 0);
assert!(t.decode(&mut b).is_err());
}
#[test]
fn oversubscribed_symbol_budget() {
let mut counts = [0u8; 16];
counts[0] = 200;
counts[1] = 200; let err = Table::build(counts, alloc::vec![0; 400]).expect_err("budget");
assert!(err.to_string().contains("more than 256 symbols"));
}
#[test]
fn symbol_count_mismatch() {
let mut counts = [0u8; 16];
counts[0] = 2;
let err = Table::build(counts, alloc::vec![0x42]).expect_err("mismatch");
assert!(err.to_string().contains("symbol count mismatch"));
}
}