use alloc::vec;
use alloc::vec::Vec;
pub const WINDOW_SIZE: usize = 32768;
const MIN_MATCH: usize = 3;
const MAX_MATCH: usize = 258;
const HASH_BUCKETS: usize = 65536;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum LzError {
UnexpectedEnd,
InvalidDistance,
LengthMismatch,
}
impl core::fmt::Display for LzError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
Self::UnexpectedEnd => f.write_str("LZ77 stream ended unexpectedly"),
Self::InvalidDistance => f.write_str("LZ77 match distance out of range"),
Self::LengthMismatch => f.write_str("LZ77 decompressed length mismatch"),
}
}
}
#[cfg(feature = "std")]
impl std::error::Error for LzError {}
#[inline(always)]
fn hash3(bytes: &[u8]) -> usize {
let a = bytes[0] as usize;
let b = bytes[1] as usize;
let c = bytes[2] as usize;
let mut hash = 0x811c_9dc5usize;
hash ^= a;
hash = hash.wrapping_mul(0x0100_0193);
hash ^= b;
hash = hash.wrapping_mul(0x0100_0193);
hash ^= c;
hash = hash.wrapping_mul(0x0100_0193);
(hash >> 16) & (HASH_BUCKETS - 1)
}
pub fn compress(input: &[u8], level: i32) -> Vec<u8> {
if input.len() < MIN_MATCH {
return Vec::new();
}
let level = level.clamp(1, 22);
let max_chain = (4 + (level as usize) * 4).min(96);
let mut head = vec![u32::MAX; HASH_BUCKETS];
let mut prev = vec![u32::MAX; input.len()];
let mut output = BitWriter::new();
output.write_u32(input.len() as u32);
let mut pos = 0usize;
while pos < input.len() {
let mut best_len = 0usize;
let mut best_dist = 0usize;
let max_len = core::cmp::min(MAX_MATCH, input.len() - pos);
if max_len >= MIN_MATCH {
let h = hash3(&input[pos..]);
let mut candidate = head[h];
let mut chain = 0usize;
let window_start = pos.saturating_sub(WINDOW_SIZE);
while candidate != u32::MAX && chain < max_chain {
let c = candidate as usize;
if c < window_start {
break;
}
let dist = pos - c;
if dist > WINDOW_SIZE {
candidate = prev[c];
chain += 1;
continue;
}
if c + max_len <= input.len() {
let mut len = 0usize;
let remaining = core::cmp::min(max_len, input.len() - c);
while len < remaining && input[c + len] == input[pos + len] {
len += 1;
}
if len > best_len {
best_len = len;
best_dist = dist;
if len == max_len {
break;
}
}
}
candidate = prev[c];
chain += 1;
}
}
if best_len >= MIN_MATCH {
output.write_bit(true);
output.write_bits((best_len - MIN_MATCH) as u32, 8);
output.write_bits((best_dist - 1) as u32, 15);
let end = core::cmp::min(pos + best_len, input.len());
for i in pos..end {
if i + MIN_MATCH <= input.len() {
let h = hash3(&input[i..]);
prev[i] = head[h];
head[h] = i as u32;
}
}
pos = end;
} else {
output.write_bit(false);
output.write_bits(input[pos] as u32, 8);
if pos + MIN_MATCH <= input.len() {
let h = hash3(&input[pos..]);
prev[pos] = head[h];
head[h] = pos as u32;
}
pos += 1;
}
}
output.finish()
}
pub fn decompress(input: &[u8], expected: usize) -> Result<Vec<u8>, LzError> {
let mut reader = BitReader::new(input);
let declared = reader.read_u32().ok_or(LzError::UnexpectedEnd)? as usize;
if declared > expected {
return Err(LzError::LengthMismatch);
}
let mut output = Vec::with_capacity(declared);
while output.len() < declared {
let is_match = reader.read_bit().ok_or(LzError::UnexpectedEnd)?;
if is_match {
let length = reader.read_bits(8).ok_or(LzError::UnexpectedEnd)? as usize + MIN_MATCH;
let distance = reader.read_bits(15).ok_or(LzError::UnexpectedEnd)? as usize + 1;
if distance > output.len() {
return Err(LzError::InvalidDistance);
}
let start = output.len();
for i in 0..length {
let src = output[start + i - distance];
output.push(src);
}
} else {
let byte = reader.read_bits(8).ok_or(LzError::UnexpectedEnd)? as u8;
output.push(byte);
}
}
if output.len() != declared {
return Err(LzError::LengthMismatch);
}
Ok(output)
}
struct BitWriter {
buffer: Vec<u8>,
bit_buffer: u64,
bit_count: u32,
}
impl BitWriter {
fn new() -> Self {
Self {
buffer: Vec::new(),
bit_buffer: 0,
bit_count: 0,
}
}
fn write_bit(&mut self, value: bool) {
self.ensure_capacity(1);
self.bit_buffer |= (value as u64) << self.bit_count;
self.bit_count += 1;
if self.bit_count == 64 {
self.flush();
}
}
fn write_bits(&mut self, value: u32, count: u32) {
debug_assert!(count <= 32);
self.ensure_capacity(count);
self.bit_buffer |= (value as u64) << self.bit_count;
self.bit_count += count;
if self.bit_count >= 64 {
self.flush();
}
}
fn ensure_capacity(&mut self, needed: u32) {
if self.bit_count + needed > 64 {
self.flush();
}
}
fn write_u32(&mut self, value: u32) {
if self.bit_count == 0 {
self.buffer.extend_from_slice(&value.to_le_bytes());
return;
}
for shift in (0..32).step_by(8) {
self.write_bits((value >> shift) & 0xff, 8);
}
}
fn flush(&mut self) {
while self.bit_count >= 8 {
self.buffer.push((self.bit_buffer & 0xff) as u8);
self.bit_buffer >>= 8;
self.bit_count -= 8;
}
}
fn finish(mut self) -> Vec<u8> {
self.flush();
if self.bit_count > 0 {
self.buffer.push((self.bit_buffer & 0xff) as u8);
}
self.buffer
}
}
struct BitReader<'a> {
bytes: &'a [u8],
byte_pos: usize,
bit_pos: u32,
}
impl<'a> BitReader<'a> {
fn new(bytes: &'a [u8]) -> Self {
Self {
bytes,
byte_pos: 0,
bit_pos: 0,
}
}
fn read_bit(&mut self) -> Option<bool> {
let byte = *self.bytes.get(self.byte_pos)?;
let bit = (byte >> self.bit_pos) & 1 != 0;
self.bit_pos += 1;
if self.bit_pos == 8 {
self.bit_pos = 0;
self.byte_pos += 1;
}
Some(bit)
}
fn read_bits(&mut self, count: u32) -> Option<u32> {
let mut value = 0u32;
for i in 0..count {
let bit = self.read_bit()?;
value |= (bit as u32) << i;
}
Some(value)
}
fn read_u32(&mut self) -> Option<u32> {
let a = self.read_bits(8)?;
let b = self.read_bits(8)?;
let c = self.read_bits(8)?;
let d = self.read_bits(8)?;
Some(a | (b << 8) | (c << 16) | (d << 24))
}
}
#[cfg(test)]
mod tests {
use super::*;
use alloc::format;
use alloc::vec;
fn roundtrip(data: &[u8], level: i32) -> Vec<u8> {
let packed = compress(data, level);
if packed.is_empty() {
return data.to_vec();
}
decompress(&packed, data.len()).expect("roundtrip")
}
#[test]
fn literals_roundtrip() {
let data = b"hello world, this is a mostly incompressible sentence!";
assert_eq!(roundtrip(data, 3), data);
}
#[test]
fn repetitive_data_roundtrip() {
let data = b"abcabcabcabcabcabcabcabcabcabcabcabc";
assert_eq!(roundtrip(data, 3), data);
let packed = compress(data, 3);
assert!(!packed.is_empty());
assert!(packed.len() < data.len());
}
#[test]
fn long_repetition_roundtrip() {
let data = vec![0xAB_u8; 100_000];
assert_eq!(roundtrip(&data, 9), data);
let packed = compress(&data, 9);
assert!(packed.len() < data.len());
}
#[test]
fn window_boundary_match() {
let mut data = vec![0_u8; WINDOW_SIZE];
data.extend_from_slice(b"marker");
data.extend_from_slice(&[0_u8; WINDOW_SIZE]);
data.extend_from_slice(b"marker");
assert_eq!(roundtrip(&data, 3), data);
}
#[test]
fn mixed_data_roundtrip() {
let mut data = Vec::new();
for i in 0..1000 {
data.extend_from_slice(format!("record-{i:04}-").as_bytes());
data.extend_from_slice(&[(i & 0xff) as u8; 8]);
data.push(0x00);
}
assert_eq!(roundtrip(&data, 6), data);
}
#[test]
fn empty_and_tiny_inputs() {
assert_eq!(roundtrip(&[], 3), []);
assert_eq!(roundtrip(b"a", 3), b"a");
assert_eq!(roundtrip(b"ab", 3), b"ab");
assert!(compress(b"ab", 3).is_empty());
}
#[test]
fn malformed_streams_are_rejected() {
assert_eq!(decompress(&[0x01], 0), Err(LzError::UnexpectedEnd));
let mut frame = Vec::new();
frame.extend_from_slice(&10u32.to_le_bytes());
frame.push(0); frame.push(0x41); assert_eq!(decompress(&frame, 1), Err(LzError::LengthMismatch));
let mut frame = Vec::new();
frame.extend_from_slice(&3u32.to_le_bytes());
frame.push(0x01); frame.push(0x00); frame.push(0x00); frame.push(0x00); assert_eq!(decompress(&frame, 3), Err(LzError::InvalidDistance));
}
#[test]
fn level_semantics_are_stable() {
let data = vec![0x5A_u8; 50_000];
let low = compress(&data, 1);
let high = compress(&data, 22);
assert!(!low.is_empty() && !high.is_empty());
assert_eq!(decompress(&low, data.len()).unwrap(), data);
assert_eq!(decompress(&high, data.len()).unwrap(), data);
}
}