use crate::{Error, Result};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum BitOrder {
LsbFirst,
MsbFirst,
}
impl BitOrder {
const fn early_change(self) -> bool {
matches!(self, Self::MsbFirst)
}
}
const MAX_WIDTH: u32 = 12;
const fn widen_at(width: u32, order: BitOrder) -> u16 {
let full = 1_u16 << width;
if order.early_change() {
full.saturating_sub(1)
} else {
full
}
}
const MAX_CODES: usize = 1 << MAX_WIDTH;
#[derive(Debug)]
struct CodeReader<'a> {
data: &'a [u8],
position: usize,
bits: u32,
count: u32,
order: BitOrder,
}
impl<'a> CodeReader<'a> {
const fn new(data: &'a [u8], order: BitOrder) -> Self {
Self {
data,
position: 0,
bits: 0,
count: 0,
order,
}
}
fn next(&mut self, width: u32) -> Option<u16> {
while self.count < width {
let &byte = self.data.get(self.position)?;
self.position += 1;
match self.order {
BitOrder::LsbFirst => self.bits |= u32::from(byte) << self.count,
BitOrder::MsbFirst => self.bits = (self.bits << 8) | u32::from(byte),
}
self.count += 8;
}
let code = match self.order {
BitOrder::LsbFirst => {
let mask = (1_u32 << width) - 1;
let value = self.bits & mask;
self.bits >>= width;
self.count -= width;
value
}
BitOrder::MsbFirst => {
let shift = self.count - width;
let value = (self.bits >> shift) & ((1_u32 << width) - 1);
self.count -= width;
self.bits &= (1_u32 << self.count).wrapping_sub(1);
value
}
};
Some(code as u16)
}
}
#[derive(Debug)]
struct CodeWriter {
out: Vec<u8>,
bits: u32,
count: u32,
order: BitOrder,
}
impl CodeWriter {
const fn new(order: BitOrder) -> Self {
Self {
out: Vec::new(),
bits: 0,
count: 0,
order,
}
}
fn push(&mut self, code: u16, width: u32) {
match self.order {
BitOrder::LsbFirst => {
self.bits |= u32::from(code) << self.count;
self.count += width;
while self.count >= 8 {
self.out.push((self.bits & 0xFF) as u8);
self.bits >>= 8;
self.count -= 8;
}
}
BitOrder::MsbFirst => {
self.bits = (self.bits << width) | u32::from(code);
self.count += width;
while self.count >= 8 {
let shift = self.count - 8;
self.out.push(((self.bits >> shift) & 0xFF) as u8);
self.count -= 8;
self.bits &= (1_u32 << self.count).wrapping_sub(1);
}
}
}
}
fn finish(mut self) -> Vec<u8> {
if self.count > 0 {
let byte = match self.order {
BitOrder::LsbFirst => (self.bits & 0xFF) as u8,
BitOrder::MsbFirst => ((self.bits << (8 - self.count)) & 0xFF) as u8,
};
self.out.push(byte);
}
self.out
}
}
#[derive(Debug)]
pub struct LzwDecoder {
order: BitOrder,
minimum_width: u32,
}
impl LzwDecoder {
pub fn new(order: BitOrder, minimum_width: u32) -> Result<Self> {
if !(2..=11).contains(&minimum_width) {
return Err(Error::malformed(
"lzw",
format!("minimum code width {minimum_width} is outside 2..=11"),
));
}
Ok(Self {
order,
minimum_width,
})
}
pub fn gif(minimum_width: u32) -> Result<Self> {
Self::new(BitOrder::LsbFirst, minimum_width)
}
#[must_use]
pub const fn tiff() -> Self {
Self {
order: BitOrder::MsbFirst,
minimum_width: 8,
}
}
pub fn decode(&self, data: &[u8], limit: usize) -> Result<Vec<u8>> {
let clear = 1_u16 << self.minimum_width;
let end = clear + 1;
let literals = clear;
let mut reader = CodeReader::new(data, self.order);
let mut out: Vec<u8> = Vec::new();
let mut prefix = vec![0_u16; MAX_CODES];
let mut suffix = vec![0_u8; MAX_CODES];
let mut scratch: Vec<u8> = Vec::with_capacity(MAX_CODES);
let mut width = self.minimum_width + 1;
let mut next = end + 1;
let mut previous: Option<u16> = None;
while let Some(code) = reader.next(width) {
if code == clear {
width = self.minimum_width + 1;
next = end + 1;
previous = None;
continue;
}
if code == end {
break;
}
scratch.clear();
let first = if code < next {
expand(code, literals, &prefix, &suffix, &mut scratch)?
} else if code == next {
let Some(previous_code) = previous else {
return Err(Error::malformed(
"lzw",
"the first code after a clear cannot be a forward reference",
));
};
let first = expand(previous_code, literals, &prefix, &suffix, &mut scratch)?;
scratch.push(first);
first
} else {
return Err(Error::malformed(
"lzw",
format!("code {code} is beyond the {next} entries defined so far"),
));
};
if out.len() + scratch.len() > limit {
return Err(Error::malformed(
"lzw",
format!("stream expands beyond the {limit} byte limit"),
));
}
out.extend_from_slice(&scratch);
let room = (next as usize) < MAX_CODES;
if let Some(previous_code) = previous.filter(|_| room) {
if let (Some(p), Some(s)) =
(prefix.get_mut(next as usize), suffix.get_mut(next as usize))
{
*p = previous_code;
*s = first;
}
next += 1;
if next >= widen_at(width, self.order) && width < MAX_WIDTH {
width += 1;
}
}
previous = Some(code);
}
Ok(out)
}
}
fn expand(
code: u16,
literals: u16,
prefix: &[u16],
suffix: &[u8],
out: &mut Vec<u8>,
) -> Result<u8> {
let start = out.len();
let mut current = code;
for _ in 0..=MAX_CODES {
if current < literals {
out.push(current as u8);
let Some(slice) = out.get_mut(start..) else {
break;
};
slice.reverse();
return Ok(current as u8);
}
let index = current as usize;
let (Some(&byte), Some(&parent)) = (suffix.get(index), prefix.get(index)) else {
return Err(Error::malformed("lzw", "code refers outside the table"));
};
out.push(byte);
current = parent;
}
Err(Error::malformed("lzw", "code chain does not terminate"))
}
#[derive(Debug)]
pub struct LzwEncoder {
order: BitOrder,
minimum_width: u32,
}
impl LzwEncoder {
pub fn new(order: BitOrder, minimum_width: u32) -> Result<Self> {
if !(2..=11).contains(&minimum_width) {
return Err(Error::malformed(
"lzw",
format!("minimum code width {minimum_width} is outside 2..=11"),
));
}
Ok(Self {
order,
minimum_width,
})
}
pub fn gif(minimum_width: u32) -> Result<Self> {
Self::new(BitOrder::LsbFirst, minimum_width)
}
#[must_use]
pub const fn tiff() -> Self {
Self {
order: BitOrder::MsbFirst,
minimum_width: 8,
}
}
#[must_use]
pub fn encode(&self, data: &[u8]) -> Vec<u8> {
let clear = 1_u16 << self.minimum_width;
let end = clear + 1;
let mut writer = CodeWriter::new(self.order);
let mut width = self.minimum_width + 1;
let mut table: std::collections::HashMap<(u16, u8), u16> = std::collections::HashMap::new();
let mut next = end + 1;
writer.push(clear, width);
let mut current: Option<u16> = None;
for &byte in data {
let combined = match current {
None => {
current = Some(u16::from(byte));
continue;
}
Some(code) => (code, byte),
};
if let Some(&found) = table.get(&combined) {
current = Some(found);
continue;
}
if let Some(code) = current {
writer.push(code, width);
}
if (next as usize) < MAX_CODES {
table.insert(combined, next);
next += 1;
if next > widen_at(width, self.order) && width < MAX_WIDTH {
width += 1;
}
} else {
writer.push(clear, width);
table.clear();
width = self.minimum_width + 1;
next = end + 1;
}
current = Some(u16::from(byte));
}
if let Some(code) = current {
writer.push(code, width);
}
writer.push(end, width);
writer.finish()
}
}
#[cfg(test)]
#[allow(
clippy::unwrap_used,
clippy::expect_used,
clippy::indexing_slicing,
clippy::panic,
reason = "tests operate on known-good values and assert shapes directly"
)]
mod tests {
use super::*;
fn round_trip(order: BitOrder, width: u32, data: &[u8]) {
let encoder = LzwEncoder::new(order, width).unwrap();
let decoder = LzwDecoder::new(order, width).unwrap();
let compressed = encoder.encode(data);
let back = decoder
.decode(&compressed, data.len() * 4 + 1024)
.unwrap_or_else(|e| panic!("{order:?} width {width}: {e}"));
assert_eq!(back, data, "{order:?} width {width} did not round-trip");
}
#[test]
fn both_dialects_round_trip_everything() {
let cases: Vec<Vec<u8>> = vec![
Vec::new(),
vec![0],
vec![1, 2, 3],
vec![7; 1000],
b"the quick brown fox jumps over the lazy dog. ".repeat(60),
(0..=255_u8).cycle().take(5000).collect(),
(0..4000).map(|i| ((i * 37) % 251) as u8).collect(),
];
for order in [BitOrder::LsbFirst, BitOrder::MsbFirst] {
for data in &cases {
round_trip(order, 8, data);
}
}
}
#[test]
fn every_gif_minimum_width_round_trips() {
for width in 2..=8_u32 {
let alphabet = 1_u16 << width;
let data: Vec<u8> = (0..3000).map(|i| ((i * 7) % alphabet) as u8).collect();
round_trip(BitOrder::LsbFirst, width, &data);
}
}
#[test]
fn the_kwkwk_case_decodes_correctly() {
for order in [BitOrder::LsbFirst, BitOrder::MsbFirst] {
for length in [3_usize, 4, 5, 10, 100, 4000] {
let data = vec![0xAB_u8; length];
round_trip(order, 8, &data);
}
round_trip(order, 8, b"ababababab");
round_trip(order, 8, b"aaaaaaaaaaaaaaaa");
}
}
#[test]
fn a_stream_longer_than_the_table_resets_cleanly() {
for order in [BitOrder::LsbFirst, BitOrder::MsbFirst] {
let data: Vec<u8> = (0..200_000_u32)
.map(|i| (i.wrapping_mul(2_654_435_761) >> 24) as u8)
.collect();
round_trip(order, 8, &data);
}
}
#[test]
fn a_forward_reference_is_an_error_not_a_panic() {
let decoder = LzwDecoder::gif(8).unwrap();
let mut writer = CodeWriter::new(BitOrder::LsbFirst);
writer.push(0x100, 9);
writer.push(0x1FF, 9);
let stream = writer.finish();
let error = decoder.decode(&stream, 1 << 20).unwrap_err();
assert!(error.detail().contains("beyond"), "{error}");
}
#[test]
fn arbitrary_bytes_never_panic() {
let mut seed = 0x1234_5678_u32;
for order in [BitOrder::LsbFirst, BitOrder::MsbFirst] {
for width in [2_u32, 8, 11] {
let decoder = LzwDecoder::new(order, width).unwrap();
for _ in 0..500 {
let len = (seed % 200) as usize + 1;
let data: Vec<u8> = (0..len)
.map(|_| {
seed = seed.wrapping_mul(1_664_525).wrapping_add(1_013_904_223);
(seed >> 16) as u8
})
.collect();
let _ = decoder.decode(&data, 1 << 16);
}
}
}
}
#[test]
fn every_truncation_of_a_valid_stream_is_an_error_or_a_short_read() {
let data = b"the quick brown fox. ".repeat(40);
let stream = LzwEncoder::gif(8).unwrap().encode(&data);
let decoder = LzwDecoder::gif(8).unwrap();
for cut in 0..stream.len() {
let _ = decoder.decode(&stream[..cut], 1 << 20);
}
}
#[test]
fn the_output_limit_is_enforced() {
let data = vec![0_u8; 1_000_000];
let stream = LzwEncoder::gif(8).unwrap().encode(&data);
assert!(stream.len() < 100_000, "the bomb should be small");
let error = LzwDecoder::gif(8)
.unwrap()
.decode(&stream, 4096)
.unwrap_err();
assert!(error.detail().contains("limit"), "{error}");
}
#[test]
fn an_out_of_range_minimum_width_is_rejected() {
for width in [0_u32, 1, 12, 99] {
assert!(
LzwDecoder::new(BitOrder::LsbFirst, width).is_err(),
"{width}"
);
assert!(
LzwEncoder::new(BitOrder::LsbFirst, width).is_err(),
"{width}"
);
}
}
#[test]
fn a_clear_code_resets_the_table_mid_stream() {
let mut writer = CodeWriter::new(BitOrder::LsbFirst);
writer.push(0x100, 9); writer.push(b'A'.into(), 9);
writer.push(b'B'.into(), 9);
writer.push(0x100, 9); writer.push(b'C'.into(), 9);
writer.push(0x101, 9); let stream = writer.finish();
let out = LzwDecoder::gif(8).unwrap().decode(&stream, 1024).unwrap();
assert_eq!(out, b"ABC");
}
#[test]
fn the_two_dialects_pack_bits_differently() {
let data = b"hello world, this is a test of bit ordering".repeat(4);
let lsb = LzwEncoder::new(BitOrder::LsbFirst, 8)
.unwrap()
.encode(&data);
let msb = LzwEncoder::new(BitOrder::MsbFirst, 8)
.unwrap()
.encode(&data);
assert_ne!(lsb, msb, "the two bit orders produced identical streams");
}
}