use crate::block::{DecompressError, MINMATCH};
use crate::sink::{Sink, SliceSink, VecSink};
use alloc::vec::Vec;
#[inline]
unsafe fn duplicate(
output_ptr: &mut *mut u8,
output_end: *mut u8,
start: *const u8,
match_length: usize,
) {
if (output_ptr.offset_from(start) as usize) < match_length + 16 - 1
|| (output_end.offset_from(*output_ptr) as usize) < match_length + 16 - 1
{
duplicate_overlapping(output_ptr, start, match_length);
} else {
debug_assert!(
output_ptr.add(match_length / 16 * 16 + ((match_length % 16) != 0) as usize * 16)
<= output_end
);
wild_copy_from_src_16(start, *output_ptr, match_length);
*output_ptr = output_ptr.add(match_length);
}
}
#[inline]
fn wild_copy_from_src_16(mut source: *const u8, mut dst_ptr: *mut u8, num_items: usize) {
unsafe {
let dst_ptr_end = dst_ptr.add(num_items);
while (dst_ptr as usize) < dst_ptr_end as usize {
core::ptr::copy_nonoverlapping(source, dst_ptr, 16);
source = source.add(16);
dst_ptr = dst_ptr.add(16);
}
}
}
#[inline]
unsafe fn duplicate_overlapping(
output_ptr: &mut *mut u8,
mut start: *const u8,
match_length: usize,
) {
output_ptr.write(0u8);
for _ in 0..match_length {
let curr = start.read();
output_ptr.write(curr);
*output_ptr = output_ptr.add(1);
start = start.add(1);
}
}
#[inline]
unsafe fn copy_from_dict(
output_base: *mut u8,
output_ptr: &mut *mut u8,
ext_dict: &[u8],
offset: usize,
match_length: usize,
) -> usize {
debug_assert!(output_ptr.offset_from(output_base) >= 0);
debug_assert!(offset > output_ptr.offset_from(output_base) as usize);
debug_assert!(ext_dict.len() + output_ptr.offset_from(output_base) as usize >= offset);
let dict_offset = ext_dict.len() + output_ptr.offset_from(output_base) as usize - offset;
let dict_match_length = match_length.min(ext_dict.len() - dict_offset);
core::ptr::copy_nonoverlapping(
ext_dict.as_ptr().add(dict_offset),
*output_ptr,
dict_match_length,
);
*output_ptr = output_ptr.add(dict_match_length);
dict_match_length
}
#[inline]
fn read_integer(input: &[u8], input_pos: &mut usize) -> Result<u32, DecompressError> {
let mut n: u32 = 0;
loop {
#[cfg(feature = "checked-decode")]
{
if *input_pos >= input.len() {
return Err(DecompressError::ExpectedAnotherByte);
}
}
let extra = *unsafe { input.get_unchecked(*input_pos) };
*input_pos += 1;
n += extra as u32;
if extra != 0xFF {
break;
}
}
Ok(n)
}
#[inline]
fn read_u16(input: &[u8], input_pos: &mut usize) -> u16 {
let mut num: u16 = 0;
unsafe {
core::ptr::copy_nonoverlapping(
input.as_ptr().add(*input_pos),
&mut num as *mut u16 as *mut u8,
2,
);
}
*input_pos += 2;
u16::from_le(num)
}
const FIT_TOKEN_MASK_LITERAL: u8 = 0b00001111;
const FIT_TOKEN_MASK_MATCH: u8 = 0b11110000;
#[test]
fn check_token() {
assert_eq!(does_token_fit(15), false);
assert_eq!(does_token_fit(14), true);
assert_eq!(does_token_fit(114), true);
assert_eq!(does_token_fit(0b11110000), false);
assert_eq!(does_token_fit(0b10110000), true);
}
#[inline]
fn does_token_fit(token: u8) -> bool {
!((token & FIT_TOKEN_MASK_LITERAL) == FIT_TOKEN_MASK_LITERAL
|| (token & FIT_TOKEN_MASK_MATCH) == FIT_TOKEN_MASK_MATCH)
}
#[inline]
pub(crate) fn decompress_internal<SINK: Sink, const USE_DICT: bool>(
input: &[u8],
output: &mut SINK,
ext_dict: &[u8],
) -> Result<usize, DecompressError> {
#[cfg(not(feature = "checked-decode"))]
{
if input.is_empty() {
return Err(DecompressError::ExpectedAnotherByte);
}
}
let ext_dict = if USE_DICT {
ext_dict
} else {
debug_assert!(ext_dict.is_empty());
&[]
};
let output_base = unsafe { output.base_mut_ptr() };
let output_end = unsafe { output_base.add(output.capacity()) };
let output_start_pos_ptr = unsafe { output.pos_mut_ptr() };
let mut output_ptr = output_start_pos_ptr;
let mut input_pos = 0;
let safe_input_pos = input
.len()
.saturating_sub(16 + 2 );
let safe_output_ptr = unsafe {
output_base.add(
output
.capacity()
.saturating_sub(16 + 18 ),
)
};
loop {
#[cfg(feature = "checked-decode")]
{
if input_pos >= input.len() {
return Err(DecompressError::ExpectedAnotherByte);
}
}
let token = unsafe { *input.get_unchecked(input_pos) };
input_pos += 1;
if does_token_fit(token) && input_pos <= safe_input_pos && output_ptr < safe_output_ptr {
let literal_length = (token >> 4) as usize;
let mut match_length = MINMATCH + (token & 0xF) as usize;
debug_assert!(
unsafe { output_ptr.add(literal_length + match_length) } <= output_end,
"{} wont fit ",
literal_length + match_length
);
#[cfg(feature = "checked-decode")]
{
if literal_length > input.len() - input_pos {
return Err(DecompressError::OffsetOutOfBounds);
}
}
unsafe {
core::ptr::copy_nonoverlapping(input.as_ptr().add(input_pos), output_ptr, 16);
}
input_pos += literal_length;
unsafe {
output_ptr = output_ptr.add(literal_length);
}
debug_assert!(input.len() - input_pos >= 2);
let offset = read_u16(input, &mut input_pos) as usize;
let output_len = unsafe { output_ptr.offset_from(output_base) as usize };
#[cfg(feature = "checked-decode")]
{
if offset > output_len + ext_dict.len() {
return Err(DecompressError::OffsetOutOfBounds);
}
}
if USE_DICT && offset > output_len {
let copied = unsafe {
copy_from_dict(output_base, &mut output_ptr, ext_dict, offset, match_length)
};
if copied == match_length {
continue;
}
match_length -= copied;
}
let start_ptr = unsafe { output_ptr.sub(offset) };
debug_assert!(start_ptr >= output_base);
debug_assert!(start_ptr < output_end);
debug_assert!(unsafe { output_end.offset_from(start_ptr) as usize } >= match_length);
if offset >= match_length {
unsafe {
core::ptr::copy(start_ptr, output_ptr, 18);
output_ptr = output_ptr.add(match_length);
}
} else {
unsafe {
duplicate_overlapping(&mut output_ptr, start_ptr, match_length);
}
}
continue;
}
let mut literal_length = (token >> 4) as usize;
if literal_length != 0 {
if literal_length == 15 {
literal_length += read_integer(input, &mut input_pos)? as usize;
}
#[cfg(feature = "checked-decode")]
{
if literal_length > input.len() - input_pos {
return Err(DecompressError::LiteralOutOfBounds);
}
if literal_length > unsafe { output_end.offset_from(output_ptr) as usize } {
return Err(DecompressError::OutputTooSmall {
expected: unsafe { output_ptr.offset_from(output_base) as usize }
+ literal_length,
actual: output.capacity(),
});
}
}
unsafe {
core::ptr::copy_nonoverlapping(
input.as_ptr().add(input_pos),
output_ptr,
literal_length,
);
output_ptr = output_ptr.add(literal_length);
}
input_pos += literal_length;
}
if input_pos >= input.len() {
break;
}
#[cfg(feature = "checked-decode")]
{
if input.len() - input_pos < 2 {
return Err(DecompressError::ExpectedAnotherByte);
}
}
let offset = read_u16(input, &mut input_pos) as usize;
let mut match_length = MINMATCH + (token & 0xF) as usize;
if match_length == MINMATCH + 15 {
match_length += read_integer(input, &mut input_pos)? as usize;
}
let output_len = unsafe { output_ptr.offset_from(output_base) as usize };
#[cfg(feature = "checked-decode")]
{
if offset > output_len + ext_dict.len() {
return Err(DecompressError::OffsetOutOfBounds);
}
if match_length > unsafe { output_end.offset_from(output_ptr) as usize } {
return Err(DecompressError::OutputTooSmall {
expected: output_len + match_length,
actual: output.capacity(),
});
}
}
if USE_DICT && offset > output_len {
let copied = unsafe {
copy_from_dict(output_base, &mut output_ptr, ext_dict, offset, match_length)
};
if copied == match_length {
continue;
}
match_length -= copied;
}
let start_ptr = unsafe { output_ptr.sub(offset) };
debug_assert!(start_ptr >= output_base);
debug_assert!(start_ptr < output_end);
debug_assert!(unsafe { output_end.offset_from(start_ptr) as usize } >= match_length);
unsafe {
duplicate(&mut output_ptr, output_end, start_ptr, match_length);
}
}
unsafe {
output.set_pos(output_ptr.offset_from(output_base) as usize);
Ok(output_ptr.offset_from(output_start_pos_ptr) as usize)
}
}
#[inline]
pub fn decompress_into(input: &[u8], output: &mut [u8]) -> Result<usize, DecompressError> {
decompress_internal::<_, false>(input, &mut SliceSink::new(output, 0), b"")
}
#[inline]
pub fn decompress_into_with_dict(
input: &[u8],
output: &mut [u8],
ext_dict: &[u8],
) -> Result<usize, DecompressError> {
decompress_internal::<_, true>(input, &mut SliceSink::new(output, 0), ext_dict)
}
#[inline]
pub fn decompress_size_prepended(input: &[u8]) -> Result<Vec<u8>, DecompressError> {
let (uncompressed_size, input) = super::uncompressed_size(input)?;
decompress(input, uncompressed_size)
}
#[inline]
pub fn decompress(input: &[u8], uncompressed_size: usize) -> Result<Vec<u8>, DecompressError> {
let mut vec: Vec<u8> = Vec::with_capacity(uncompressed_size);
let decomp_len =
decompress_internal::<_, false>(input, &mut VecSink::new(&mut vec, 0, 0), b"")?;
if decomp_len != uncompressed_size {
return Err(DecompressError::UncompressedSizeDiffers {
expected: uncompressed_size,
actual: decomp_len,
});
}
Ok(vec)
}
#[inline]
pub fn decompress_size_prepended_with_dict(
input: &[u8],
ext_dict: &[u8],
) -> Result<Vec<u8>, DecompressError> {
let (uncompressed_size, input) = super::uncompressed_size(input)?;
decompress_with_dict(input, uncompressed_size, ext_dict)
}
#[inline]
pub fn decompress_with_dict(
input: &[u8],
uncompressed_size: usize,
ext_dict: &[u8],
) -> Result<Vec<u8>, DecompressError> {
let mut vec: Vec<u8> = Vec::with_capacity(uncompressed_size);
let decomp_len =
decompress_internal::<_, true>(input, &mut VecSink::new(&mut vec, 0, 0), ext_dict)?;
if decomp_len != uncompressed_size {
return Err(DecompressError::UncompressedSizeDiffers {
expected: uncompressed_size,
actual: decomp_len,
});
}
Ok(vec)
}
#[cfg(test)]
mod test {
use super::*;
#[test]
fn all_literal() {
assert_eq!(decompress(&[0x30, b'a', b'4', b'9'], 3).unwrap(), b"a49");
}
#[cfg(feature = "checked-decode")]
#[test]
fn offset_oob() {
decompress(&[0x10, b'a', 2, 0], 4).unwrap_err();
decompress(&[0x40, b'a', 1, 0], 4).unwrap_err();
}
}