use super::super::blocks::literals_section::{LiteralsSection, LiteralsSectionType};
use super::buffer_backend::WILDCOPY_OVERLENGTH;
use super::scratch::HuffmanScratch;
use crate::bit_io::BitReaderReversed;
#[cfg(all(target_arch = "x86_64", feature = "kernel-avx2"))]
use crate::cpu_kernel::Avx2Kernel;
#[cfg(all(
any(target_arch = "x86", target_arch = "x86_64"),
feature = "kernel-bmi2"
))]
use crate::cpu_kernel::Bmi2Kernel;
#[cfg(all(target_arch = "aarch64", feature = "kernel-neon"))]
use crate::cpu_kernel::NeonKernel;
#[cfg(all(
target_arch = "aarch64",
feature = "kernel-sve",
any(feature = "std", target_feature = "sve"),
))]
use crate::cpu_kernel::SveKernel;
#[cfg(all(target_arch = "x86_64", feature = "kernel-vbmi2"))]
use crate::cpu_kernel::Vbmi2Kernel;
#[cfg(test)]
use crate::cpu_kernel::detect_cpu_kernel;
use crate::cpu_kernel::{CpuKernel, CpuKernelTag, ScalarKernel};
use crate::decoding::dictionary::Dictionary;
use crate::decoding::errors::DecompressLiteralsError;
use crate::huff0::HuffmanDecoder;
use alloc::vec::Vec;
#[cfg(test)]
pub fn decode_literals(
section: &LiteralsSection,
scratch: &mut HuffmanScratch,
dict: Option<&Dictionary>,
source: &[u8],
target: &mut Vec<u8>,
) -> Result<u32, DecompressLiteralsError> {
match section.ls_type {
LiteralsSectionType::Raw => {
target.extend(&source[0..section.regenerated_size as usize]);
Ok(section.regenerated_size)
}
LiteralsSectionType::RLE => {
target.resize(target.len() + section.regenerated_size as usize, source[0]);
Ok(1)
}
LiteralsSectionType::Compressed | LiteralsSectionType::Treeless => {
let bytes_read =
decompress_literals(section, scratch, dict, source, target, detect_cpu_kernel())?;
Ok(bytes_read)
}
}
}
pub struct LiteralsView<'a> {
pub data: &'a [u8],
pub len: usize,
pub bytes_used: u32,
}
#[allow(clippy::too_many_arguments)]
pub fn decode_literals_zerocopy<'a>(
section: &LiteralsSection,
scratch: &mut HuffmanScratch,
dict: Option<&Dictionary>,
source: &'a [u8],
payload_len: usize,
needs_slack: bool,
target: &'a mut Vec<u8>,
kernel: CpuKernelTag,
) -> Result<LiteralsView<'a>, DecompressLiteralsError> {
let base = target.len();
match section.ls_type {
LiteralsSectionType::Raw => {
let n = section.regenerated_size as usize;
if source.len() < n {
return Err(DecompressLiteralsError::MissingBytesForLiterals {
got: source.len(),
needed: n,
});
}
if !needs_slack || source.len() >= n + WILDCOPY_OVERLENGTH {
let end = if needs_slack {
n + WILDCOPY_OVERLENGTH
} else {
n
};
return Ok(LiteralsView {
data: &source[..end],
len: n,
bytes_used: section.regenerated_size,
});
}
target.extend_from_slice(&source[..n]);
pad_with_slack(target, base, n, bytes_used_raw(section), needs_slack)
}
LiteralsSectionType::RLE => {
if source.is_empty() {
return Err(DecompressLiteralsError::MissingBytesForLiterals { got: 0, needed: 1 });
}
let n = section.regenerated_size as usize;
target.resize(base + n, source[0]);
pad_with_slack(target, base, n, 1, needs_slack)
}
LiteralsSectionType::Compressed | LiteralsSectionType::Treeless => {
let payload = &source[..payload_len.min(source.len())];
let bytes_used = decompress_literals(section, scratch, dict, payload, target, kernel)?;
let n = target.len() - base;
pad_with_slack(target, base, n, bytes_used, needs_slack)
}
}
}
fn bytes_used_raw(section: &LiteralsSection) -> u32 {
section.regenerated_size
}
fn pad_with_slack(
target: &mut Vec<u8>,
base: usize,
n: usize,
bytes_used: u32,
needs_slack: bool,
) -> Result<LiteralsView<'_>, DecompressLiteralsError> {
if needs_slack {
target.resize(base + n + WILDCOPY_OVERLENGTH, 0);
}
Ok(LiteralsView {
data: &target[base..],
len: n,
bytes_used,
})
}
fn decompress_literals(
section: &LiteralsSection,
scratch: &mut HuffmanScratch,
dict: Option<&Dictionary>,
source: &[u8],
target: &mut Vec<u8>,
kernel: CpuKernelTag,
) -> Result<u32, DecompressLiteralsError> {
match kernel {
#[cfg(all(target_arch = "x86_64", feature = "kernel-vbmi2"))]
CpuKernelTag::Vbmi2 => unsafe {
decompress_literals_vbmi2(section, scratch, dict, source, target)
},
#[cfg(all(target_arch = "x86_64", feature = "kernel-avx2"))]
CpuKernelTag::Avx2 => unsafe {
decompress_literals_avx2(section, scratch, dict, source, target)
},
#[cfg(all(
any(target_arch = "x86", target_arch = "x86_64"),
feature = "kernel-bmi2"
))]
CpuKernelTag::Bmi2 => unsafe {
decompress_literals_bmi2(section, scratch, dict, source, target)
},
#[cfg(all(target_arch = "aarch64", feature = "kernel-neon"))]
CpuKernelTag::Neon => {
decompress_literals_impl::<NeonKernel>(section, scratch, dict, source, target)
}
#[cfg(all(
target_arch = "aarch64",
feature = "kernel-sve",
any(feature = "std", target_feature = "sve"),
))]
CpuKernelTag::Sve => {
decompress_literals_impl::<SveKernel>(section, scratch, dict, source, target)
}
_ => decompress_literals_impl::<ScalarKernel>(section, scratch, dict, source, target),
}
}
#[cfg(all(target_arch = "x86_64", feature = "kernel-avx2"))]
#[target_feature(enable = "bmi2,avx2")]
unsafe fn decompress_literals_avx2(
section: &LiteralsSection,
scratch: &mut HuffmanScratch,
dict: Option<&Dictionary>,
source: &[u8],
target: &mut Vec<u8>,
) -> Result<u32, DecompressLiteralsError> {
decompress_literals_impl::<Avx2Kernel>(section, scratch, dict, source, target)
}
#[cfg(all(
any(target_arch = "x86", target_arch = "x86_64"),
feature = "kernel-bmi2"
))]
#[target_feature(enable = "bmi2")]
unsafe fn decompress_literals_bmi2(
section: &LiteralsSection,
scratch: &mut HuffmanScratch,
dict: Option<&Dictionary>,
source: &[u8],
target: &mut Vec<u8>,
) -> Result<u32, DecompressLiteralsError> {
decompress_literals_impl::<Bmi2Kernel>(section, scratch, dict, source, target)
}
#[cfg(all(target_arch = "x86_64", feature = "kernel-vbmi2"))]
#[target_feature(enable = "avx512vbmi2,avx512f,avx512vl,avx512bw,bmi2,avx2")]
unsafe fn decompress_literals_vbmi2(
section: &LiteralsSection,
scratch: &mut HuffmanScratch,
dict: Option<&Dictionary>,
source: &[u8],
target: &mut Vec<u8>,
) -> Result<u32, DecompressLiteralsError> {
decompress_literals_impl::<Vbmi2Kernel>(section, scratch, dict, source, target)
}
fn decompress_literals_impl<K: CpuKernel>(
section: &LiteralsSection,
scratch: &mut HuffmanScratch,
dict: Option<&Dictionary>,
source: &[u8],
target: &mut Vec<u8>,
) -> Result<u32, DecompressLiteralsError> {
use DecompressLiteralsError as err;
let compressed_size = section.compressed_size.ok_or(err::MissingCompressedSize)? as usize;
let num_streams = section.num_streams.ok_or(err::MissingNumStreams)?;
let base = target.len();
let regen = section.regenerated_size as usize;
target.reserve(regen);
let source = source
.get(..compressed_size)
.ok_or(err::MissingBytesForLiterals {
got: source.len(),
needed: compressed_size,
})?;
let mut bytes_read = 0;
match section.ls_type {
LiteralsSectionType::Compressed => {
bytes_read += scratch.table.build_decoder(source)?;
scratch.mark_table_local();
vprintln!("Built huffman table using {} bytes", bytes_read);
}
LiteralsSectionType::Treeless if scratch.huf_table(dict).max_num_bits == 0 => {
return Err(err::UninitializedHuffmanTable);
}
_ => { }
}
let source = &source[bytes_read as usize..];
let table = scratch.huf_table(dict);
if num_streams == 4 {
if source.len() < 6 {
return Err(err::MissingBytesForJumpHeader { got: source.len() });
}
let jump1 = source[0] as usize + ((source[1] as usize) << 8);
let jump2 = jump1 + source[2] as usize + ((source[3] as usize) << 8);
let jump3 = jump2 + source[4] as usize + ((source[5] as usize) << 8);
bytes_read += 6;
let source = &source[6..];
if source.len() < jump3 {
return Err(err::MissingBytesForLiterals {
got: source.len(),
needed: jump3,
});
}
let streams: [&[u8]; 4] = [
&source[..jump1],
&source[jump1..jump2],
&source[jump2..jump3],
&source[jump3..],
];
let mut decoders: [HuffmanDecoder<'_>; 4] = [
HuffmanDecoder::new(table),
HuffmanDecoder::new(table),
HuffmanDecoder::new(table),
HuffmanDecoder::new(table),
];
let mut brs: [BitReaderReversed<'_, K>; 4] = [
BitReaderReversed::<K>::new(streams[0]),
BitReaderReversed::<K>::new(streams[1]),
BitReaderReversed::<K>::new(streams[2]),
BitReaderReversed::<K>::new(streams[3]),
];
for i in 0..4 {
let mut skipped_bits = 0;
loop {
let val = brs[i].get_bits(1);
skipped_bits += 1;
if val == 1 || skipped_bits > 8 {
break;
}
}
if skipped_bits > 8 {
return Err(DecompressLiteralsError::ExtraPadding { skipped_bits });
}
decoders[i].init_state(&mut brs[i]);
}
let max_bits = table.max_num_bits as isize;
let seg = regen.div_ceil(4);
let target_ptr: *mut u8 = target.as_mut_ptr();
let limit = base + regen;
let starts: [usize; 4] = [
base,
(base + seg).min(limit),
(base + 2 * seg).min(limit),
(base + 3 * seg).min(limit),
];
let ends: [usize; 4] = [starts[1], starts[2], starts[3], limit];
let mut cursors = starts;
let max_num_bits = table.max_num_bits;
let symbols_per_burst: usize = (63 - 8) / max_num_bits as usize;
let burst_bits = (symbols_per_burst * max_num_bits as usize) as u8;
let table_shift = (64 - max_num_bits) as u32;
let packed = table.packed_decode.as_slice();
let min_seg_len = (ends[0] - starts[0])
.min(ends[1] - starts[1])
.min(ends[2] - starts[2])
.min(ends[3] - starts[3]);
let burst_eligible = symbols_per_burst >= 1 && min_seg_len >= symbols_per_burst;
let cursor_burst_ceil = (starts[0] + min_seg_len).saturating_sub(symbols_per_burst);
let bounds = LoopBounds {
symbols_per_burst,
burst_bits,
table_shift,
cursor_burst_ceil,
burst_eligible,
alloc_upper_bound: base + regen,
};
unsafe {
run_4stream_burst_loop(
&mut decoders,
&mut brs,
target_ptr,
packed,
&mut cursors,
&bounds,
);
}
let group_bits = 4 * max_num_bits;
for i in 0..4 {
while cursors[i] + 4 <= ends[i] {
brs[i].ensure_bits(group_bits);
for _ in 0..4 {
let byte = decoders[i].decode_symbol_and_advance_no_refill(&mut brs[i]);
unsafe {
target_ptr.add(cursors[i]).write(byte);
}
cursors[i] += 1;
}
}
if cursors[i] < ends[i] {
brs[i].ensure_bits(group_bits);
while cursors[i] < ends[i] {
let byte = decoders[i].decode_symbol_and_advance_no_refill(&mut brs[i]);
unsafe {
target_ptr.add(cursors[i]).write(byte);
}
cursors[i] += 1;
}
}
if brs[i].bits_remaining() != -max_bits {
return Err(DecompressLiteralsError::BitstreamReadMismatch {
read_til: brs[i].bits_remaining(),
expected: -max_bits,
});
}
}
let decoded: usize = cursors.iter().zip(starts.iter()).map(|(c, s)| c - s).sum();
if decoded != regen {
return Err(DecompressLiteralsError::DecodedLiteralCountMismatch {
decoded,
expected: regen,
});
}
unsafe {
target.set_len(base + regen);
}
bytes_read += source.len() as u32;
} else {
assert!(num_streams == 1);
let mut decoder = HuffmanDecoder::new(table);
let mut br = BitReaderReversed::<K>::new(source);
let mut skipped_bits = 0;
loop {
let val = br.get_bits(1);
skipped_bits += 1;
if val == 1 || skipped_bits > 8 {
break;
}
}
if skipped_bits > 8 {
return Err(DecompressLiteralsError::ExtraPadding { skipped_bits });
}
decoder.init_state(&mut br);
while br.bits_remaining() > -(table.max_num_bits as isize) {
target.push(decoder.decode_symbol_and_advance(&mut br));
}
let expected = -(table.max_num_bits as isize);
if br.bits_remaining() != expected {
target.truncate(base);
return Err(DecompressLiteralsError::BitstreamReadMismatch {
read_til: br.bits_remaining(),
expected,
});
}
bytes_read += source.len() as u32;
}
if target.len() != base + regen {
let decoded = target.len() - base;
target.truncate(base);
return Err(DecompressLiteralsError::DecodedLiteralCountMismatch {
decoded,
expected: regen,
});
}
Ok(bytes_read)
}
#[derive(Copy, Clone)]
struct LoopBounds {
symbols_per_burst: usize,
burst_bits: u8,
table_shift: u32,
cursor_burst_ceil: usize,
burst_eligible: bool,
alloc_upper_bound: usize,
}
#[inline(always)]
unsafe fn run_4stream_burst_loop<K: CpuKernel>(
decoders: &mut [HuffmanDecoder<'_>; 4],
brs: &mut [BitReaderReversed<'_, K>; 4],
target_ptr: *mut u8,
packed: &[u16],
cursors: &mut [usize; 4],
bounds: &LoopBounds,
) {
let LoopBounds {
symbols_per_burst,
burst_bits,
table_shift,
cursor_burst_ceil,
burst_eligible,
alloc_upper_bound,
} = *bounds;
let max_num_bits = (64 - table_shift) as u8;
if !burst_eligible {
return;
}
debug_assert!(
cursor_burst_ceil + symbols_per_burst <= alloc_upper_bound,
"caller must size the target allocation so the lockstep-advanced \
cursors stay within bounds across a full burst",
);
let mut b0 = (brs[0].bit_container | 1) << (brs[0].bits_consumed - max_num_bits);
let mut b1 = (brs[1].bit_container | 1) << (brs[1].bits_consumed - max_num_bits);
let mut b2 = (brs[2].bit_container | 1) << (brs[2].bits_consumed - max_num_bits);
let mut b3 = (brs[3].bit_container | 1) << (brs[3].bits_consumed - max_num_bits);
let mut ip0 = brs[0].index;
let mut ip1 = brs[1].index;
let mut ip2 = brs[2].index;
let mut ip3 = brs[3].index;
let mut c0 = cursors[0];
let mut c1 = cursors[1];
let mut c2 = cursors[2];
let mut c3 = cursors[3];
let src0 = brs[0].source;
let src1 = brs[1].source;
let src2 = brs[2].source;
let src3 = brs[3].source;
macro_rules! decode1 {
($b:ident, $c:ident) => {{
let idx = ($b >> table_shift) as usize;
let entry = unsafe { *packed.get_unchecked(idx) };
unsafe { target_ptr.add($c).write((entry & 0xFF) as u8) };
$c += 1;
$b <<= (entry >> 8) & 0xFF;
}};
}
macro_rules! reload1 {
($b:ident, $ip:ident, $src:expr) => {{
let ctz = $b.trailing_zeros();
$ip -= (ctz >> 3) as usize;
let nb_bits = (ctz & 7) as u8;
let new_window = u64::from_le_bytes(unsafe {
$src.get_unchecked($ip..$ip + 8)
.try_into()
.unwrap_unchecked()
});
$b = (new_window | 1) << nb_bits;
}};
}
macro_rules! burst {
($n:literal) => {{
for _ in 0..$n {
decode1!(b0, c0);
decode1!(b1, c1);
decode1!(b2, c2);
decode1!(b3, c3);
}
}};
}
let bytes_per_iter_upper = (8 + burst_bits as usize) / 8;
let mut any_iter = false;
while c0 <= cursor_burst_ceil {
let min_ip = ip0.min(ip1).min(ip2).min(ip3);
if min_ip < bytes_per_iter_upper {
break;
}
any_iter = true;
match symbols_per_burst {
5 => burst!(5),
6 => burst!(6),
7 => burst!(7),
_ => {
for _ in 0..symbols_per_burst {
decode1!(b0, c0);
decode1!(b1, c1);
decode1!(b2, c2);
decode1!(b3, c3);
}
}
}
reload1!(b0, ip0, src0);
reload1!(b1, ip1, src1);
reload1!(b2, ip2, src2);
reload1!(b3, ip3, src3);
}
cursors[0] = c0;
cursors[1] = c1;
cursors[2] = c2;
cursors[3] = c3;
if !any_iter {
return;
}
macro_rules! writeback {
($i:literal, $b:ident, $ip:ident, $src:expr) => {{
brs[$i].index = $ip;
brs[$i].bit_container = u64::from_le_bytes(unsafe {
$src.get_unchecked($ip..$ip + 8)
.try_into()
.unwrap_unchecked()
});
brs[$i].bits_consumed = $b.trailing_zeros() as u8 + max_num_bits;
decoders[$i].state = $b >> table_shift;
}};
}
writeback!(0, b0, ip0, src0);
writeback!(1, b1, ip1, src1);
writeback!(2, b2, ip2, src2);
writeback!(3, b3, ip3, src3);
}
#[cfg(test)]
mod zerocopy_robustness_tests;
#[cfg(test)]
mod burst_gate_tests;