diff --git a/rust/src/zstd_compress.rs b/rust/src/zstd_compress.rs index f86a6356a..cc554cccc 100644 --- a/rust/src/zstd_compress.rs +++ b/rust/src/zstd_compress.rs @@ -17,7 +17,8 @@ use crate::errors::{ERR_isError, ZstdErrorCode, ERROR}; #[cfg(not(test))] use crate::zstd_compress_api::ZSTD_compressBound; use crate::zstd_compress_frame::{ - write_raw_block, ZSTD_rust_writeFrameHeader, ZSTD_writeLastEmptyBlock, + write_raw_block, ZSTD_rust_writeBlockHeader, ZSTD_rust_writeFrameHeader, + ZSTD_writeLastEmptyBlock, }; use crate::zstd_compress_literals::min_gain; use crate::zstd_compress_params::{ @@ -25,8 +26,9 @@ use crate::zstd_compress_params::{ ZSTD_RUST_CPM_NO_ATTACH_DICT, ZSTD_RUST_PS_DISABLE, }; use crate::zstd_compress_sequences::SeqDef; -use crate::zstd_compress_stats::{SeqStore_t, ZSTD_compressedBlockState_t}; -use crate::zstd_compress_superblock::ZSTD_rust_compressSuperBlock; +use crate::zstd_compress_stats::{ + SeqStore_t, ZSTD_compressedBlockState_t, ZSTD_rust_entropyCompressSeqStore, ZSTD_rust_isRLE, +}; use std::ffi::c_void; use std::mem::{size_of, MaybeUninit}; use std::os::raw::{c_int, c_uint}; @@ -34,6 +36,8 @@ use std::ptr; #[cfg(not(test))] unsafe extern "C" { + fn ZSTD_createCCtx() -> *mut c_void; + fn ZSTD_freeCCtx(cctx: *mut c_void) -> usize; fn ZSTD_rust_resetCCtxForSimpleCompression(cctx: *mut c_void) -> usize; fn ZSTD_rust_prepareCCtxForSimpleCompression( cctx: *mut c_void, @@ -82,6 +86,7 @@ const ZSTD_BLOCKSIZE_MAX: usize = 1 << 17; const ZSTD_CONTENTSIZE_UNKNOWN: u64 = u64::MAX; const ZSTD_TARGET_CBLOCK_BSS_COMPRESS: c_int = 0; const ZSTD_BLOCK_HEADER_SIZE: usize = 3; +const MIN_COMPRESSIBLE_BLOCK_SIZE: usize = 2 + ZSTD_BLOCK_HEADER_SIZE + 1 + 1; const ZSTD_ROWSIZE: usize = 16; const ZSTD_WINDOW_START_INDEX: u32 = 2; const ZSTD_DUBT_UNSORTED_MARK: u32 = 1; @@ -876,10 +881,9 @@ fn ceil_log2(size: usize) -> u32 { /// Compress one frame using the already migrated block leaves. /// -/// This path deliberately starts a fresh match table for each 128 KiB block. -/// That keeps the Rust-owned context independent from the still-private C -/// `ZSTD_MatchState_t` window while producing ordinary zstd blocks that any -/// decoder can consume. +/// This path owns the match tables and carries their history across the +/// 128 KiB block boundaries, while keeping the private C +/// `ZSTD_MatchState_t` window out of the Rust ABI. unsafe fn compress_frame( dst: *mut c_void, dst_capacity: usize, @@ -947,10 +951,7 @@ unsafe fn compress_frame( } let first_block_size = src_size.min(ZSTD_BLOCKSIZE_MAX); - /* A table larger than the block cannot contain a useful index for this - * independent-block path. This also keeps one-shot small inputs from - * allocating the table selected for a huge source-size hint. */ - let matcher_hash_log = cparams.hashLog.min(ceil_log2(first_block_size).max(6)); + let matcher_hash_log = cparams.hashLog; let matcher_chain_log = cparams.chainLog.min(ceil_log2(first_block_size).max(6)); let hash_size = match checked_table_size(matcher_hash_log) { Some(size) => size, @@ -994,12 +995,32 @@ unsafe fn compress_frame( let source = src.cast::(); let output = dst.cast::(); + let matcher_base = source.wrapping_sub(ZSTD_WINDOW_START_INDEX as usize); let mut input_offset = 0usize; let mut output_offset = header_size; while input_offset < src_size { let block_size = (src_size - input_offset).min(ZSTD_BLOCKSIZE_MAX); let block_src = unsafe { source.add(input_offset) }; let block_end = unsafe { block_src.add(block_size) }; + let last_block = u32::from(input_offset + block_size == src_size); + + if block_size < MIN_COMPRESSIBLE_BLOCK_SIZE { + let written = unsafe { + write_raw_block( + output.add(output_offset), + dst_capacity.saturating_sub(output_offset), + block_src, + block_size, + last_block, + ) + }; + if ERR_isError(written) { + return written; + } + output_offset += written; + input_offset += block_size; + continue; + } seq_store.sequences = seq_store.sequencesStart; seq_store.lit = seq_store.litStart; @@ -1015,8 +1036,8 @@ unsafe fn compress_frame( crate::zstd_double_fast::ZSTD_rust_compressBlock_doubleFast( hash_table.as_mut_ptr(), chain_table.as_mut_ptr(), - block_src, - 0, + matcher_base, + ZSTD_WINDOW_START_INDEX, 0, matcher_hash_log, matcher_chain_log, @@ -1032,8 +1053,8 @@ unsafe fn compress_frame( unsafe { crate::zstd_fast::ZSTD_rust_compressBlock_fast( hash_table.as_mut_ptr(), - block_src, - 0, + matcher_base, + ZSTD_WINDOW_START_INDEX, 0, matcher_hash_log, cparams.minMatch, @@ -1057,34 +1078,40 @@ unsafe fn compress_frame( seq_store.lit = seq_store.lit.add(last_literals); } } - next_block.rep = reps; - - let last_block = u32::from(input_offset + block_size == src_size); let remaining_capacity = dst_capacity.saturating_sub(output_offset); - let mut written = unsafe { - ZSTD_rust_compressSuperBlock( - (seq_store as *mut SeqStore_t).cast(), - (&prev_block as *const ZSTD_compressedBlockState_t).cast(), - (&mut next_block as *mut ZSTD_compressedBlockState_t).cast(), + if remaining_capacity < ZSTD_BLOCK_HEADER_SIZE { + return ERROR(ZstdErrorCode::DstSizeTooSmall); + } + let mut compressed_size = unsafe { + ZSTD_rust_entropyCompressSeqStore( + (seq_store as *const SeqStore_t).cast(), + ptr::addr_of!(prev_block.entropy), + ptr::addr_of_mut!(next_block.entropy), cparams.strategy, disable_literal_compression, + output.add(output_offset + ZSTD_BLOCK_HEADER_SIZE).cast(), + remaining_capacity - ZSTD_BLOCK_HEADER_SIZE, + block_size, workspace.as_mut_ptr().cast(), TMP_WORKSPACE_SIZE, 0, /* BMI2 is optional; portable Rust leaf path */ - cparams.windowLog, - 0, /* no target compressed block size */ - output.add(output_offset).cast(), - remaining_capacity, - block_src.cast(), - block_size, - last_block, ) }; - if ERR_isError(written) { - return written; + if ERR_isError(compressed_size) { + return compressed_size; } - if written == 0 { - written = unsafe { + if input_offset != 0 + && compressed_size < 25 + && unsafe { ZSTD_rust_isRLE(block_src, block_size) } != 0 + { + unsafe { + *output.add(output_offset + ZSTD_BLOCK_HEADER_SIZE) = *block_src; + } + compressed_size = 1; + } + + let written = if compressed_size == 0 { + unsafe { write_raw_block( output.add(output_offset), remaining_capacity, @@ -1092,19 +1119,29 @@ unsafe fn compress_frame( block_size, last_block, ) - }; - if ERR_isError(written) { - return written; } - } - if written > remaining_capacity { - return ERROR(ZstdErrorCode::DstSizeTooSmall); + } else { + unsafe { + ZSTD_rust_writeBlockHeader( + output.add(output_offset).cast(), + compressed_size, + block_size, + last_block, + ) + }; + if compressed_size > 1 { + next_block.rep = reps; + } + compressed_size + ZSTD_BLOCK_HEADER_SIZE + }; + if ERR_isError(written) { + return written; } output_offset += written; input_offset += block_size; - std::mem::swap(&mut prev_block, &mut next_block); - hash_table.fill(0); - chain_table.fill(0); + if compressed_size > 1 { + std::mem::swap(&mut prev_block, &mut next_block); + } } output_offset @@ -1119,6 +1156,30 @@ pub unsafe extern "C" fn ZSTD_compress( src_size: usize, compression_level: c_int, ) -> usize { + #[cfg(not(test))] + { + let strategy = unsafe { ZSTD_rust_compressCCtxStrategy(src_size, compression_level) }; + if strategy != ZSTD_FAST && strategy != ZSTD_DFAST { + let cctx = unsafe { ZSTD_createCCtx() }; + if cctx.is_null() { + return ERROR(ZstdErrorCode::MemoryAllocation); + } + let result = unsafe { + ZSTD_compress_usingDict( + cctx, + dst, + dst_capacity, + src, + src_size, + ptr::null(), + 0, + compression_level, + ) + }; + unsafe { ZSTD_freeCCtx(cctx) }; + return result; + } + } unsafe { compress_frame(dst, dst_capacity, src, src_size, compression_level) } } @@ -2284,6 +2345,18 @@ mod tests { } } + #[test] + fn one_shot_promotes_nonfirst_rle_blocks() { + let mut input = vec![b'B'; 256 * 1024 - 2]; + input.extend(std::iter::repeat(b'A').take(100 * 1024)); + + let compressed = compress_input(&input, 1); + assert!(compressed.len() <= 46); + if let Some(restored) = system_round_trip(&compressed) { + assert_eq!(restored, input); + } + } + #[test] fn public_error_paths_match_size_t_error_contract() { let mut output = [0u8; 64];