hfmn 0.1.0

A flexible Huffman coding implementation
Documentation
use std::fmt::Debug;

use crate::code::Bit;
use crate::{code::HfmnCode, errors::Error};
use rustc_hash::FxHashMap as HashMap;

#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord)]
enum TableEntry<T> {
    Code { symbol: T, length: u8 },
    TooLong,
}

#[derive(Debug, Clone, PartialEq, Eq)]
pub struct DecodeTables<T> {
    tier1_table: [TableEntry<T>; 256],
    full_table: HashMap<HfmnCode, TableEntry<T>>,
}

impl<T: Clone + Debug> DecodeTables<T> {
    pub fn new(encode_map: &HashMap<T, HfmnCode>) -> Self {
        let mut tier1_table = [const { TableEntry::TooLong }; 256];
        let mut full_table = HashMap::default();

        for (k, v) in encode_map {
            if v.len() <= 8 {
                let all_last_chunks = get_all_paddings::<8>(v);
                for chunk in all_last_chunks {
                    tier1_table[*chunk.inner()] = TableEntry::Code {
                        symbol: k.clone(),
                        length: v.len(),
                    };
                }
            } else {
                tier1_table[*v.first_n_bits(8).inner()] = TableEntry::TooLong;
                full_table.insert(
                    *v,
                    TableEntry::Code {
                        symbol: k.clone(),
                        length: v.len(),
                    },
                );
            }
        }

        Self {
            tier1_table,
            full_table,
        }
    }

    pub fn decode(&self, bytes: &[u8], num_symbols: usize) -> Result<Vec<&T>, Error> {
        let mut start_idx = 0;
        let mut codes = Vec::new();
        while start_idx < (bytes.len() - 1) * 8 {
            match &self.tier1_table[get_8_bits(bytes, start_idx)] {
                TableEntry::Code { symbol, length } => {
                    codes.push(symbol);
                    start_idx += *length as usize;
                }
                TableEntry::TooLong => {
                    for i in 0..(std::mem::size_of::<usize>() - 8) {
                        let idx = start_idx + 8 + i;
                        // if idx > bits.len() {
                        //     return Err(Error::InvalidCode);
                        // }
                        match self.full_table.get(&index_by_bits(bytes, start_idx, idx)) {
                            Some(entry) => match entry {
                                TableEntry::Code { symbol, length } => {
                                    codes.push(symbol);
                                    start_idx += *length as usize;
                                }
                                TableEntry::TooLong => unreachable!(),
                            },
                            None => return Err(Error::InvalidCode),
                        }
                    }
                }
            }
        }

        while start_idx < bytes.len() * 8 {
            if codes.len() >= num_symbols {
                break;
            }
            let local_bits = index_by_bits(bytes, start_idx, bytes.len() * 8);

            match &self.tier1_table[*local_bits.inner()] {
                TableEntry::Code { symbol, length } => {
                    codes.push(symbol);
                    start_idx += *length as usize;
                }
                TableEntry::TooLong => {
                    return Err(Error::InvalidCode);
                }
            }
        }

        Ok(codes)
    }
}

#[allow(clippy::cast_possible_truncation)]
fn index_by_bits(bytes: &[u8], start_idx: usize, end_idx: usize) -> HfmnCode {
    let length = end_idx - start_idx;
    let pushed_length = (1 + end_idx / 8 - start_idx / 8) * 8;

    assert!(length < std::mem::size_of::<usize>() * 8);

    let mut symbol = HfmnCode::new();
    *symbol.length_mut() = length as u8;

    for byte in &bytes[start_idx / 8..=(end_idx / 8)] {
        *symbol.inner_mut() <<= 8;
        *symbol.inner_mut() |= *byte as usize;
    }

    *symbol.inner_mut() <<= std::mem::size_of::<usize>() * 8 - pushed_length + start_idx % 8;
    *symbol.inner_mut() >>= std::mem::size_of::<usize>() * 8 - length;

    symbol
}

const fn get_8_bits(bytes: &[u8], start_idx: usize) -> usize {
    let byte_idx = start_idx / 8;
    let bit_offset = start_idx % 8;

    let combined = ((bytes[byte_idx] as usize) << 8) | (bytes[byte_idx + 1] as usize);
    (combined >> (8 - bit_offset)) & 0b1111_1111
}

fn get_all_paddings<const N: u8>(code: &HfmnCode) -> Vec<HfmnCode> {
    let padding_length = N - code.len();

    let mut permuted_codes = vec![*code];
    let mut new_codes = Vec::new();

    for _ in 0..padding_length {
        for p_code in &mut permuted_codes {
            let mut left = *p_code;
            left.push(Bit::Zero);

            let mut right = *p_code;
            right.push(Bit::One);

            new_codes.push(left);
            new_codes.push(right);
        }

        permuted_codes = new_codes;
        new_codes = Vec::new();
    }

    permuted_codes
}

#[cfg(test)]
mod tests {
    use super::*;

    #[test]
    fn check_paddings() {
        let base_code = HfmnCode::from(0b11_1111, 6);

        let paddings = get_all_paddings::<8>(&base_code);

        let expected = [0b1111_1100, 0b1111_1101, 0b1111_1110, 0b1111_1111];

        assert_eq!(
            paddings,
            expected
                .into_iter()
                .map(|e| HfmnCode::from(e, 8))
                .collect::<Vec<HfmnCode>>()
        );
    }

    #[test]
    fn check_construction() {
        let mut test_map = HashMap::default();

        let a_code = HfmnCode::from(0, 1);
        let b_code = HfmnCode::from(0b1_1111_1110, 9);

        let mut first_pad_a_code = a_code;
        first_pad_a_code.resize(8);

        test_map.insert('a', a_code);
        test_map.insert('b', b_code);

        let table = DecodeTables::<char>::new(&test_map);

        assert_eq!(
            *table.tier1_table.get(*first_pad_a_code.inner()).unwrap(),
            TableEntry::Code {
                symbol: 'a',
                length: 1
            },
        );

        assert_eq!(
            *table
                .tier1_table
                .get(*b_code.first_n_bits(8).inner())
                .unwrap(),
            TableEntry::TooLong,
        );

        let _ = test_map.remove(&'a');
        for (symbol, code) in test_map {
            assert_eq!(
                table.full_table.get(&code),
                Some(&TableEntry::Code {
                    symbol,
                    length: code.len()
                })
            );
        }
    }

    #[test]
    fn check_bit_access() {
        let bytes = vec![0b0100_1101, 0b1000_0001];

        assert_eq!(
            index_by_bits(&bytes, 0, 8),
            HfmnCode::from(bytes[0] as usize, 8)
        );
    }

    #[test]
    fn check_offset_bit_access() {
        let bytes = vec![0b0101_1111, 0b1011_0111];

        assert_eq!(index_by_bits(&bytes, 3, 11), HfmnCode::from(0b1111_1101, 8));
    }
}