Skip to main content

zrip_core/huffman/
encode.rs

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}