Skip to main content

huffman_rust/
canonical.rs

1use 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        // Build the canonical code book
33        let code_book = canonical_code_book(&code_lengths);
34
35        // Build the lookup tree
36        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        // Keep track of state
47        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        // Read was empty
59        if bytes_read == 0 {
60            return Err(From::from("Read was empty"));
61        }
62
63        // Create a huffman from the frequencies
64        let huff_tree = HuffmanTree::new(&freq_table)
65            .ok_or("Could not create buffman tree")?;
66
67        // Get code lengths from huffman tree
68        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            // Find the lookup entry
112            let (&min_code, entry) = self.lookup.range((Bound::Unbounded, Bound::Included(code)))
113                .next_back()
114                .ok_or("File corrupt")?;
115
116            // Index into the entry
117            let index = (code - min_code) >> (64 - entry.length);
118
119            // Lookup the index in the entry
120            buf[0] = entry.codes[index as usize];
121
122            // Write out the byte
123            write.write(&buf)?;
124
125            // Clear the first entry.length bits and left shift the code
126            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    // Sort by code_length and then by symbol
141    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    // Current code
147    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    // Group by lengths
178    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    // Create the entries to put into the tree
199    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}