use otf_pixels_core::{PixelsError, Result};
const LOOKUP_BITS: u32 = 8;
#[derive(Debug, Clone)]
pub struct HuffmanTable {
lookup: [u16; 1 << LOOKUP_BITS],
mincode: [i32; 17],
maxcode: [i32; 17],
valptr: [i32; 17],
values: Vec<u8>,
}
impl HuffmanTable {
pub fn new(counts: &[u8; 16], values: Vec<u8>) -> Result<Self> {
let total: usize = counts.iter().map(|&c| c as usize).sum();
if total != values.len() {
return Err(PixelsError::malformed(
"jpeg",
format!(
"Huffman table declares {total} codes but carries {} symbols",
values.len()
),
));
}
if total == 0 {
return Err(PixelsError::malformed("jpeg", "Huffman table is empty"));
}
let mut mincode = [0_i32; 17];
let mut maxcode = [-1_i32; 17];
let mut valptr = [0_i32; 17];
let mut lookup = [0_u16; 1 << LOOKUP_BITS];
let mut code: i32 = 0;
let mut index: i32 = 0;
for length in 1..=16_usize {
let count = i32::from(*counts.get(length - 1).unwrap_or(&0));
if let (Some(min), Some(max), Some(ptr)) = (
mincode.get_mut(length),
maxcode.get_mut(length),
valptr.get_mut(length),
) {
*min = code;
*ptr = index;
*max = if count > 0 { code + count - 1 } else { -1 };
}
if length <= LOOKUP_BITS as usize {
let shift = LOOKUP_BITS as usize - length;
for step in 0..count {
let Some(&symbol) = values.get((index + step) as usize) else {
break;
};
let prefix = ((code + step) as usize) << shift;
for slot in prefix..prefix + (1 << shift) {
if let Some(entry) = lookup.get_mut(slot) {
*entry = ((length as u16) << 8) | u16::from(symbol);
}
}
}
}
code += count;
index += count;
if code > (1 << length) {
return Err(PixelsError::malformed(
"jpeg",
format!("Huffman table is over-subscribed at code length {length}"),
));
}
code <<= 1;
}
Ok(Self {
lookup,
mincode,
maxcode,
valptr,
values,
})
}
#[must_use]
pub fn lookup(&self, prefix: u8) -> Option<(u32, u8)> {
match self.lookup.get(prefix as usize).copied().unwrap_or(0) {
0 => None,
entry => Some((u32::from(entry >> 8), (entry & 0xFF) as u8)),
}
}
#[must_use]
pub fn resolve(&self, length: usize, code: i32) -> Option<u8> {
let max = *self.maxcode.get(length)?;
if max < 0 || code > max {
return None;
}
let min = *self.mincode.get(length)?;
let base = *self.valptr.get(length)?;
let at = base.checked_add(code.checked_sub(min)?)?;
self.values.get(usize::try_from(at).ok()?).copied()
}
#[must_use]
pub fn max_length(&self) -> usize {
(1..=16)
.rev()
.find(|&l| self.maxcode.get(l).is_some_and(|&m| m >= 0))
.unwrap_or(16)
}
}
#[derive(Debug, Clone)]
pub struct HuffmanEncoder {
codes: [(u16, u8); 256],
}
impl HuffmanEncoder {
pub fn new(counts: &[u8; 16], values: &[u8]) -> Result<Self> {
HuffmanTable::new(counts, values.to_vec())?;
let mut codes = [(0_u16, 0_u8); 256];
let mut code: u16 = 0;
let mut index = 0_usize;
for length in 1..=16_usize {
let count = usize::from(*counts.get(length - 1).unwrap_or(&0));
for step in 0..count {
if let Some(&symbol) = values.get(index + step) {
if let Some(slot) = codes.get_mut(symbol as usize) {
*slot = (code, length as u8);
}
}
code = code.wrapping_add(1);
}
index += count;
code <<= 1;
}
Ok(Self { codes })
}
#[must_use]
pub fn code(&self, symbol: u8) -> Option<(u32, u32)> {
match self.codes.get(symbol as usize).copied() {
Some((_, 0)) | None => None,
Some((code, length)) => Some((u32::from(code), u32::from(length))),
}
}
}
#[cfg(test)]
#[allow(
clippy::unwrap_used,
clippy::expect_used,
clippy::indexing_slicing,
reason = "tests operate on known-good values and assert shapes directly"
)]
mod tests {
use super::*;
use otf_pixels_core::ErrorCode;
fn simple() -> HuffmanTable {
let mut counts = [0_u8; 16];
counts[0] = 1;
counts[1] = 1;
counts[2] = 1;
HuffmanTable::new(&counts, vec![b'a', b'b', b'c']).unwrap()
}
#[test]
fn canonical_codes_resolve_to_their_symbols() {
let table = simple();
assert_eq!(table.resolve(1, 0b0), Some(b'a'));
assert_eq!(table.resolve(2, 0b10), Some(b'b'));
assert_eq!(table.resolve(3, 0b110), Some(b'c'));
assert_eq!(table.resolve(3, 0b111), None);
assert_eq!(table.resolve(4, 0), None);
}
#[test]
fn the_lookup_table_agrees_with_the_slow_path() {
let table = simple();
assert_eq!(table.lookup(0b0000_0000), Some((1, b'a')));
assert_eq!(table.lookup(0b0111_1111), Some((1, b'a')));
assert_eq!(table.lookup(0b1011_1111), Some((2, b'b')));
assert_eq!(table.lookup(0b1101_0101), Some((3, b'c')));
assert_eq!(table.lookup(0b1110_0000), None);
}
#[test]
fn long_codes_fall_out_of_the_lookup_table() {
let counts = [1_u8; 16];
let values: Vec<u8> = (0..16).collect();
let table = HuffmanTable::new(&counts, values).unwrap();
assert_eq!(table.max_length(), 16);
assert_eq!(table.resolve(9, 0b1_1111_1110), Some(8));
assert_eq!(table.lookup(0b1111_1111), None);
}
#[test]
fn the_encoding_and_decoding_tables_are_inverses() {
use crate::tables::{
CHROMA_AC_COUNTS, CHROMA_AC_VALUES, CHROMA_DC_COUNTS, CHROMA_DC_VALUES, LUMA_AC_COUNTS,
LUMA_AC_VALUES, LUMA_DC_COUNTS, LUMA_DC_VALUES,
};
for (counts, values) in [
(&LUMA_DC_COUNTS, &LUMA_DC_VALUES[..]),
(&CHROMA_DC_COUNTS, &CHROMA_DC_VALUES[..]),
(&LUMA_AC_COUNTS, &LUMA_AC_VALUES[..]),
(&CHROMA_AC_COUNTS, &CHROMA_AC_VALUES[..]),
] {
let encoder = HuffmanEncoder::new(counts, values).unwrap();
let decoder = HuffmanTable::new(counts, values.to_vec()).unwrap();
for &symbol in values {
let (code, length) = encoder.code(symbol).expect("every symbol has a code");
assert_eq!(
decoder.resolve(length as usize, code as i32),
Some(symbol),
"symbol {symbol:#04x} coded as {code:#x}/{length} bits"
);
}
let undefined = (0..=255_u8).find(|s| !values.contains(s));
if let Some(symbol) = undefined {
assert_eq!(encoder.code(symbol), None, "{symbol:#04x}");
}
}
}
#[test]
fn mismatched_counts_and_symbols_are_malformed() {
let mut counts = [0_u8; 16];
counts[0] = 2;
assert_eq!(
HuffmanTable::new(&counts, vec![b'a']).unwrap_err().code(),
ErrorCode::Malformed
);
assert_eq!(
HuffmanTable::new(&[0; 16], vec![]).unwrap_err().code(),
ErrorCode::Malformed
);
}
#[test]
fn over_subscribed_tables_are_rejected() {
let mut counts = [0_u8; 16];
counts[0] = 3;
assert_eq!(
HuffmanTable::new(&counts, vec![b'a', b'b', b'c'])
.unwrap_err()
.code(),
ErrorCode::Malformed
);
}
#[test]
fn a_full_table_of_short_codes_is_accepted() {
let mut counts = [0_u8; 16];
counts[7] = 255;
let values: Vec<u8> = (0..255).collect();
let table = HuffmanTable::new(&counts, values).unwrap();
assert_eq!(table.lookup(0), Some((8, 0)));
assert_eq!(table.lookup(254), Some((8, 254)));
}
}