Skip to main content

atmos/os_lib/webp/
huffman.rs

1use alloc::vec;
2use alloc::vec::Vec;
3// Rudimentary utility for reading Canonical Huffman Codes.
4// Based off <https://github.com/webmproject/libwebp/blob/7f8472a610b61ec780ef0a8873cd954ac512a505/src/utils/huffman.c>
5
6use crate::os_lib::webp::io::BufRead;
7
8use crate::os_lib::webp::decoder::DecodingError;
9
10use super::lossless::BitReader;
11
12const MAX_ALLOWED_CODE_LENGTH: usize = 15;
13const MAX_TABLE_BITS: u8 = 10;
14
15#[derive(Clone, Copy, Debug, PartialEq, Eq)]
16enum HuffmanTreeNode {
17    Branch(usize), //offset in vector to children
18    Leaf(u16),     //symbol stored in leaf
19    Empty,
20}
21
22#[derive(Clone, Debug)]
23enum HuffmanTreeInner {
24    Single(u16),
25    Tree {
26        tree: Vec<HuffmanTreeNode>,
27        table: Vec<u32>,
28        table_mask: u16,
29    },
30}
31
32/// Huffman tree
33#[derive(Clone, Debug)]
34pub(crate) struct HuffmanTree(HuffmanTreeInner);
35
36impl Default for HuffmanTree {
37    fn default() -> Self {
38        Self(HuffmanTreeInner::Single(0))
39    }
40}
41
42impl HuffmanTree {
43    /// Builds a tree implicitly, just from code lengths
44    pub(crate) fn build_implicit(code_lengths: Vec<u16>) -> Result<Self, DecodingError> {
45        // Count symbols and build histogram
46        let mut num_symbols = 0;
47        let mut code_length_hist = [0; MAX_ALLOWED_CODE_LENGTH + 1];
48        for &length in code_lengths.iter().filter(|&&x| x != 0) {
49            code_length_hist[usize::from(length)] += 1;
50            num_symbols += 1;
51        }
52
53        // Handle special cases
54        if num_symbols == 0 {
55            return Err(DecodingError::HuffmanError);
56        } else if num_symbols == 1 {
57            let root_symbol = code_lengths.iter().position(|&x| x != 0).unwrap() as u16;
58            return Ok(Self::build_single_node(root_symbol));
59        };
60
61        // Assign codes
62        let mut curr_code = 0;
63        let mut next_codes = [0; MAX_ALLOWED_CODE_LENGTH + 1];
64        let max_code_length = code_length_hist.iter().rposition(|&x| x != 0).unwrap() as u16;
65        for code_len in 1..usize::from(max_code_length) + 1 {
66            next_codes[code_len] = curr_code;
67            curr_code = (curr_code + code_length_hist[code_len]) << 1;
68        }
69
70        // Confirm that the huffman tree is valid
71        if curr_code != 2 << max_code_length {
72            return Err(DecodingError::HuffmanError);
73        }
74
75        // Calculate table/tree parameters
76        let table_bits = max_code_length.min(u16::from(MAX_TABLE_BITS));
77        let table_size = (1 << table_bits) as usize;
78        let table_mask = table_size as u16 - 1;
79        let tree_size = code_length_hist[table_bits as usize + 1..=max_code_length as usize]
80            .iter()
81            .sum::<u16>() as usize;
82
83        // Populate decoding table
84        let mut tree = Vec::with_capacity(2 * tree_size);
85        let mut table = vec![0; table_size];
86        for (symbol, &length) in code_lengths.iter().enumerate() {
87            if length == 0 {
88                continue;
89            }
90
91            let code = next_codes[length as usize];
92            next_codes[length as usize] += 1;
93
94            if length <= table_bits {
95                let mut j = (u16::reverse_bits(code) >> (16 - length)) as usize;
96                let entry = (u32::from(length) << 16) | symbol as u32;
97                while j < table_size {
98                    table[j] = entry;
99                    j += 1 << length as usize;
100                }
101            } else {
102                let table_index =
103                    ((u16::reverse_bits(code) >> (16 - length)) & table_mask) as usize;
104                let table_value = table[table_index];
105
106                debug_assert_eq!(table_value >> 16, 0);
107
108                let mut node_index = if table_value == 0 {
109                    let node_index = tree.len();
110                    table[table_index] = (node_index + 1) as u32;
111                    tree.push(HuffmanTreeNode::Empty);
112                    node_index
113                } else {
114                    (table_value - 1) as usize
115                };
116
117                let code = usize::from(code);
118                for depth in (0..length - table_bits).rev() {
119                    let node = tree[node_index];
120
121                    let offset = match node {
122                        HuffmanTreeNode::Empty => {
123                            // Turns a node from empty into a branch and assigns its children
124                            let offset = tree.len() - node_index;
125                            tree[node_index] = HuffmanTreeNode::Branch(offset);
126                            tree.push(HuffmanTreeNode::Empty);
127                            tree.push(HuffmanTreeNode::Empty);
128                            offset
129                        }
130                        HuffmanTreeNode::Leaf(_) => return Err(DecodingError::HuffmanError),
131                        HuffmanTreeNode::Branch(offset) => offset,
132                    };
133
134                    node_index += offset + ((code >> depth) & 1);
135                }
136
137                match tree[node_index] {
138                    HuffmanTreeNode::Empty => {
139                        tree[node_index] = HuffmanTreeNode::Leaf(symbol as u16);
140                    }
141                    HuffmanTreeNode::Leaf(_) => return Err(DecodingError::HuffmanError),
142                    HuffmanTreeNode::Branch(_offset) => return Err(DecodingError::HuffmanError),
143                }
144            }
145        }
146
147        Ok(Self(HuffmanTreeInner::Tree {
148            tree,
149            table,
150            table_mask,
151        }))
152    }
153
154    pub(crate) const fn build_single_node(symbol: u16) -> Self {
155        Self(HuffmanTreeInner::Single(symbol))
156    }
157
158    pub(crate) fn build_two_node(zero: u16, one: u16) -> Self {
159        Self(HuffmanTreeInner::Tree {
160            tree: vec![
161                HuffmanTreeNode::Leaf(zero),
162                HuffmanTreeNode::Leaf(one),
163                HuffmanTreeNode::Empty,
164            ],
165            table: vec![(1 << 16) | u32::from(zero), (1 << 16) | u32::from(one)],
166            table_mask: 0x1,
167        })
168    }
169
170    pub(crate) const fn is_single_node(&self) -> bool {
171        matches!(self.0, HuffmanTreeInner::Single(_))
172    }
173
174    #[inline(never)]
175    fn read_symbol_slowpath<R: BufRead>(
176        tree: &[HuffmanTreeNode],
177        mut v: usize,
178        start_index: usize,
179        bit_reader: &mut BitReader<R>,
180    ) -> Result<u16, DecodingError> {
181        let mut depth = MAX_TABLE_BITS;
182        let mut index = start_index;
183        loop {
184            match &tree[index] {
185                HuffmanTreeNode::Branch(children_offset) => {
186                    index += children_offset + (v & 1);
187                    depth += 1;
188                    v >>= 1;
189                }
190                HuffmanTreeNode::Leaf(symbol) => {
191                    bit_reader.consume(depth)?;
192                    return Ok(*symbol);
193                }
194                HuffmanTreeNode::Empty => return Err(DecodingError::HuffmanError),
195            }
196        }
197    }
198
199    /// Reads a symbol using the bit reader.
200    ///
201    /// You must call call `bit_reader.fill()` before calling this function or it may erroroneosly
202    /// detect the end of the stream and return a bitstream error.
203    pub(crate) fn read_symbol<R: BufRead>(
204        &self,
205        bit_reader: &mut BitReader<R>,
206    ) -> Result<u16, DecodingError> {
207        match &self.0 {
208            HuffmanTreeInner::Tree {
209                tree,
210                table,
211                table_mask,
212            } => {
213                let v = bit_reader.peek_full() as u16;
214                let entry = table[(v & table_mask) as usize];
215                if entry >> 16 != 0 {
216                    bit_reader.consume((entry >> 16) as u8)?;
217                    return Ok(entry as u16);
218                }
219
220                Self::read_symbol_slowpath(
221                    tree,
222                    (v >> MAX_TABLE_BITS) as usize,
223                    ((entry & 0xffff) - 1) as usize,
224                    bit_reader,
225                )
226            }
227            HuffmanTreeInner::Single(symbol) => Ok(*symbol),
228        }
229    }
230
231    /// Peek at the next symbol in the bitstream if it can be read with only a primary table lookup.
232    ///
233    /// Returns a tuple of the codelength and symbol value. This function may return wrong
234    /// information if there aren't enough bits in the bit reader to read the next symbol.
235    pub(crate) fn peek_symbol<R: BufRead>(&self, bit_reader: &BitReader<R>) -> Option<(u8, u16)> {
236        match &self.0 {
237            HuffmanTreeInner::Tree {
238                table, table_mask, ..
239            } => {
240                let v = bit_reader.peek_full() as u16;
241                let entry = table[(v & table_mask) as usize];
242                if entry >> 16 != 0 {
243                    return Some(((entry >> 16) as u8, entry as u16));
244                }
245                None
246            }
247            HuffmanTreeInner::Single(symbol) => Some((0, *symbol)),
248        }
249    }
250}