Skip to main content

boytacean_encoding/
huffman.rs

1use std::{
2    cmp::Ordering,
3    collections::BinaryHeap,
4    io::{Cursor, Read},
5    mem::size_of,
6};
7
8use boytacean_common::error::Error;
9
10use crate::codec::Codec;
11
12#[derive(Debug, Eq, PartialEq)]
13struct Node {
14    frequency: u32,
15    character: Option<u8>,
16    left: Option<Box<Node>>,
17    right: Option<Box<Node>>,
18}
19
20impl Ord for Node {
21    fn cmp(&self, other: &Self) -> Ordering {
22        other.frequency.cmp(&self.frequency)
23    }
24}
25
26impl PartialOrd for Node {
27    fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
28        Some(self.cmp(other))
29    }
30}
31
32pub struct Huffman;
33
34impl Huffman {
35    fn build_frequency(data: &[u8]) -> [u32; 256] {
36        let mut frequency_map = [0_u32; 256];
37        for &byte in data {
38            frequency_map[byte as usize] += 1;
39        }
40        frequency_map
41    }
42
43    fn build_tree(frequency_map: &[u32; 256]) -> Option<Box<Node>> {
44        let mut heap: BinaryHeap<Box<Node>> = BinaryHeap::new();
45
46        for (byte, &frequency) in frequency_map.iter().enumerate() {
47            if frequency == 0 {
48                continue;
49            }
50            heap.push(Box::new(Node {
51                frequency,
52                character: Some(byte as u8),
53                left: None,
54                right: None,
55            }));
56        }
57
58        while heap.len() > 1 {
59            let left = heap.pop().unwrap();
60            let right = heap.pop().unwrap();
61
62            let merged = Box::new(Node {
63                frequency: left.frequency + right.frequency,
64                character: None,
65                left: Some(left),
66                right: Some(right),
67            });
68
69            heap.push(merged);
70        }
71
72        heap.pop()
73    }
74
75    fn build_codes(node: &Node, prefix: Vec<u8>, codes: &mut [Vec<u8>]) {
76        if let Some(character) = node.character {
77            codes[character as usize] = prefix;
78        } else {
79            if let Some(ref left) = node.left {
80                let mut left_prefix = prefix.clone();
81                left_prefix.push(0);
82                Self::build_codes(left, left_prefix, codes);
83            }
84            if let Some(ref right) = node.right {
85                let mut right_prefix = prefix;
86                right_prefix.push(1);
87                Self::build_codes(right, right_prefix, codes);
88            }
89        }
90    }
91
92    fn encode_data(data: &[u8], codes: &[Vec<u8>]) -> Vec<u8> {
93        let mut bit_buffer = Vec::new();
94        let mut current_byte = 0u8;
95        let mut bit_count = 0;
96
97        for &byte in data {
98            let code = &codes[byte as usize];
99            for &bit in code {
100                current_byte <<= 1;
101                if bit == 1 {
102                    current_byte |= 1;
103                }
104                bit_count += 1;
105
106                if bit_count == 8 {
107                    bit_buffer.push(current_byte);
108                    current_byte = 0;
109                    bit_count = 0;
110                }
111            }
112        }
113
114        if bit_count > 0 {
115            current_byte <<= 8 - bit_count;
116            bit_buffer.push(current_byte);
117        }
118
119        bit_buffer
120    }
121
122    fn decode_data(encoded: &[u8], root: &Node, data_length: u64) -> Vec<u8> {
123        let mut decoded = Vec::new();
124        let mut current_node = root;
125        let mut bit_index = 0;
126
127        for &byte in encoded {
128            if decoded.len() as u64 == data_length {
129                break;
130            }
131
132            for bit_offset in (0..8).rev() {
133                let bit = (byte >> bit_offset) & 1;
134                current_node = if bit == 0 {
135                    current_node.left.as_deref().unwrap()
136                } else {
137                    current_node.right.as_deref().unwrap()
138                };
139
140                if let Some(character) = current_node.character {
141                    decoded.push(character);
142                    current_node = root;
143                }
144
145                if decoded.len() as u64 == data_length {
146                    break;
147                }
148
149                bit_index += 1;
150                if bit_index == encoded.len() * 8 {
151                    break;
152                }
153            }
154        }
155
156        decoded
157    }
158
159    fn encode_tree(node: &Node) -> Vec<u8> {
160        let mut result = Vec::new();
161        if let Some(character) = node.character {
162            result.push(1);
163            result.push(character);
164        } else {
165            result.push(0);
166            if let Some(ref left) = node.left {
167                result.extend(Self::encode_tree(left));
168            }
169            if let Some(ref right) = node.right {
170                result.extend(Self::encode_tree(right));
171            }
172        }
173        result
174    }
175
176    fn decode_tree(data: &mut &[u8]) -> Box<Node> {
177        let mut node = Box::new(Node {
178            frequency: 0,
179            character: None,
180            left: None,
181            right: None,
182        });
183
184        if data[0] == 1 {
185            node.character = Some(data[1]);
186            *data = &data[2..];
187        } else {
188            *data = &data[1..];
189            node.left = Some(Self::decode_tree(data));
190            node.right = Some(Self::decode_tree(data));
191        }
192        node
193    }
194}
195
196impl Codec for Huffman {
197    type EncodeOptions = ();
198    type DecodeOptions = ();
199
200    fn encode(data: &[u8], _options: &Self::EncodeOptions) -> Result<Vec<u8>, Error> {
201        let frequency_map = Self::build_frequency(data);
202        let tree = Self::build_tree(&frequency_map)
203            .ok_or(Error::CustomError(String::from("Failed to build tree")))?;
204
205        let mut codes = vec![Vec::new(); 256];
206        Self::build_codes(&tree, Vec::new(), &mut codes);
207
208        let encoded_tree = Self::encode_tree(&tree);
209        let encoded_data = Self::encode_data(data, &codes);
210        let tree_length = encoded_tree.len() as u32;
211        let data_length = data.len() as u64;
212
213        let mut result = Vec::new();
214        result.extend(tree_length.to_be_bytes());
215        result.extend(encoded_tree);
216        result.extend(data_length.to_be_bytes());
217        result.extend(encoded_data);
218
219        Ok(result)
220    }
221
222    fn decode(data: &[u8], _options: &Self::DecodeOptions) -> Result<Vec<u8>, Error> {
223        let mut reader = Cursor::new(data);
224
225        let mut buffer = [0x00; size_of::<u32>()];
226        reader.read_exact(&mut buffer)?;
227        let tree_length = u32::from_be_bytes(buffer);
228
229        let mut buffer = vec![0; tree_length as usize];
230        reader.read_exact(&mut buffer)?;
231        let tree = Self::decode_tree(&mut buffer.as_slice());
232
233        let mut buffer = [0x00; size_of::<u64>()];
234        reader.read_exact(&mut buffer)?;
235        let data_length = u64::from_be_bytes(buffer);
236
237        let mut buffer =
238            vec![0; data.len() - size_of::<u32>() - tree_length as usize - size_of::<u64>()];
239        reader.read_exact(&mut buffer)?;
240
241        let result = Self::decode_data(&buffer, &tree, data_length);
242
243        Ok(result)
244    }
245}
246
247pub fn encode_huffman(data: &[u8]) -> Result<Vec<u8>, Error> {
248    Huffman::encode(data, &())
249}
250
251pub fn decode_huffman(data: &[u8]) -> Result<Vec<u8>, Error> {
252    Huffman::decode(data, &())
253}
254
255#[cfg(test)]
256mod tests {
257    use super::{decode_huffman, encode_huffman};
258
259    #[test]
260    fn test_huffman_encoding() {
261        let data = b"this is an example for huffman encoding, huffman encoding, huffman encoding";
262        let encoded = encode_huffman(data).unwrap();
263        let decoded = decode_huffman(&encoded).unwrap();
264        assert_eq!(data.to_vec(), decoded);
265        assert_eq!(encoded.len(), 109);
266        assert_eq!(decoded.len(), 75);
267    }
268}