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;
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));
}
}