use super::bitstream::BitStream;
use crate::Error;
const LITERAL_ALPHABET: usize = 256 + 32;
const LITERAL_LEN_START: usize = 256;
const OFFSET_ALPHABET: usize = 32;
const MIN_MATCH_LEN: usize = 3;
const LITERAL_INITIAL_DELAY: usize = LITERAL_ALPHABET / 4;
const LITERAL_DELAY: usize = LITERAL_ALPHABET * 12;
const OFFSET_INITIAL_DELAY: usize = OFFSET_ALPHABET / 4;
const OFFSET_DELAY: usize = OFFSET_ALPHABET * 12;
const MAX_CODE_BITS: usize = 16;
#[derive(Clone, Copy)]
struct Node {
frequency: u32,
child_left: usize,
child_right: usize,
parent: usize,
search_me: bool,
}
struct HuffTree {
nodes: Vec<Node>,
alphabet: usize,
root: usize,
fully_active: bool,
increment: usize,
remaining: usize,
}
fn err() -> Error {
Error::compression_error()
}
fn add(a: usize, b: usize) -> Result<usize, Error> {
a.checked_add(b).ok_or_else(err)
}
fn sub(a: usize, b: usize) -> Result<usize, Error> {
a.checked_sub(b).ok_or_else(err)
}
fn shl1(shift: usize) -> Result<usize, Error> {
1usize
.checked_shl(u32::try_from(shift).map_err(|_err| err())?)
.ok_or_else(err)
}
impl HuffTree {
fn new(alphabet: usize, initial_delay: usize) -> Result<Self, Error> {
let total = sub(add(alphabet, alphabet)?, 1)?;
let root = sub(total, 1)?;
let mut nodes = vec![
Node {
frequency: 1,
child_left: 0,
child_right: 0,
parent: 0,
search_me: false,
};
total
];
for (index, node) in nodes.iter_mut().enumerate().take(alphabet) {
node.child_left = index;
node.child_right = index;
}
let mut tree = Self {
nodes,
alphabet,
root,
fully_active: false,
increment: initial_delay,
remaining: initial_delay,
};
tree.generate(0)?;
Ok(tree)
}
fn node(&self, index: usize) -> Result<&Node, Error> {
self.nodes.get(index).ok_or_else(err)
}
fn node_mut(&mut self, index: usize) -> Result<&mut Node, Error> {
self.nodes.get_mut(index).ok_or_else(err)
}
fn generate(&mut self, freq_mod: u32) -> Result<(), Error> {
loop {
for index in 0..self.alphabet {
self.node_mut(index)?.search_me = true;
}
for index in self.alphabet..self.nodes.len() {
self.node_mut(index)?.search_me = false;
}
let end = add(self.root, 1)?;
let mut next_blank = self.alphabet;
while next_blank != end {
let (mut b1, mut b2) = (0usize, 0usize);
let (mut b1_freq, mut b2_freq) = (u32::MAX, u32::MAX);
for index in 0..next_blank {
let node = self.node(index)?;
if node.search_me && node.frequency < b2_freq {
if node.frequency < b1_freq {
b2 = b1;
b2_freq = b1_freq;
b1 = index;
b1_freq = node.frequency;
} else {
b2 = index;
b2_freq = node.frequency;
}
}
}
let combined = b1_freq.wrapping_add(b2_freq);
self.node_mut(b1)?.search_me = false;
self.node_mut(b1)?.parent = next_blank;
self.node_mut(b2)?.search_me = false;
self.node_mut(b2)?.parent = next_blank;
let blank = self.node_mut(next_blank)?;
blank.frequency = combined;
blank.search_me = true;
blank.child_left = b1;
blank.child_right = b2;
next_blank = add(next_blank, 1)?;
}
if self.max_depth_exceeded()? {
for index in 0..self.alphabet {
let node = self.node_mut(index)?;
node.frequency = (node.frequency >> 2).wrapping_add(1);
}
continue;
}
break;
}
if freq_mod != 0 {
for index in 0..self.alphabet {
let node = self.node_mut(index)?;
node.frequency = (node.frequency >> freq_mod).wrapping_add(1);
}
}
Ok(())
}
fn max_depth_exceeded(&self) -> Result<bool, Error> {
for leaf in 0..self.alphabet {
let mut parent = leaf;
let mut depth = 0usize;
while parent != self.root {
depth = add(depth, 1)?;
if depth > self.nodes.len() {
return Err(err());
}
parent = self.node(parent)?.parent;
}
if depth > MAX_CODE_BITS {
return Ok(true);
}
}
Ok(false)
}
fn read_symbol(&self, bits: &mut BitStream<'_>) -> Result<usize, Error> {
let mut code = self.root;
let mut steps = 0usize;
loop {
let node = self.node(code)?;
if node.child_left == code {
return Ok(code);
}
code = if bits.read_bits(1)? == 0 {
node.child_left
} else {
node.child_right
};
steps = add(steps, 1)?;
if steps > self.nodes.len() {
return Err(err());
}
}
}
fn read_adaptive(&mut self, bits: &mut BitStream<'_>, delay: usize) -> Result<usize, Error> {
let symbol = self.read_symbol(bits)?;
let node = self.node_mut(symbol)?;
node.frequency = node.frequency.wrapping_add(1);
self.remaining = sub(self.remaining, 1)?;
if self.remaining == 0 {
if self.fully_active {
self.remaining = delay;
self.generate(1)?;
} else {
self.increment = add(self.increment, self.initial_delay())?;
if self.increment >= delay {
self.fully_active = true;
}
self.remaining = self.initial_delay();
self.generate(0)?;
}
}
Ok(symbol)
}
fn initial_delay(&self) -> usize {
if self.alphabet == OFFSET_ALPHABET {
OFFSET_INITIAL_DELAY
} else {
LITERAL_INITIAL_DELAY
}
}
}
fn read_length(bits: &mut BitStream<'_>, symbol: usize) -> Result<usize, Error> {
if symbol <= 263 {
return sub(symbol, LITERAL_LEN_START);
}
let code = sub(symbol, 264)?;
let extra_bits = add(code >> 2, 1)?;
let msb_value = shl1(add(extra_bits, 2)?)?;
let low = code & 0x0003;
let value = bits.read_bits(extra_bits)? as usize;
add(
add(value, msb_value)?,
low.checked_shl(u32::try_from(extra_bits).map_err(|_err| err())?)
.ok_or_else(err)?,
)
}
fn read_offset(tree: &mut HuffTree, bits: &mut BitStream<'_>) -> Result<usize, Error> {
let code = tree.read_adaptive(bits, OFFSET_DELAY)?;
if code <= 3 {
return Ok(code);
}
let code = sub(code, 4)?;
let extra_bits = add(code >> 1, 1)?;
let msb_value = shl1(add(extra_bits, 1)?)?;
let low = code & 0x0001;
let value = bits.read_bits(extra_bits)? as usize;
add(
add(value, msb_value)?,
low.checked_shl(u32::try_from(extra_bits).map_err(|_err| err())?)
.ok_or_else(err)?,
)
}
pub fn decompress_jb01(payload: &[u8], output_size: usize) -> Result<Vec<u8>, Error> {
let mut bits = BitStream::new(payload);
let mut literals = HuffTree::new(LITERAL_ALPHABET, LITERAL_INITIAL_DELAY)?;
let mut offsets = HuffTree::new(OFFSET_ALPHABET, OFFSET_INITIAL_DELAY)?;
let mut output = Vec::with_capacity(output_size.min(1024 * 1024));
while output.len() < output_size {
let symbol = literals.read_adaptive(&mut bits, LITERAL_DELAY)?;
if symbol < LITERAL_LEN_START {
let byte = u8::try_from(symbol).map_err(|_err| err())?;
output.push(byte);
} else {
let match_len = add(read_length(&mut bits, symbol)?, MIN_MATCH_LEN)?;
let offset = read_offset(&mut offsets, &mut bits)?;
if offset == 0 || offset > output.len() {
return Err(err());
}
let end = add(output.len(), match_len)?;
if end > output_size {
return Err(err());
}
for _ in 0..match_len {
let source = sub(output.len(), offset)?;
let byte = *output.get(source).ok_or_else(err)?;
output.push(byte);
}
}
}
Ok(output)
}