1#![forbid(unsafe_code)]
2
3#[cfg(feature = "alloc")]
4use alloc::vec;
5#[cfg(feature = "alloc")]
6use alloc::vec::Vec;
7
8use super::primitives;
9use crate::huffman::{MAX_BITS, MAX_SYMBOL_VALUE};
10
11#[derive(Clone)]
12pub struct HuffmanEncodeTable {
13 codes: [u16; MAX_SYMBOL_VALUE + 1],
14 num_bits: [u8; MAX_SYMBOL_VALUE + 1],
15 weights: Vec<u8>,
16 max_symbol: u8,
17 table_log: u8,
18}
19
20#[cfg(feature = "alloc")]
21impl HuffmanEncodeTable {
22 pub fn from_data(data: &[u8]) -> Option<Self> {
23 if data.is_empty() {
24 return None;
25 }
26
27 let mut freqs = [0u32; MAX_SYMBOL_VALUE + 1];
28 let mut max_sym = 0u8;
29 for &b in data {
30 freqs[b as usize] += 1;
31 if b > max_sym {
32 max_sym = b;
33 }
34 }
35
36 let num_symbols = max_sym as usize + 1;
37 let active_count = freqs[..num_symbols].iter().filter(|&&f| f > 0).count();
38 if active_count < 2 {
39 return None;
40 }
41
42 if max_sym as usize > 128 {
43 return None;
44 }
45
46 let (weights, table_log) = compute_huffman_weights(&freqs, num_symbols)?;
47 let (codes, num_bits) = build_encode_codes(&weights, table_log);
48
49 Some(Self {
50 codes,
51 num_bits,
52 weights,
53 max_symbol: max_sym,
54 table_log,
55 })
56 }
57
58 pub fn from_decode_table(
59 decode_table: &[super::HuffmanDecodeEntry],
60 table_log: u8,
61 ) -> Option<Self> {
62 let table_size = 1usize << table_log;
63 if decode_table.len() < table_size {
64 return None;
65 }
66
67 let mut num_bits_per_sym = [0u8; MAX_SYMBOL_VALUE + 1];
68 let mut max_sym = 0u8;
69 let mut seen = [false; MAX_SYMBOL_VALUE + 1];
70 for entry in &decode_table[..table_size] {
71 let s = entry.symbol;
72 if !seen[s as usize] {
73 num_bits_per_sym[s as usize] = entry.num_bits;
74 seen[s as usize] = true;
75 if s > max_sym {
76 max_sym = s;
77 }
78 }
79 }
80
81 if max_sym as usize > 128 {
82 return None;
83 }
84
85 let num_symbols = max_sym as usize + 1;
86 let active_count = seen[..num_symbols].iter().filter(|&&s| s).count();
87 if active_count < 2 {
88 return None;
89 }
90
91 let mut weights = vec![0u8; num_symbols];
92 for s in 0..num_symbols {
93 if seen[s] {
94 weights[s] = table_log + 1 - num_bits_per_sym[s];
95 }
96 }
97
98 let (codes, num_bits) = build_encode_codes(&weights, table_log);
99
100 Some(Self {
101 codes,
102 num_bits,
103 weights,
104 max_symbol: max_sym,
105 table_log,
106 })
107 }
108
109 pub fn table_log(&self) -> u8 {
110 self.table_log
111 }
112
113 pub fn can_encode(&self, data: &[u8]) -> bool {
114 for &b in data {
115 if self.num_bits[b as usize] == 0 {
116 return false;
117 }
118 }
119 true
120 }
121
122 pub fn serialize_weights(&self) -> Vec<u8> {
123 let explicit = &self.weights[..self.max_symbol as usize];
124 let num_symbols = explicit.len();
125
126 let mut out = Vec::with_capacity(1 + num_symbols.div_ceil(2));
127 out.push((num_symbols + 127) as u8);
128 let num_bytes = num_symbols.div_ceil(2);
129 for i in 0..num_bytes {
130 let hi = explicit.get(i * 2).copied().unwrap_or(0);
131 let lo = explicit.get(i * 2 + 1).copied().unwrap_or(0);
132 out.push((hi << 4) | lo);
133 }
134 out
135 }
136
137 pub fn encode_single_stream(&self, data: &[u8]) -> Vec<u8> {
138 let mut buf = Vec::with_capacity(data.len() + 8);
139 self.encode_single_stream_into(data, &mut buf);
140 buf
141 }
142
143 pub fn encode_single_stream_into(&self, data: &[u8], buf: &mut Vec<u8>) {
144 let tl = self.table_log as usize;
145 let unroll: usize = (32usize).checked_div(tl).unwrap_or(1).max(2);
146
147 let mut bitstream = primitives::BitstreamScratch::new(buf, data.len() + 16);
148 let mut bits: u64 = 0;
149 let mut bits_used: u8 = 0;
150 let mut wpos: usize = 0;
151
152 macro_rules! flush_bits {
153 () => {
154 bitstream.flush(wpos, bits);
155 let nb = (bits_used >> 3) as usize;
156 wpos += nb;
157 bits >>= nb << 3;
158 bits_used &= 7;
159 };
160 }
161
162 let mut pos = data.len();
163 while pos >= unroll {
164 pos -= unroll;
165 for j in 0..unroll {
166 let b = data[pos + (unroll - 1 - j)];
167 let c = self.codes[b as usize] as u64;
168 let n = self.num_bits[b as usize];
169 bits |= c << bits_used;
170 bits_used += n;
171 }
172 if bits_used >= 32 {
173 flush_bits!();
174 }
175 }
176 while pos > 0 {
177 pos -= 1;
178 let b = data[pos];
179 let c = self.codes[b as usize] as u64;
180 let n = self.num_bits[b as usize];
181 bits |= c << bits_used;
182 bits_used += n;
183 if bits_used >= 32 {
184 flush_bits!();
185 }
186 }
187
188 bits |= 1u64 << bits_used;
189 bits_used += 1;
190 while bits_used > 0 {
191 bitstream.write_byte(wpos, bits as u8);
192 wpos += 1;
193 bits >>= 8;
194 bits_used = bits_used.saturating_sub(8);
195 }
196 bitstream.finish(wpos);
197 }
198
199 pub fn encode_4_streams(&self, data: &[u8]) -> Vec<u8> {
200 let mut out = Vec::new();
201 self.encode_4_streams_into(data, &mut out, &mut Vec::new());
202 out
203 }
204
205 pub fn encode_4_streams_into(&self, data: &[u8], out: &mut Vec<u8>, stream_buf: &mut Vec<u8>) {
206 let seg = data.len().div_ceil(4);
207 let s1 = &data[..seg.min(data.len())];
208 let s2 = &data[seg.min(data.len())..(seg * 2).min(data.len())];
209 let s3 = &data[(seg * 2).min(data.len())..(seg * 3).min(data.len())];
210 let s4 = &data[(seg * 3).min(data.len())..];
211
212 out.clear();
213 out.extend_from_slice(&[0u8; 6]);
214
215 self.encode_single_stream_into(s1, stream_buf);
216 let e1_len = stream_buf.len();
217 out.extend_from_slice(stream_buf);
218
219 self.encode_single_stream_into(s2, stream_buf);
220 let e2_len = stream_buf.len();
221 out.extend_from_slice(stream_buf);
222
223 self.encode_single_stream_into(s3, stream_buf);
224 let e3_len = stream_buf.len();
225 out.extend_from_slice(stream_buf);
226
227 self.encode_single_stream_into(s4, stream_buf);
228 out.extend_from_slice(stream_buf);
229
230 out[0..2].copy_from_slice(&(e1_len as u16).to_le_bytes());
231 out[2..4].copy_from_slice(&(e2_len as u16).to_le_bytes());
232 out[4..6].copy_from_slice(&(e3_len as u16).to_le_bytes());
233 }
234
235 pub fn compressed_size_single(&self, data: &[u8]) -> usize {
236 let total_bits: usize = data
237 .iter()
238 .map(|&b| self.num_bits[b as usize] as usize)
239 .sum();
240 (total_bits + 8) / 8
241 }
242}
243
244fn compute_huffman_weights(freqs: &[u32], num_symbols: usize) -> Option<(Vec<u8>, u8)> {
245 use alloc::collections::BinaryHeap;
246 use core::cmp::Reverse;
247
248 let active: Vec<(u64, usize)> = freqs[..num_symbols]
249 .iter()
250 .enumerate()
251 .filter(|(_, f)| **f > 0)
252 .map(|(s, &f)| (f as u64, s))
253 .collect();
254
255 if active.len() < 2 {
256 return None;
257 }
258
259 let n = active.len();
260
261 let max_nodes = 2 * n;
262 let mut parent = vec![usize::MAX; max_nodes];
263
264 let mut heap: BinaryHeap<Reverse<(u64, usize)>> = BinaryHeap::with_capacity(n);
265 for (i, &(f, _)) in active.iter().enumerate() {
266 heap.push(Reverse((f, i)));
267 }
268
269 for next_id in n..n + (n - 1) {
270 let Reverse((f1, n1)) = heap.pop().unwrap();
271 let Reverse((f2, n2)) = heap.pop().unwrap();
272 parent[n1] = next_id;
273 parent[n2] = next_id;
274 heap.push(Reverse((f1 + f2, next_id)));
275 }
276
277 let mut bit_lengths = vec![0u8; num_symbols];
278 for (i, &(_, sym)) in active.iter().enumerate().take(n) {
279 let mut depth = 0u8;
280 let mut node = i;
281 while parent[node] != usize::MAX {
282 depth += 1;
283 node = parent[node];
284 }
285 bit_lengths[sym] = depth;
286 }
287
288 let max_bl = *bit_lengths.iter().max().unwrap();
289 if max_bl == 0 || max_bl > MAX_BITS {
290 return None;
291 }
292
293 let table_log = max_bl;
294 let mut weights = vec![0u8; num_symbols];
295 for (s, &bl) in bit_lengths.iter().enumerate() {
296 if bl > 0 {
297 weights[s] = table_log + 1 - bl;
298 }
299 }
300
301 Some((weights, table_log))
302}
303
304fn build_encode_codes(
305 weights: &[u8],
306 table_log: u8,
307) -> ([u16; MAX_SYMBOL_VALUE + 1], [u8; MAX_SYMBOL_VALUE + 1]) {
308 let mut codes = [0u16; MAX_SYMBOL_VALUE + 1];
309 let mut num_bits = [0u8; MAX_SYMBOL_VALUE + 1];
310
311 let max_w = table_log + 1;
312 let mut rank_count = [0u32; MAX_BITS as usize + 2];
313
314 for (s, &w) in weights.iter().enumerate() {
315 if w > 0 && w <= max_w {
316 num_bits[s] = table_log + 1 - w;
317 rank_count[w as usize] += 1;
318 }
319 }
320
321 let mut rank_start = [0u32; MAX_BITS as usize + 2];
322 let mut cumul = 0u32;
323 for w in 1..=max_w {
324 rank_start[w as usize] = cumul;
325 cumul += rank_count[w as usize] * (1u32 << (w - 1));
326 }
327
328 for (s, &w) in weights.iter().enumerate() {
329 if w == 0 {
330 continue;
331 }
332 let start = rank_start[w as usize];
333 codes[s] = (start >> (w - 1)) as u16;
334 rank_start[w as usize] += 1u32 << (w - 1);
335 }
336
337 (codes, num_bits)
338}
339
340#[cfg(test)]
341mod tests {
342 use super::*;
343
344 #[test]
345 fn from_decode_table_roundtrip() {
346 let data = b"hello world hello world hello world!";
347 let original = HuffmanEncodeTable::from_data(data).unwrap();
348 let weights_raw = original.serialize_weights();
349
350 let (parsed_weights, _) =
351 crate::huffman::weights::parse_huffman_weights(&weights_raw).unwrap();
352 let (decode_table, decode_log) =
353 crate::huffman::weights::build_huffman_decode_table(&parsed_weights).unwrap();
354
355 let rebuilt = HuffmanEncodeTable::from_decode_table(&decode_table, decode_log).unwrap();
356 assert_eq!(original.table_log, rebuilt.table_log);
357 assert_eq!(original.max_symbol, rebuilt.max_symbol);
358 assert_eq!(original.weights, rebuilt.weights);
359 assert_eq!(original.num_bits, rebuilt.num_bits);
360 assert_eq!(original.codes, rebuilt.codes);
361
362 let encoded = rebuilt.encode_single_stream(data);
363 let decoded = crate::huffman::decode::decode_single_stream(
364 &decode_table,
365 decode_log,
366 &encoded,
367 data.len(),
368 )
369 .unwrap();
370 assert_eq!(decoded, data);
371 }
372
373 #[test]
374 fn from_decode_table_skewed() {
375 let mut data = vec![0u8; 900];
376 data.extend(vec![1u8; 80]);
377 data.extend(vec![2u8; 15]);
378 data.extend(vec![3u8; 5]);
379 let original = HuffmanEncodeTable::from_data(&data).unwrap();
380 let weights_raw = original.serialize_weights();
381
382 let (parsed_weights, _) =
383 crate::huffman::weights::parse_huffman_weights(&weights_raw).unwrap();
384 let (decode_table, decode_log) =
385 crate::huffman::weights::build_huffman_decode_table(&parsed_weights).unwrap();
386
387 let rebuilt = HuffmanEncodeTable::from_decode_table(&decode_table, decode_log).unwrap();
388 assert_eq!(original.codes, rebuilt.codes);
389 assert_eq!(original.num_bits, rebuilt.num_bits);
390
391 let encoded = rebuilt.encode_single_stream(&data);
392 let decoded = crate::huffman::decode::decode_single_stream(
393 &decode_table,
394 decode_log,
395 &encoded,
396 data.len(),
397 )
398 .unwrap();
399 assert_eq!(decoded, data);
400 }
401
402 #[test]
403 fn roundtrip_simple() {
404 let data = b"hello world hello world hello world!";
405 let table = HuffmanEncodeTable::from_data(data).unwrap();
406 let weights_raw = table.serialize_weights();
407 let encoded = table.encode_single_stream(data);
408
409 let (parsed_weights, _) =
410 crate::huffman::weights::parse_huffman_weights(&weights_raw).unwrap();
411 let (decode_table, decode_log) =
412 crate::huffman::weights::build_huffman_decode_table(&parsed_weights).unwrap();
413 let decoded = crate::huffman::decode::decode_single_stream(
414 &decode_table,
415 decode_log,
416 &encoded,
417 data.len(),
418 )
419 .unwrap();
420 assert_eq!(decoded, data);
421 }
422
423 #[test]
424 fn roundtrip_4_streams() {
425 let data: Vec<u8> = b"ABCDEFGH".iter().cycle().take(1024).copied().collect();
426 let table = HuffmanEncodeTable::from_data(&data).unwrap();
427 let weights_raw = table.serialize_weights();
428 let encoded = table.encode_4_streams(&data);
429
430 let (parsed_weights, _) =
431 crate::huffman::weights::parse_huffman_weights(&weights_raw).unwrap();
432 let (decode_table, decode_log) =
433 crate::huffman::weights::build_huffman_decode_table(&parsed_weights).unwrap();
434 let decoded = crate::huffman::decode::decode_4_streams(
435 &decode_table,
436 decode_log,
437 &encoded,
438 data.len(),
439 )
440 .unwrap();
441 assert_eq!(decoded, data);
442 }
443
444 #[test]
445 fn roundtrip_all_bytes() {
446 let data: Vec<u8> = (0u8..=127).cycle().take(4096).collect();
447 let table = HuffmanEncodeTable::from_data(&data).unwrap();
448 let weights_raw = table.serialize_weights();
449 let encoded = table.encode_single_stream(&data);
450
451 let (parsed_weights, _) =
452 crate::huffman::weights::parse_huffman_weights(&weights_raw).unwrap();
453 let (decode_table, decode_log) =
454 crate::huffman::weights::build_huffman_decode_table(&parsed_weights).unwrap();
455 let decoded = crate::huffman::decode::decode_single_stream(
456 &decode_table,
457 decode_log,
458 &encoded,
459 data.len(),
460 )
461 .unwrap();
462 assert_eq!(decoded, data);
463 }
464
465 #[test]
466 fn skewed_distribution() {
467 let mut data = vec![0u8; 900];
468 data.extend(vec![1u8; 80]);
469 data.extend(vec![2u8; 15]);
470 data.extend(vec![3u8; 5]);
471 let table = HuffmanEncodeTable::from_data(&data).unwrap();
472 assert!(table.num_bits[0] < table.num_bits[3]);
473 let weights_raw = table.serialize_weights();
474 let encoded = table.encode_single_stream(&data);
475
476 let (parsed_weights, _) =
477 crate::huffman::weights::parse_huffman_weights(&weights_raw).unwrap();
478 let (decode_table, decode_log) =
479 crate::huffman::weights::build_huffman_decode_table(&parsed_weights).unwrap();
480 let decoded = crate::huffman::decode::decode_single_stream(
481 &decode_table,
482 decode_log,
483 &encoded,
484 data.len(),
485 )
486 .unwrap();
487 assert_eq!(decoded, data);
488 }
489
490 #[test]
491 fn two_symbols() {
492 let mut data = vec![0u8; 500];
493 data.extend(vec![1u8; 500]);
494 let table = HuffmanEncodeTable::from_data(&data).unwrap();
495 assert_eq!(table.num_bits[0], 1);
496 assert_eq!(table.num_bits[1], 1);
497 let encoded = table.encode_single_stream(&data);
498
499 let weights_raw = table.serialize_weights();
500 let (parsed_weights, _) =
501 crate::huffman::weights::parse_huffman_weights(&weights_raw).unwrap();
502 let (decode_table, decode_log) =
503 crate::huffman::weights::build_huffman_decode_table(&parsed_weights).unwrap();
504 let decoded = crate::huffman::decode::decode_single_stream(
505 &decode_table,
506 decode_log,
507 &encoded,
508 data.len(),
509 )
510 .unwrap();
511 assert_eq!(decoded, data);
512 }
513}