diff --git a/lib/decompress/zstd_decompress_block.c b/lib/decompress/zstd_decompress_block.c index d8b0c9ca4..81dfdbfaa 100644 --- a/lib/decompress/zstd_decompress_block.c +++ b/lib/decompress/zstd_decompress_block.c @@ -20,6 +20,19 @@ #include "zstd_decompress_internal.h" #include "zstd_decompress_block.h" +#if defined(ZSTD_FORCE_DECOMPRESS_SEQUENCES_SHORT) && \ + defined(ZSTD_FORCE_DECOMPRESS_SEQUENCES_LONG) +#error "Cannot force the use of the short and the long ZSTD_decompressSequences variants!" +#endif + +#if defined(ZSTD_FORCE_DECOMPRESS_SEQUENCES_SHORT) +#define ZSTD_RUST_SEQUENCE_DECODER_MODE 1 +#elif defined(ZSTD_FORCE_DECOMPRESS_SEQUENCES_LONG) +#define ZSTD_RUST_SEQUENCE_DECODER_MODE 2 +#else +#define ZSTD_RUST_SEQUENCE_DECODER_MODE 0 +#endif + typedef char ZSTD_rust_block_seq_symbol_layout[(sizeof(ZSTD_seqSymbol) == 8) ? 1 : -1]; typedef char ZSTD_rust_block_entropy_rep_offset[ (offsetof(ZSTD_entropyDTables_t, rep) == 26652) ? 1 : -1]; @@ -53,6 +66,7 @@ typedef struct { ZSTD_litLocation_e* litBufferLocation; BYTE* litExtraBuffer; size_t litExtraBufferSize; + int sequenceDecoderMode; } ZSTD_rustBlockCtx; static ZSTD_rustBlockCtx ZSTD_rust_block_context(ZSTD_DCtx* dctx) @@ -84,6 +98,7 @@ static ZSTD_rustBlockCtx ZSTD_rust_block_context(ZSTD_DCtx* dctx) ctx.litBufferLocation = &dctx->litBufferLocation; ctx.litExtraBuffer = dctx->litExtraBuffer; ctx.litExtraBufferSize = ZSTD_LITBUFFEREXTRASIZE; + ctx.sequenceDecoderMode = ZSTD_RUST_SEQUENCE_DECODER_MODE; return ctx; } diff --git a/rust/src/zstd_decompress_block.rs b/rust/src/zstd_decompress_block.rs index 8ea71ff5c..0f2e07603 100644 --- a/rust/src/zstd_decompress_block.rs +++ b/rust/src/zstd_decompress_block.rs @@ -57,6 +57,15 @@ const ML_FSE_LOG: u32 = 9; const ZSTD_HUFFDTABLE_CAPACITY_LOG: usize = 12; const HUF_DTABLE_SIZE: usize = 1 + (1 << ZSTD_HUFFDTABLE_CAPACITY_LOG); const ZSTD_BUILD_FSE_TABLE_WKSP_SIZE_U32: usize = 157; +const STREAM_ACCUMULATOR_MIN_32: usize = 25; +const STREAM_ACCUMULATOR_MIN_64: usize = 57; +const SEQUENCE_DECODER_RUNTIME: c_int = 0; +const SEQUENCE_DECODER_FORCE_SHORT: c_int = 1; +const SEQUENCE_DECODER_FORCE_LONG: c_int = 2; +const LONG_OFFSET_ADDITIONAL_BITS: u8 = 22; +const LONG_OFFSET_HISTORY_THRESHOLD: usize = 1 << 24; +const LONG_OFFSET_MIN_SHARE_32: u32 = 20; +const LONG_OFFSET_MIN_SHARE_64: u32 = 7; const LL_BASE: [u32; MAX_LL + 1] = [ 0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 18, 20, 22, 24, 28, 32, 40, 48, 64, @@ -138,6 +147,7 @@ pub struct ZSTD_rustBlockCtx { lit_buffer_location: *mut c_int, lit_extra_buffer: *mut u8, lit_extra_buffer_size: usize, + sequence_decoder_mode: c_int, } #[repr(C)] @@ -155,6 +165,18 @@ struct block_properties_t { orig_size: u32, } +#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)] +struct ZSTD_offsetInfo { + long_offset_share: u32, + max_nb_additional_bits: u8, +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +enum SequenceDecoder { + Short, + Long, +} + #[derive(Clone, Copy)] struct ZSTD_fseState { state: usize, @@ -213,6 +235,80 @@ unsafe fn set_seq_header(table: *mut ZSTD_seqSymbol, header: ZSTD_seqSymbol_head }; } +#[inline] +fn stream_accumulator_min() -> usize { + if MEM_32bits() { + STREAM_ACCUMULATOR_MIN_32 + } else { + STREAM_ACCUMULATOR_MIN_64 + } +} + +/// Port of the C decoder's `ZSTD_getOffsetInfo()` policy helper. +/// +/// `long_offset_share` is scaled to the full offset-table log, so the caller +/// can compare tables with different table logs using the same thresholds. +/// The table is valid when this helper is reached from sequence-header +/// decoding; the defensive checks keep the helper harmless in focused tests +/// and preserve the C zero-sequence behavior. +#[inline] +unsafe fn get_offset_info(off_table: *const ZSTD_seqSymbol, nb_seq: c_int) -> ZSTD_offsetInfo { + if nb_seq == 0 || off_table.is_null() { + return ZSTD_offsetInfo::default(); + } + let table_log = unsafe { seq_header(off_table).table_log }; + if table_log > OFF_FSE_LOG { + return ZSTD_offsetInfo::default(); + } + let table_size = 1usize << table_log; + let table = unsafe { off_table.add(1) }; + let mut info = ZSTD_offsetInfo::default(); + for index in 0..table_size { + let additional_bits = unsafe { (*table.add(index)).nb_additional_bits }; + info.max_nb_additional_bits = info.max_nb_additional_bits.max(additional_bits); + if additional_bits > LONG_OFFSET_ADDITIONAL_BITS { + info.long_offset_share = info.long_offset_share.wrapping_add(1); + } + } + info.long_offset_share <<= OFF_FSE_LOG - table_log; + info +} + +#[inline] +fn choose_sequence_decoder( + mode: c_int, + ddict_is_cold: c_int, + potential_long_offsets: bool, + total_history_size: usize, + nb_seq: c_int, + offset_info: ZSTD_offsetInfo, +) -> SequenceDecoder { + match mode { + SEQUENCE_DECODER_FORCE_SHORT => return SequenceDecoder::Short, + SEQUENCE_DECODER_FORCE_LONG => return SequenceDecoder::Long, + SEQUENCE_DECODER_RUNTIME => {} + _ => {} + } + + let mut use_prefetch_decoder = ddict_is_cold != 0; + if !use_prefetch_decoder + && (potential_long_offsets + || (total_history_size > LONG_OFFSET_HISTORY_THRESHOLD && nb_seq > 8)) + { + let minimum_share = if MEM_32bits() { + LONG_OFFSET_MIN_SHARE_32 + } else { + LONG_OFFSET_MIN_SHARE_64 + }; + use_prefetch_decoder = offset_info.long_offset_share >= minimum_share; + } + if use_prefetch_decoder { + SequenceDecoder::Long + } else { + SequenceDecoder::Short + } +} + #[inline] unsafe fn copy_bytes(dst: *mut u8, src: *const u8, len: usize) { if len != 0 { @@ -276,6 +372,17 @@ unsafe fn block_size_max(ctx: *const ZSTD_rustBlockCtx) -> usize { value } +#[inline] +unsafe fn total_history_size( + ctx: *const ZSTD_rustBlockCtx, + dst: *mut u8, + dst_capacity: usize, +) -> usize { + let block_size = min(dst_capacity, unsafe { block_size_max(ctx) }); + let history_end = unsafe { ptr_add(dst, block_size) }; + unsafe { address_distance(history_end.cast_const(), *(*ctx).virtual_start) } +} + unsafe fn allocate_literals_buffer( ctx: *mut ZSTD_rustBlockCtx, dst: *mut u8, @@ -1399,6 +1506,7 @@ unsafe fn decompress_sequences( sequence_size: usize, nb_seq: c_int, long_offsets: bool, + use_prefetch_decoder: bool, ) -> usize { if nb_seq < 0 { return ERROR(ZstdErrorCode::CorruptionDetected); @@ -1460,8 +1568,13 @@ unsafe fn decompress_sequences( { return error; } + let mut prefetch_pos = 0usize; for remaining in (1..=nb_seq as usize).rev() { let mut sequence = unsafe { decode_sequence(&mut state, long_offsets, remaining == 1) }; + if use_prefetch_decoder { + prefetch_pos = + unsafe { ZSTD_prefetchMatch(prefetch_pos, sequence, prefix_start, dict_end) }; + } if location == ZSTD_SPLIT && sequence.lit_length > unsafe { address_distance(lit_limit, lit_ptr) } { @@ -1594,9 +1707,8 @@ pub unsafe extern "C" fn ZSTD_rust_decompressBlock_internal( } input = unsafe { input.add(lit_size) }; let remaining = src_size - lit_size; - let history_end = unsafe { ptr_add(dst.cast(), min(dst_capacity, block_max)) }; - let history_size = unsafe { address_distance(history_end.cast_const(), *(*ctx).virtual_start) }; - let long_offsets = MEM_32bits() && history_size > max_short_offset(); + let history_size = unsafe { total_history_size(ctx, dst.cast(), dst_capacity) }; + let potential_long_offsets = MEM_32bits() && history_size > max_short_offset(); let mut nb_seq = 0 as c_int; let header_size = unsafe { ZSTD_rust_decodeSeqHeaders(ctx, &mut nb_seq, input.cast(), remaining) }; @@ -1608,6 +1720,27 @@ pub unsafe extern "C" fn ZSTD_rust_decompressBlock_internal( if (dst.is_null() || dst_capacity == 0) && nb_seq > 0 { return ERROR(ZstdErrorCode::DstSizeTooSmall); } + let offset_info = if potential_long_offsets + || (unsafe { *(*ctx).ddict_is_cold } == 0 + && history_size > LONG_OFFSET_HISTORY_THRESHOLD + && nb_seq > 8) + { + unsafe { get_offset_info(*(*ctx).oft_ptr, nb_seq) } + } else { + ZSTD_offsetInfo::default() + }; + let mut long_offsets = potential_long_offsets; + if long_offsets && offset_info.max_nb_additional_bits <= stream_accumulator_min() as u8 { + long_offsets = false; + } + let decoder = choose_sequence_decoder( + unsafe { (*ctx).sequence_decoder_mode }, + unsafe { *(*ctx).ddict_is_cold }, + potential_long_offsets, + history_size, + nb_seq, + offset_info, + ); unsafe { *(*ctx).ddict_is_cold = 0 }; unsafe { decompress_sequences( @@ -1618,6 +1751,7 @@ pub unsafe extern "C" fn ZSTD_rust_decompressBlock_internal( sequence_size, nb_seq, long_offsets, + decoder == SequenceDecoder::Long, ) } } @@ -1697,4 +1831,163 @@ mod tests { .iter() .all(|symbol| symbol.base_value == 0x3F_FFFD)); } + + #[test] + fn offset_info_scales_long_share_and_handles_invalid_tables() { + let mut table = [ZSTD_seqSymbol::default(); 1 + (1 << 5)]; + unsafe { + set_seq_header( + table.as_mut_ptr(), + ZSTD_seqSymbol_header { + fast_mode: 0, + table_log: 5, + }, + ); + } + table[1].nb_additional_bits = 23; + table[2].nb_additional_bits = 31; + + let info = unsafe { get_offset_info(table.as_ptr(), 1) }; + assert_eq!(info.max_nb_additional_bits, 31); + assert_eq!(info.long_offset_share, 2 << (OFF_FSE_LOG - 5)); + assert_eq!( + unsafe { get_offset_info(table.as_ptr(), 0) }, + ZSTD_offsetInfo::default() + ); + + unsafe { + set_seq_header( + table.as_mut_ptr(), + ZSTD_seqSymbol_header { + fast_mode: 0, + table_log: OFF_FSE_LOG + 1, + }, + ); + } + assert_eq!( + unsafe { get_offset_info(table.as_ptr(), 1) }, + ZSTD_offsetInfo::default() + ); + } + + #[test] + fn sequence_decoder_policy_preserves_forced_and_runtime_variants() { + let minimum_share = if MEM_32bits() { + LONG_OFFSET_MIN_SHARE_32 + } else { + LONG_OFFSET_MIN_SHARE_64 + }; + let info = ZSTD_offsetInfo { + long_offset_share: minimum_share, + max_nb_additional_bits: 31, + }; + + assert_eq!( + choose_sequence_decoder(SEQUENCE_DECODER_FORCE_SHORT, 1, true, usize::MAX, 9, info,), + SequenceDecoder::Short + ); + assert_eq!( + choose_sequence_decoder( + SEQUENCE_DECODER_FORCE_LONG, + 0, + false, + 0, + 0, + ZSTD_offsetInfo::default(), + ), + SequenceDecoder::Long + ); + assert_eq!( + choose_sequence_decoder( + SEQUENCE_DECODER_RUNTIME, + 1, + false, + 0, + 0, + ZSTD_offsetInfo::default(), + ), + SequenceDecoder::Long + ); + assert_eq!( + choose_sequence_decoder( + SEQUENCE_DECODER_RUNTIME, + 0, + false, + LONG_OFFSET_HISTORY_THRESHOLD + 1, + 9, + info, + ), + SequenceDecoder::Long + ); + assert_eq!( + choose_sequence_decoder( + SEQUENCE_DECODER_RUNTIME, + 0, + false, + LONG_OFFSET_HISTORY_THRESHOLD, + 8, + info, + ), + SequenceDecoder::Short + ); + } + + #[test] + fn block_header_edges_keep_rle_reserved_and_truncated_contracts() { + let mut properties = block_properties_t { + block_type: -1, + last_block: 99, + orig_size: 99, + }; + let truncated = [0u8; 2]; + let truncated_result = unsafe { + ZSTD_getcBlockSize( + truncated.as_ptr().cast(), + truncated.len(), + (&mut properties as *mut block_properties_t).cast(), + ) + }; + assert_eq!( + crate::errors::ERR_getErrorCode(truncated_result), + ZstdErrorCode::SrcSizeWrong as i32 + ); + assert_eq!(properties.block_type, -1); + + let rle_header = (0x1234u32 << 3) | (1 << 1) | 1; + let rle = [ + rle_header as u8, + (rle_header >> 8) as u8, + (rle_header >> 16) as u8, + ]; + let rle_result = unsafe { + ZSTD_getcBlockSize( + rle.as_ptr().cast(), + rle.len(), + (&mut properties as *mut block_properties_t).cast(), + ) + }; + assert_eq!(rle_result, 1); + assert_eq!(properties.block_type, SET_RLE); + assert_eq!(properties.last_block, 1); + assert_eq!(properties.orig_size, 0x1234); + + let reserved_header = 3u32 << 1; + let reserved = [ + reserved_header as u8, + (reserved_header >> 8) as u8, + (reserved_header >> 16) as u8, + ]; + let reserved_result = unsafe { + ZSTD_getcBlockSize( + reserved.as_ptr().cast(), + reserved.len(), + (&mut properties as *mut block_properties_t).cast(), + ) + }; + assert_eq!( + crate::errors::ERR_getErrorCode(reserved_result), + ZstdErrorCode::CorruptionDetected as i32 + ); + assert_eq!(properties.block_type, 3); + } }