huffman_rust/
canonical.rs1use std::io::{Read, Write};
2use std::collections::{Bound, HashMap, BTreeMap};
3use std::result::Result;
4use std::error::Error;
5
6use super::*;
7
8const MAX_U64_MASK: u64 = 1 << 63;
9
10pub type CodeBook = HashMap<u8, Vec<bool>>;
11
12#[derive(Debug)]
13pub struct LookupEntry {
14 length: u8,
15 codes: Vec<u8>,
16}
17
18impl LookupEntry {
19 pub fn new(length: u8, codes: Vec<u8>) -> LookupEntry {
20 LookupEntry {length, codes}
21 }
22}
23
24pub struct CanonicalTree {
25 pub bytes: u64,
26 pub code_book: CodeBook,
27 lookup: BTreeMap<u64, LookupEntry>,
28}
29
30impl CanonicalTree {
31 pub fn new(bytes: u64, code_lengths: Vec<(u8, u8)>) -> CanonicalTree {
32 let code_book = canonical_code_book(&code_lengths);
34
35 let lookup = lookup_tree(&code_book);
37
38 CanonicalTree {
39 bytes,
40 code_book,
41 lookup,
42 }
43 }
44
45 pub fn from_read<R: Read>(read: R) -> Result<CanonicalTree, Box<Error>> {
46 let mut bytes_read: u64 = 0;
48 let mut freq_table: [u64; NUM_BYTES] = [0; NUM_BYTES];
49
50 for byte in read.bytes() {
51 if bytes_read == u64::max_value() {
52 return Err(From::from(format!("Cannot read file larger than {} bytes", u64::max_value())));
53 }
54 bytes_read += 1;
55 freq_table[byte? as usize] += 1;
56 }
57
58 if bytes_read == 0 {
60 return Err(From::from("Read was empty"));
61 }
62
63 let huff_tree = HuffmanTree::new(&freq_table)
65 .ok_or("Could not create buffman tree")?;
66
67 let code_lengths = huff_tree.get_code_lengths();
69
70 Ok(CanonicalTree::new(bytes_read, code_lengths))
71 }
72
73 pub fn encode<R: Read, W: Write>(&self, read: & mut R, write: & mut W) -> Result<(), Box<Error>> {
74 let mut bit_writer = BitWriter::new(write);
75
76 for byte_res in read.bytes() {
77 let byte = byte_res?;
78 let code = self.code_book.get(&byte)
79 .ok_or(format!("Symbol {} not found in code book", byte))?;
80
81 bit_writer.write_bits(&code)?;
82 }
83
84 Ok(())
85 }
86
87 pub fn decode<R: Read, W: Write>(&self, read: & mut R, write: & mut W) -> Result<(), Box<Error>> {
88 let mut bit_reader = BitReader::new(read);
89
90 let mut buf: [u8; 1] = [0; 1];
91 let mut code: u64 = 0;
92 let mut mask: u64 = MAX_U64_MASK;
93 let mut offset: u64 = 0;
94
95 loop {
96 if let Some(bit) = bit_reader.read_bit()? {
97 if bit {
98 code |= mask;
99 }
100
101 mask >>= 1;
102 offset += 1;
103
104 if mask > 0 {
105 continue;
106 }
107 } else if offset == 0 {
108 return Ok(())
109 }
110
111 let (&min_code, entry) = self.lookup.range((Bound::Unbounded, Bound::Included(code)))
113 .next_back()
114 .ok_or("File corrupt")?;
115
116 let index = (code - min_code) >> (64 - entry.length);
118
119 buf[0] = entry.codes[index as usize];
121
122 write.write(&buf)?;
124
125 mask = MAX_U64_MASK;
127 for _ in 0..entry.length {
128 code &= !mask;
129 mask >>= 1;
130 }
131
132 code <<= entry.length;
133 offset -= entry.length as u64;
134 mask = 1 << entry.length as u64 - 1;
135 }
136 }
137}
138
139pub fn canonical_code_book(code_lengths: &[(u8, u8)]) -> CodeBook {
140 let mut sorted = Vec::from(code_lengths);
142 sorted.sort_by_key(|&(symbol, length)| (length, symbol));
143
144 let mut result = HashMap::new();
145
146 let mut code: u64 = 0;
148
149 let mut iter = sorted.iter().peekable();
150 while let Some(&(symbol, length)) = iter.next() {
151 result.insert(symbol, code_to_vec(length, code));
152
153 if let Some(&&(_symbol_next, length_next)) = iter.peek() {
154 code = (code + 1) << (length_next - length);
155 }
156 }
157
158 result
159}
160
161#[inline]
162fn code_to_vec(length: u8, code: u64) -> Vec<bool> {
163 let mut vec = Vec::with_capacity(length as usize);
164 let mut mask = 1 << ((length - 1) as u64);
165
166 for _ in 0..(length as u64) {
167 vec.push((mask & code) != 0);
168 mask >>= 1;
169 }
170
171 vec
172}
173
174pub fn lookup_tree(code_book: &CodeBook) -> BTreeMap<u64, LookupEntry> {
175 let mut tree = BTreeMap::new();
176
177 let mut map: HashMap<usize, Vec<(u8, u64)>> = HashMap::new();
179
180 for (&symbol, code_vec) in code_book.iter() {
181 let vec = map.entry(code_vec.len())
182 .or_insert(Vec::new());
183
184 let mut mask: u64 = MAX_U64_MASK;
185 let mut code: u64 = 0;
186
187 for &bit in code_vec.iter() {
188 if bit {
189 code |= mask;
190 }
191
192 mask >>= 1;
193 }
194
195 vec.push((symbol, code));
196 }
197
198 for (&length, &ref vec) in map.iter() {
200 let min_code = vec.iter()
201 .map(|&(_symbol, code)| code)
202 .min()
203 .expect(&format!("No codes for length {}", length));
204
205 let mut symbols: Vec<u8> = vec.iter()
206 .map(|&(symbol, _code)| symbol)
207 .collect();
208 symbols.sort();
209
210 let entry = LookupEntry::new(length as u8, symbols);
211
212 tree.insert(min_code, entry);
213 }
214
215 tree
216}
217
218#[cfg(test)]
219mod tests {
220 use super::*;
221 use std::io::Cursor;
222 use std::vec::Vec;
223
224 #[test]
225 fn test_small_sample_string() {
226 let text = "a small sample string";
227
228 assert!(encode_decode_test(text.as_bytes()));
229 }
230
231 fn encode_decode_test(text: &[u8]) -> bool {
232 let mut encoded_cursor = Cursor::new(text);
233 let tree = CanonicalTree::from_read(&mut encoded_cursor).unwrap();
234 encoded_cursor = Cursor::new(text);
235
236 let mut encoded = Vec::new();
237
238 tree.encode(&mut encoded_cursor, &mut encoded).unwrap();
239
240 let mut decoded = Vec::new();
241
242 tree.decode(&mut Cursor::new(encoded), &mut decoded).unwrap();
243
244 decoded == text
245 }
246}