use crate::{Error, Result};
#[cfg(feature = "alloc")]
extern crate alloc;
#[allow(dead_code)]
pub const DEFAULT_LEVEL: i32 = 3;
const ZSTD_MAGIC: [u8; 4] = [0x28, 0xB5, 0x2F, 0xFD];
const ZSTD_MAX_BLOCK_SIZE: usize = 128 * 1024;
const FHD_SINGLE_SEGMENT: u8 = 0x20;
const FHD_DICT_ID_FLAG_MASK: u8 = 0x03;
const FHD_CONTENT_SIZE_FLAG_MASK: u8 = 0xC0;
#[cfg(feature = "alloc")]
pub fn compress(data: &[u8], level: i32) -> Result<alloc::vec::Vec<u8>> {
let level = level.clamp(1, 22);
oxiarc_zstd::compress_with_level(data, level).map_err(|e| Error::OwnedCustom {
message: alloc::format!("Zstd compression error: {e}"),
})
}
#[cfg(feature = "alloc")]
pub fn decompress(data: &[u8], max_output: usize) -> Result<alloc::vec::Vec<u8>> {
validate_zstd_frame_bound(data, max_output)?;
oxiarc_zstd::decompress(data).map_err(|e| Error::OwnedCustom {
message: alloc::format!("Zstd decompression error: {e}"),
})
}
#[cfg(feature = "alloc")]
fn validate_zstd_frame_bound(data: &[u8], max_output: usize) -> Result<()> {
if data.len() < 5 {
return Err(Error::UnexpectedEnd {
additional: 5usize.saturating_sub(data.len()),
});
}
if data[0..4] != ZSTD_MAGIC {
return Err(Error::InvalidData {
message: "invalid Zstd frame magic",
});
}
let descriptor = data[4];
let single_segment = (descriptor & FHD_SINGLE_SEGMENT) != 0;
let dict_id_flag = descriptor & FHD_DICT_ID_FLAG_MASK;
let content_size_flag = (descriptor & FHD_CONTENT_SIZE_FLAG_MASK) >> 6;
let mut pos = 5usize;
if !single_segment {
if data.len() <= pos {
return Err(Error::UnexpectedEnd { additional: 1 });
}
pos += 1;
}
let dict_id_bytes = match dict_id_flag {
0 => 0usize,
1 => 1,
2 => 2,
_ => 4,
};
pos = pos.checked_add(dict_id_bytes).ok_or(Error::InvalidData {
message: "malformed Zstd frame header",
})?;
if data.len() < pos {
return Err(Error::UnexpectedEnd {
additional: pos - data.len(),
});
}
let content_size = if single_segment || content_size_flag != 0 {
let size_bytes = match content_size_flag {
0 => 1usize, 1 => 2,
2 => 4,
_ => 8,
};
let end = pos.checked_add(size_bytes).ok_or(Error::InvalidData {
message: "malformed Zstd frame header",
})?;
if data.len() < end {
return Err(Error::UnexpectedEnd {
additional: end - data.len(),
});
}
let size = match size_bytes {
1 => data[pos] as u64,
2 => u16::from_le_bytes([data[pos], data[pos + 1]]) as u64 + 256,
4 => {
u32::from_le_bytes([data[pos], data[pos + 1], data[pos + 2], data[pos + 3]]) as u64
}
_ => u64::from_le_bytes([
data[pos],
data[pos + 1],
data[pos + 2],
data[pos + 3],
data[pos + 4],
data[pos + 5],
data[pos + 6],
data[pos + 7],
]),
};
pos = end;
size
} else {
return Err(Error::InvalidData {
message: "Zstd frame omits content size; refusing to decompress (possible decompression bomb)",
});
};
let max_output_u64 = max_output as u64;
if content_size > max_output_u64 {
return Err(Error::LimitExceeded {
limit: max_output_u64,
found: content_size,
});
}
let regen_ceiling = content_size.saturating_add(ZSTD_MAX_BLOCK_SIZE as u64);
let mut total: u64 = 0;
loop {
if data.len() < pos + 3 {
return Err(Error::UnexpectedEnd {
additional: (pos + 3).saturating_sub(data.len()),
});
}
let block_header = u32::from_le_bytes([data[pos], data[pos + 1], data[pos + 2], 0]);
pos += 3;
let last_block = (block_header & 1) != 0;
let block_type = (block_header >> 1) & 0x03;
let block_size = ((block_header >> 3) & 0x1F_FFFF) as usize;
if block_size > ZSTD_MAX_BLOCK_SIZE {
return Err(Error::InvalidData {
message: "Zstd block size exceeds maximum",
});
}
let (regen_bound, stored) = match block_type {
0 => (block_size as u64, block_size), 1 => (block_size as u64, 1usize), 2 => (ZSTD_MAX_BLOCK_SIZE as u64, block_size), _ => {
return Err(Error::InvalidData {
message: "Zstd reserved block type",
});
}
};
total = total.saturating_add(regen_bound);
if total > regen_ceiling {
return Err(Error::LimitExceeded {
limit: regen_ceiling,
found: total,
});
}
pos = pos.checked_add(stored).ok_or(Error::InvalidData {
message: "malformed Zstd block header",
})?;
if data.len() < pos {
return Err(Error::UnexpectedEnd {
additional: pos - data.len(),
});
}
if last_block {
break;
}
}
Ok(())
}
#[cfg(feature = "alloc")]
#[allow(dead_code)]
pub fn compress_default(data: &[u8]) -> Result<alloc::vec::Vec<u8>> {
compress(data, DEFAULT_LEVEL)
}
#[cfg(test)]
mod tests {
use super::*;
#[cfg(feature = "alloc")]
const TEST_CAP: usize = 8 * 1024 * 1024;
#[cfg(feature = "alloc")]
#[test]
fn test_compress_decompress() {
let data = b"Hello, World! This is a test of Zstd compression.";
let compressed = compress(data, DEFAULT_LEVEL).expect("compress failed");
let decompressed = decompress(&compressed, TEST_CAP).expect("decompress failed");
assert_eq!(data.as_slice(), decompressed.as_slice());
}
#[cfg(feature = "alloc")]
#[test]
fn test_empty_data() {
let data: &[u8] = b"";
let compressed = compress(data, DEFAULT_LEVEL).expect("compress failed");
let decompressed = decompress(&compressed, TEST_CAP).expect("decompress failed");
assert_eq!(data, decompressed.as_slice());
}
#[cfg(feature = "alloc")]
#[test]
fn test_large_data() {
let data: alloc::vec::Vec<u8> = (0..100000).map(|i| (i % 256) as u8).collect();
let compressed = compress(&data, DEFAULT_LEVEL).expect("compress failed");
let decompressed = decompress(&compressed, TEST_CAP).expect("decompress failed");
assert_eq!(data, decompressed);
assert!(compressed.len() < data.len());
}
#[cfg(feature = "alloc")]
#[test]
fn test_compression_levels() {
let data: alloc::vec::Vec<u8> = (0..10000).map(|i| (i % 256) as u8).collect();
for level in [1, 3, 9, 19, 22] {
let compressed = compress(&data, level).expect("compress failed");
let decompressed = decompress(&compressed, TEST_CAP).expect("decompress failed");
assert_eq!(data, decompressed);
}
}
#[cfg(feature = "alloc")]
#[test]
fn test_rle_bomb_rejected_by_cap() {
let mut frame = alloc::vec::Vec::new();
frame.extend_from_slice(&ZSTD_MAGIC);
frame.push(FHD_SINGLE_SEGMENT); frame.push(1u8);
let regen: u32 = ZSTD_MAX_BLOCK_SIZE as u32;
for i in 0..8u32 {
let last = if i == 7 { 1 } else { 0 };
let block_header: u32 = (regen << 3) | (1 << 1) | last; frame.push((block_header & 0xFF) as u8);
frame.push(((block_header >> 8) & 0xFF) as u8);
frame.push(((block_header >> 16) & 0xFF) as u8);
frame.push(0xAB); }
let err = decompress(&frame, 64 * 1024).expect_err("bomb must be rejected");
assert!(matches!(err, Error::LimitExceeded { .. }));
}
#[cfg(feature = "alloc")]
#[test]
fn test_omitted_content_size_rejected() {
let mut frame = alloc::vec::Vec::new();
frame.extend_from_slice(&ZSTD_MAGIC);
frame.push(0u8); frame.push(0u8); let block_header: u32 = 1; frame.push((block_header & 0xFF) as u8);
frame.push(((block_header >> 8) & 0xFF) as u8);
frame.push(((block_header >> 16) & 0xFF) as u8);
let err = decompress(&frame, TEST_CAP).expect_err("omitted content size must be rejected");
assert!(matches!(err, Error::InvalidData { .. }));
}
#[cfg(feature = "alloc")]
#[test]
fn test_level_clamping() {
let data = b"test data";
let _ = compress(data, 0).expect("compress failed");
let _ = compress(data, 30).expect("compress failed");
}
}