use crate::bits::ReverseBitReader;
use crate::error::Error;
use crate::fse;
const MAX_TABLE_LOG: u32 = 12;
const MAX_WEIGHT: u32 = MAX_TABLE_LOG;
const WEIGHT_FSE_MAX_LOG: u32 = 6;
#[derive(Clone, Copy)]
struct HufEntry {
symbol: u8,
nb_bits: u8,
}
#[derive(Clone)]
pub(crate) struct HuffmanTable {
table_log: u32,
entries: Vec<HufEntry>,
}
pub(crate) fn read_table(src: &[u8]) -> Result<(HuffmanTable, usize), Error> {
let (weights, table_log, consumed) = read_weights(src)?;
Ok((build(&weights, table_log)?, consumed))
}
pub(crate) fn read_weights(src: &[u8]) -> Result<(Vec<u8>, u32, usize), Error> {
let header = *src
.first()
.ok_or(Error::Corrupted("missing huffman header"))?;
let (mut weights, consumed) = if header < 128 {
let payload = src
.get(1..1 + header as usize)
.ok_or(Error::Corrupted("huffman weights overrun input"))?;
let nc = fse::read_ncount(payload, 255, WEIGHT_FSE_MAX_LOG)?;
let table = fse::build_dtable(&nc.counts, nc.table_log)?;
let stream = &payload[nc.bytes_consumed..];
let weights = fse::decode_interleaved(&table, stream, 255)?;
(weights, 1 + header as usize)
} else {
let n = (header - 127) as usize;
let n_bytes = n.div_ceil(2);
let payload = src
.get(1..1 + n_bytes)
.ok_or(Error::Corrupted("huffman weights overrun input"))?;
let mut weights = Vec::with_capacity(n);
for i in 0..n {
let b = payload[i / 2];
weights.push(if i % 2 == 0 { b >> 4 } else { b & 0x0F });
}
(weights, 1 + n_bytes)
};
let mut total: u32 = 0;
for &w in &weights {
if u32::from(w) > MAX_WEIGHT {
return Err(Error::Corrupted("huffman weight too large"));
}
if w > 0 {
total += 1u32 << (w - 1);
}
}
if total == 0 {
return Err(Error::Corrupted("huffman table has no weights"));
}
let table_log = 32 - total.leading_zeros();
if table_log > MAX_TABLE_LOG {
return Err(Error::Corrupted("huffman code lengths too long"));
}
let rest = (1u32 << table_log) - total;
if !rest.is_power_of_two() {
return Err(Error::Corrupted(
"huffman weights do not sum to a power of two",
));
}
let last_weight = rest.trailing_zeros() + 1;
weights.push(last_weight as u8);
if weights.len() > 256 {
return Err(Error::Corrupted("too many huffman symbols"));
}
Ok((weights, table_log, consumed))
}
fn build(weights: &[u8], table_log: u32) -> Result<HuffmanTable, Error> {
let table_size = 1usize << table_log;
let mut rank_count = [0usize; (MAX_WEIGHT + 1) as usize + 1];
for &w in weights {
rank_count[w as usize] += 1;
}
let mut rank_next = [0usize; (MAX_WEIGHT + 1) as usize + 1];
let mut cur = 0usize;
for w in 1..=MAX_WEIGHT as usize {
rank_next[w] = cur;
cur += rank_count[w] << (w - 1);
}
debug_assert_eq!(cur, table_size, "weight sum must fill the table exactly");
let mut entries = vec![
HufEntry {
symbol: 0,
nb_bits: 0
};
table_size
];
for (symbol, &w) in weights.iter().enumerate() {
if w == 0 {
continue;
}
let w = w as usize;
let len = 1usize << (w - 1);
let nb_bits = (table_log + 1 - w as u32) as u8;
for e in &mut entries[rank_next[w]..rank_next[w] + len] {
e.symbol = symbol as u8;
e.nb_bits = nb_bits;
}
rank_next[w] += len;
}
Ok(HuffmanTable { table_log, entries })
}
fn decode_stream_into(
table: &HuffmanTable,
src: &[u8],
count: usize,
out: &mut Vec<u8>,
) -> Result<(), Error> {
let mut br = ReverseBitReader::new(src)?;
let mask = table.entries.len() - 1;
for _ in 0..count {
let idx = br.peek(table.table_log) as usize & mask;
let e = table.entries[idx];
br.consume(u32::from(e.nb_bits));
out.push(e.symbol);
}
if !br.finished_exactly() {
return Err(Error::Corrupted("huffman stream not fully consumed"));
}
Ok(())
}
pub(crate) fn decode_single_stream(
table: &HuffmanTable,
src: &[u8],
regenerated_size: usize,
) -> Result<Vec<u8>, Error> {
let mut out = Vec::with_capacity(regenerated_size);
decode_stream_into(table, src, regenerated_size, &mut out)?;
Ok(out)
}
pub(crate) fn decode_four_streams(
table: &HuffmanTable,
src: &[u8],
regenerated_size: usize,
) -> Result<Vec<u8>, Error> {
let jump = src
.get(..6)
.ok_or(Error::Corrupted("missing huffman jump table"))?;
let s1 = u16::from_le_bytes([jump[0], jump[1]]) as usize;
let s2 = u16::from_le_bytes([jump[2], jump[3]]) as usize;
let s3 = u16::from_le_bytes([jump[4], jump[5]]) as usize;
let rest = &src[6..];
if s1 + s2 + s3 >= rest.len() {
return Err(Error::Corrupted("huffman jump table overruns input"));
}
let segment = regenerated_size.div_ceil(4);
let last = regenerated_size as i64 - 3 * segment as i64;
if last <= 0 {
return Err(Error::Corrupted("four-stream literals too short"));
}
let mut out = Vec::with_capacity(regenerated_size);
decode_stream_into(table, &rest[..s1], segment, &mut out)?;
decode_stream_into(table, &rest[s1..s1 + s2], segment, &mut out)?;
decode_stream_into(table, &rest[s1 + s2..s1 + s2 + s3], segment, &mut out)?;
decode_stream_into(table, &rest[s1 + s2 + s3..], last as usize, &mut out)?;
Ok(out)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn direct_weights_roundtrip() {
let desc = [127 + 3, 0x22, 0x20];
let (table, consumed) = read_table(&desc).unwrap();
assert_eq!(consumed, 3);
assert_eq!(table.table_log, 3);
assert!(table.entries.iter().all(|e| e.nb_bits == 2));
}
#[test]
fn rejects_invalid_weight_sum() {
let desc = [127 + 3, 0x22, 0x10];
assert!(read_table(&desc).is_err());
}
}