use std::collections::HashMap;
pub(super) struct HuffmanCoder {
codes: Vec<String>,
}
struct Node {
freq: u32,
leftmost: &'static str,
leaves: Vec<(usize, String)>,
}
impl HuffmanCoder {
pub fn new(freqs: &[u16], symbols: &[&'static str]) -> HuffmanCoder {
let mut codes = vec![String::new(); freqs.len()];
let mut pool: Vec<Node> = freqs
.iter()
.enumerate()
.filter(|&(_, &f)| f > 0)
.map(|(i, &f)| Node {
freq: u32::from(f),
leftmost: symbols[i],
leaves: vec![(i, String::new())],
})
.collect();
if pool.is_empty() {
return HuffmanCoder { codes };
}
if pool.len() == 1 {
codes[pool[0].leaves[0].0] = "0".to_string();
return HuffmanCoder { codes };
}
while pool.len() > 1 {
let a = pop_min(&mut pool);
let b = pop_min(&mut pool);
let mut leaves = Vec::with_capacity(a.leaves.len() + b.leaves.len());
for (i, code) in a.leaves {
let mut c = String::with_capacity(code.len() + 1);
c.push('0');
c.push_str(&code);
leaves.push((i, c));
}
for (i, code) in b.leaves {
let mut c = String::with_capacity(code.len() + 1);
c.push('1');
c.push_str(&code);
leaves.push((i, c));
}
pool.push(Node {
freq: a.freq + b.freq,
leftmost: a.leftmost,
leaves,
});
}
for (i, code) in pool.pop().unwrap().leaves {
codes[i] = code;
}
HuffmanCoder { codes }
}
pub fn encode(&self, symbol_index: usize) -> &str {
&self.codes[symbol_index]
}
pub fn can_encode(&self, symbol_index: usize) -> bool {
symbol_index < self.codes.len() && !self.codes[symbol_index].is_empty()
}
pub fn decode(&self, data: &[bool], pos: &mut usize) -> Option<usize> {
for (i, code) in self.codes.iter().enumerate() {
if code.is_empty() || *pos + code.len() > data.len() {
continue;
}
if code
.bytes()
.zip(&data[*pos..])
.all(|(c, &bit)| (c == b'1') == bit)
{
*pos += code.len();
return Some(i);
}
}
None
}
}
fn pop_min(pool: &mut Vec<Node>) -> Node {
let mut best = 0;
for i in 1..pool.len() {
let (a, b) = (&pool[i], &pool[best]);
if a.freq < b.freq || (a.freq == b.freq && a.leftmost < b.leftmost) {
best = i;
}
}
pool.swap_remove(best)
}
pub(super) struct Trie {
data: &'static [u8],
pub symbols: &'static [&'static str],
k: usize,
max_depth: usize,
}
impl Trie {
pub fn new(data: &'static [u8], symbols: &'static [&'static str]) -> Trie {
let k = symbols.len();
let expected = (1 + k + k * k) * k * 2;
assert_eq!(
data.len(),
expected,
"trie table size mismatch for k={k}: expected {expected}, got {}",
data.len()
);
Trie {
data,
symbols,
k,
max_depth: 2,
}
}
fn frequencies(&self, node: usize) -> Vec<u16> {
let base = node * self.k * 2;
(0..self.k)
.map(|i| u16::from_be_bytes([self.data[base + i * 2], self.data[base + i * 2 + 1]]))
.collect()
}
fn child(&self, parent: usize, symbol_index: usize) -> usize {
self.k * parent + 1 + symbol_index
}
}
pub(super) struct MultiCoder {
trie: Trie,
index: HashMap<char, usize>,
cache: std::sync::Mutex<HashMap<usize, std::sync::Arc<HuffmanCoder>>>,
}
impl MultiCoder {
pub fn new(trie: Trie) -> MultiCoder {
let index = trie
.symbols
.iter()
.enumerate()
.map(|(i, s)| (s.chars().next().unwrap(), i))
.collect();
MultiCoder {
trie,
index,
cache: std::sync::Mutex::new(HashMap::new()),
}
}
pub fn symbol_index(&self, c: char) -> Option<usize> {
self.index.get(&c).copied()
}
fn coder(&self, node: usize) -> std::sync::Arc<HuffmanCoder> {
let mut cache = self.cache.lock().expect("coder cache poisoned");
cache
.entry(node)
.or_insert_with(|| {
std::sync::Arc::new(HuffmanCoder::new(
&self.trie.frequencies(node),
self.trie.symbols,
))
})
.clone()
}
fn advance(&self, node: usize, depth: usize, symbol_index: usize) -> (usize, usize) {
if depth < self.trie.max_depth {
(self.trie.child(node, symbol_index), depth + 1)
} else {
let prev = (node - 1) % self.trie.k;
(self.trie.child(1 + prev, symbol_index), depth)
}
}
pub fn encode(&self, text: &str, start_ctx: &str) -> Option<String> {
let (mut node, mut depth) = (0usize, 0usize);
for c in start_ctx.chars() {
let idx = self.symbol_index(c)?;
(node, depth) = self.advance(node, depth, idx);
}
let mut bits = String::new();
for c in text.chars() {
let idx = self.symbol_index(c)?;
let hc = self.coder(node);
if !hc.can_encode(idx) {
return None;
}
bits.push_str(hc.encode(idx));
(node, depth) = self.advance(node, depth, idx);
}
Some(bits)
}
pub fn decode(
&self,
data: &[bool],
pos: &mut usize,
stop: Option<char>,
start_ctx: &str,
) -> String {
let (mut node, mut depth) = (0usize, 0usize);
for c in start_ctx.chars() {
let Some(idx) = self.symbol_index(c) else {
return String::new();
};
(node, depth) = self.advance(node, depth, idx);
}
let mut out = String::new();
while *pos < data.len() {
let hc = self.coder(node);
let Some(idx) = hc.decode(data, pos) else {
break;
};
let sym = self.trie.symbols[idx];
out.push_str(sym);
if stop.is_some_and(|s| sym.starts_with(s)) {
break;
}
(node, depth) = self.advance(node, depth, idx);
}
out
}
}