use crate::lossless::bit_io::writer::BitWriter;
use crate::lossless::constants::{
CODE_LENGTH_CODE_ORDER, CODE_LENGTH_CODES, CODE_LENGTH_EXTRA_BITS, DEFAULT_CODE_LENGTH,
};
use crate::lossless::huffman::build::build_code_lengths;
use crate::lossless::huffman::canonical::emit_codes;
use crate::lossless::prelude::*;
const MIN_CODE_LENGTH_CODES: usize = 4;
const CODE_LENGTH_CODE_LIMIT: u32 = 7;
const SIMPLE_SYMBOL_LIMIT: u32 = 256;
struct Token {
code: u32,
extra: u32,
}
pub(crate) fn write_huffman_code(bw: &mut BitWriter, lengths: &[u32]) {
let mut count = 0usize;
let mut sym = [0u32; 2];
for (index, &len) in lengths.iter().enumerate() {
if len != 0 {
if count < 2 {
sym[count] = u32::try_from(index).unwrap_or(u32::MAX);
}
count += 1;
if count == 3 {
break; }
}
}
if count <= 2 && sym[0] < SIMPLE_SYMBOL_LIMIT && sym[1] < SIMPLE_SYMBOL_LIMIT {
write_simple_code(bw, count, sym[0], sym[1]);
} else {
write_full_code(bw, lengths);
}
}
fn write_simple_code(bw: &mut BitWriter, count: usize, sym0: u32, sym1: u32) {
bw.write_bits(1, 1); bw.write_bits(u32::from(count == 2), 1); if sym0 <= 1 {
bw.write_bits(0, 1); bw.write_bits(sym0, 1);
} else {
bw.write_bits(1, 1); bw.write_bits(sym0, 8);
}
if count == 2 {
bw.write_bits(sym1, 8);
}
}
fn write_full_code(bw: &mut BitWriter, lengths: &[u32]) {
bw.write_bits(0, 1);
let tokens = compress_code_lengths(lengths);
let mut hist = [0u32; CODE_LENGTH_CODES];
for token in &tokens {
hist[token.code as usize] += 1;
}
let cl_lengths = build_code_lengths(&hist, CODE_LENGTH_CODE_LIMIT);
let mut codes_to_store = CODE_LENGTH_CODES;
while codes_to_store > MIN_CODE_LENGTH_CODES
&& cl_lengths[usize::from(CODE_LENGTH_CODE_ORDER[codes_to_store - 1])] == 0
{
codes_to_store -= 1;
}
bw.write_bits(
u32::try_from(codes_to_store - MIN_CODE_LENGTH_CODES).unwrap_or(0),
4,
);
for &order in CODE_LENGTH_CODE_ORDER.iter().take(codes_to_store) {
bw.write_bits(cl_lengths[usize::from(order)], 3);
}
let cl_codes = emit_codes(&cl_lengths);
bw.write_bits(0, 1);
for token in &tokens {
let (code, emit_len) = cl_codes[token.code as usize];
bw.write_bits(code, emit_len);
if token.code >= 16 {
let width = u32::from(CODE_LENGTH_EXTRA_BITS[(token.code - 16) as usize]);
bw.write_bits(token.extra, width);
}
}
}
#[must_use]
fn compress_code_lengths(lengths: &[u32]) -> Vec<Token> {
let mut tokens = Vec::new();
let mut prev_value = DEFAULT_CODE_LENGTH;
let mut i = 0usize;
while i < lengths.len() {
let value = lengths[i];
let mut k = i + 1;
while k < lengths.len() && lengths[k] == value {
k += 1;
}
let runs = u32::try_from(k - i).unwrap_or(u32::MAX);
if value == 0 {
code_repeated_zeros(runs, &mut tokens);
} else {
code_repeated_values(runs, value, prev_value, &mut tokens);
prev_value = value;
}
i = k;
}
tokens
}
fn code_repeated_zeros(mut repetitions: u32, tokens: &mut Vec<Token>) {
while repetitions >= 1 {
if repetitions < 3 {
for _ in 0..repetitions {
tokens.push(Token { code: 0, extra: 0 });
}
return;
}
if repetitions < 11 {
tokens.push(Token {
code: 17,
extra: repetitions - 3,
});
return;
}
if repetitions < 139 {
tokens.push(Token {
code: 18,
extra: repetitions - 11,
});
return;
}
tokens.push(Token {
code: 18,
extra: 0x7f, });
repetitions -= 138;
}
}
fn code_repeated_values(
mut repetitions: u32,
value: u32,
prev_value: u32,
tokens: &mut Vec<Token>,
) {
if value != prev_value {
tokens.push(Token {
code: value,
extra: 0,
});
repetitions -= 1;
}
while repetitions >= 1 {
if repetitions < 3 {
for _ in 0..repetitions {
tokens.push(Token {
code: value,
extra: 0,
});
}
return;
}
if repetitions < 7 {
tokens.push(Token {
code: 16,
extra: repetitions - 3,
});
return;
}
tokens.push(Token {
code: 16,
extra: 3, });
repetitions -= 6;
}
}
#[cfg(test)]
mod tests {
use super::write_huffman_code;
use crate::lossless::bit_io::reader::BitReader;
use crate::lossless::bit_io::writer::BitWriter;
use crate::lossless::huffman::canonical::emit_codes;
use crate::lossless::huffman::decode::read_huffman_code;
const SENTINEL: u32 = 0b1_0110;
const SENTINEL_BITS: u32 = 5;
fn round_trip(lengths: &[u32]) {
let mut bw = BitWriter::new();
write_huffman_code(&mut bw, lengths);
let codes = emit_codes(lengths);
for (sym, &len) in lengths.iter().enumerate() {
if len != 0 {
let (code, emit_len) = codes[sym];
bw.write_bits(code, emit_len);
}
}
bw.write_bits(SENTINEL, SENTINEL_BITS);
let bytes = bw.into_bytes();
let mut br = BitReader::new(&bytes);
let table = read_huffman_code(&mut br, lengths.len())
.expect("write_huffman_code must emit a decodable prefix code");
for (sym, &len) in lengths.iter().enumerate() {
if len != 0 {
assert_eq!(
table.read_symbol(&mut br) as usize,
sym,
"symbol {sym} must round-trip"
);
}
}
assert_eq!(
br.read_bits(SENTINEL_BITS),
SENTINEL,
"reader must consume exactly the bits the writer produced"
);
}
#[test]
fn single_symbol_costs_zero_bits_per_occurrence() {
let mut lengths = vec![0u32; 32];
lengths[5] = 1;
round_trip(&lengths);
}
#[test]
fn two_symbols_one_bit_first_symbol() {
let mut lengths = vec![0u32; 16];
lengths[1] = 1;
lengths[5] = 1;
round_trip(&lengths);
}
#[test]
fn two_symbols_eight_bit_first_symbol() {
let mut lengths = vec![0u32; 256];
lengths[3] = 1;
lengths[200] = 1;
round_trip(&lengths);
}
#[test]
fn symbol_at_or_above_256_forces_full_form() {
let mut lengths = vec![0u32; 280];
lengths[5] = 1;
lengths[260] = 1;
round_trip(&lengths);
}
#[test]
fn mixed_lengths_use_the_full_form() {
round_trip(&[1, 2, 3, 3]);
}
#[test]
fn green_alphabet_with_trailing_zeros() {
let mut lengths = vec![0u32; 280];
lengths[..4].copy_from_slice(&[1, 2, 3, 3]);
round_trip(&lengths);
}
#[test]
fn all_length_eight_is_a_single_token_kind() {
round_trip(&[8u32; 256]);
}
#[test]
fn all_zero_alphabet_decodes_symbol_zero_with_no_bits() {
let lengths = [0u32; 40];
let mut bw = BitWriter::new();
write_huffman_code(&mut bw, &lengths);
bw.write_bits(SENTINEL, SENTINEL_BITS); let bytes = bw.into_bytes();
let mut br = BitReader::new(&bytes);
let table = read_huffman_code(&mut br, lengths.len()).expect("simple code");
assert_eq!(table.read_symbol(&mut br), 0);
assert_eq!(br.read_bits(SENTINEL_BITS), SENTINEL);
}
#[derive(Default)]
struct BitBuf {
bytes: Vec<u8>,
acc: u32,
n: u32,
}
impl BitBuf {
fn put(&mut self, value: u32, bits: u32) {
self.acc |= value << self.n;
self.n += bits;
while self.n >= 8 {
self.bytes.push((self.acc & 0xff) as u8);
self.acc >>= 8;
self.n -= 8;
}
}
fn finish(mut self) -> Vec<u8> {
if self.n > 0 {
self.bytes.push((self.acc & 0xff) as u8);
}
self.bytes
}
}
fn put_simple_code(b: &mut BitBuf, symbol: u32) {
b.put(1, 1); b.put(0, 1); if symbol <= 1 {
b.put(0, 1); b.put(symbol, 1);
} else {
b.put(1, 1); b.put(symbol, 8);
}
}
#[test]
fn single_symbol_at_index_256_forces_full_form() {
let mut lengths = vec![0u32; 257];
lengths[256] = 1;
round_trip(&lengths);
}
#[test]
fn second_symbol_at_index_256_forces_full_form() {
let mut lengths = vec![0u32; 257];
lengths[5] = 1;
lengths[256] = 1;
round_trip(&lengths);
}
#[test]
fn simple_single_symbol_matches_decoder_oracle() {
for symbol in [0u32, 1, 5, 200] {
let mut lengths = vec![0u32; 256];
lengths[symbol as usize] = 1;
let mut bw = BitWriter::new();
write_huffman_code(&mut bw, &lengths);
let ours = bw.into_bytes();
let mut oracle = BitBuf::default();
put_simple_code(&mut oracle, symbol);
assert_eq!(ours, oracle.finish(), "symbol {symbol}");
}
}
}
#[cfg(test)]
mod proptests {
use super::write_huffman_code;
use crate::lossless::bit_io::reader::BitReader;
use crate::lossless::bit_io::writer::BitWriter;
use crate::lossless::huffman::build::build_code_lengths;
use crate::lossless::huffman::canonical::emit_codes;
use crate::lossless::huffman::decode::read_huffman_code;
use proptest::prelude::*;
proptest! {
#[test]
fn any_built_code_round_trips(
weights in proptest::collection::vec(0u32..12, 1..300usize)
) {
let lengths = build_code_lengths(&weights, 15);
let mut bw = BitWriter::new();
write_huffman_code(&mut bw, &lengths);
let codes = emit_codes(&lengths);
for (sym, &len) in lengths.iter().enumerate() {
if len != 0 {
let (code, emit_len) = codes[sym];
bw.write_bits(code, emit_len);
}
}
let bytes = bw.into_bytes();
let mut br = BitReader::new(&bytes);
let table = read_huffman_code(&mut br, lengths.len());
prop_assert!(table.is_some());
let table = table.expect("checked is_some above");
for (sym, &len) in lengths.iter().enumerate() {
if len != 0 {
prop_assert_eq!(table.read_symbol(&mut br) as usize, sym);
}
}
}
}
}