diff --git a/lib/compress/zstd_ldm.c b/lib/compress/zstd_ldm.c index c97512fc5..037bb6e61 100644 --- a/lib/compress/zstd_ldm.c +++ b/lib/compress/zstd_ldm.c @@ -262,7 +262,8 @@ size_t ZSTD_rust_ldm_generateSequences( const void* src, size_t srcSize, int overflowCorrectFrequently); size_t ZSTD_rust_ldm_blockCompress( void* rawSeqStore, void* blockContext, void* seqStore, U32 rep[ZSTD_REP_NUM], - const void* src, size_t srcSize, U32 minMatch, int strategy); + const void* src, size_t srcSize, U32 minMatch, int strategy, + int useRowMatchFinder); U32 ZSTD_rust_ldm_limitTableUpdate(U32 curr, U32 nextToUpdate); enum { @@ -280,15 +281,43 @@ size_t ZSTD_ldm_rust_compressLiterals( const void* src, size_t srcSize); void ZSTD_ldm_rust_setLdmSeqStore(void* blockContext, const void* rawSeqStore); +typedef size_t (*ZSTD_rust_ldm_select_block_compressor_f)( + void* blockContext, int useRowMatchFinder); + const U64* ZSTD_ldm_rust_gearTable(void) { return ZSTD_ldm_gearTab; } typedef struct { - ZSTD_MatchState_t* ms; + void* callbackContext; + ZSTD_rust_ldm_select_block_compressor_f selectBlockCompressor; ZSTD_BlockCompressor_f blockCompressor; } ZSTD_rust_ldm_block_context; +typedef char ZSTD_rust_ldm_block_context_layout[ + (offsetof(ZSTD_rust_ldm_block_context, callbackContext) == 0 + && offsetof(ZSTD_rust_ldm_block_context, selectBlockCompressor) + == sizeof(void*) + && offsetof(ZSTD_rust_ldm_block_context, blockCompressor) + == 2 * sizeof(void*) + && sizeof(ZSTD_rust_ldm_block_context) == 3 * sizeof(void*)) + ? 1 : -1]; + +/* Rust owns when block-compressor selection happens and how its status is + * handled. This callback retains the dictionary-mode inspection and the + * private MatchState/compressor selection in C. */ +static size_t ZSTD_rust_ldm_selectBlockCompressor( + void* blockContext, int useRowMatchFinder) +{ + ZSTD_rust_ldm_block_context* const context = + (ZSTD_rust_ldm_block_context*)blockContext; + ZSTD_MatchState_t* const ms = (ZSTD_MatchState_t*)context->callbackContext; + context->blockCompressor = ZSTD_selectBlockCompressor( + ms->cParams.strategy, (ZSTD_ParamSwitch_e)useRowMatchFinder, + ZSTD_matchState_dictMode(ms)); + assert(context->blockCompressor != NULL); + return context->blockCompressor == NULL ? ERROR(GENERIC) : 0; +} static void ZSTD_rust_ldm_fillFastTables(ZSTD_MatchState_t* ms, const BYTE* end) { @@ -313,10 +342,10 @@ static void ZSTD_rust_ldm_fillFastTables(ZSTD_MatchState_t* ms, const BYTE* end) void ZSTD_ldm_rust_prepareBlock(void* blockContext, const void* anchor) { ZSTD_rust_ldm_block_context* const context = (ZSTD_rust_ldm_block_context*)blockContext; - U32 const curr = (U32)((const BYTE*)anchor - context->ms->window.base); - context->ms->nextToUpdate = ZSTD_rust_ldm_limitTableUpdate( - curr, context->ms->nextToUpdate); - ZSTD_rust_ldm_fillFastTables(context->ms, (const BYTE*)anchor); + ZSTD_MatchState_t* const ms = (ZSTD_MatchState_t*)context->callbackContext; + U32 const curr = (U32)((const BYTE*)anchor - ms->window.base); + ms->nextToUpdate = ZSTD_rust_ldm_limitTableUpdate(curr, ms->nextToUpdate); + ZSTD_rust_ldm_fillFastTables(ms, (const BYTE*)anchor); } size_t ZSTD_ldm_rust_compressLiterals( @@ -324,13 +353,16 @@ size_t ZSTD_ldm_rust_compressLiterals( const void* src, size_t srcSize) { ZSTD_rust_ldm_block_context* const context = (ZSTD_rust_ldm_block_context*)blockContext; - return context->blockCompressor(context->ms, (SeqStore_t*)seqStore, rep, src, srcSize); + ZSTD_MatchState_t* const ms = (ZSTD_MatchState_t*)context->callbackContext; + assert(context->blockCompressor != NULL); + return context->blockCompressor(ms, (SeqStore_t*)seqStore, rep, src, srcSize); } void ZSTD_ldm_rust_setLdmSeqStore(void* blockContext, const void* rawSeqStore) { ZSTD_rust_ldm_block_context* const context = (ZSTD_rust_ldm_block_context*)blockContext; - context->ms->ldmSeqStore = (const RawSeqStore_t*)rawSeqStore; + ZSTD_MatchState_t* const ms = (ZSTD_MatchState_t*)context->callbackContext; + ms->ldmSeqStore = (const RawSeqStore_t*)rawSeqStore; } size_t ZSTD_ldm_getTableSize(ldmParams_t params) @@ -372,11 +404,11 @@ size_t ZSTD_ldm_blockCompress(RawSeqStore_t* rawSeqStore, void const* src, size_t srcSize) { ZSTD_rust_ldm_block_context context; - context.ms = ms; - context.blockCompressor = ZSTD_selectBlockCompressor( - ms->cParams.strategy, useRowMatchFinder, ZSTD_matchState_dictMode(ms)); - assert(context.blockCompressor != NULL); + context.callbackContext = ms; + context.selectBlockCompressor = ZSTD_rust_ldm_selectBlockCompressor; + context.blockCompressor = NULL; return ZSTD_rust_ldm_blockCompress( rawSeqStore, &context, seqStore, rep, src, srcSize, - ms->cParams.minMatch, (int)ms->cParams.strategy); + ms->cParams.minMatch, (int)ms->cParams.strategy, + (int)useRowMatchFinder); } diff --git a/rust/src/zstd_ldm.rs b/rust/src/zstd_ldm.rs index 0ec11914a..274fbb2ba 100644 --- a/rust/src/zstd_ldm.rs +++ b/rust/src/zstd_ldm.rs @@ -15,7 +15,7 @@ use crate::xxhash::XXH64; use crate::zstd_compress_params::ZSTD_compressionParameters; use crate::zstd_fast::store_seq_opaque; use std::ffi::c_void; -use std::mem::size_of; +use std::mem::{offset_of, size_of}; use std::os::raw::c_int; const LDM_BATCH_SIZE: usize = 64; @@ -35,6 +35,23 @@ const LDM_FAST_TABLES_FAST: c_int = 1; const LDM_FAST_TABLES_DFAST: c_int = 2; const LDM_FAST_TABLES_INVALID: c_int = -1; +type LdmSelectBlockCompressorFn = unsafe extern "C" fn(*mut c_void, c_int) -> usize; + +/// Prefix shared with the C-owned block context. The trailing compressor +/// function pointer remains private to C; Rust only invokes the selection +/// callback and consumes its status before entering block processing. +#[repr(C)] +struct LdmBlockContextProjection { + callback_context: *mut c_void, + select_block_compressor: Option, +} + +const _: () = { + assert!(offset_of!(LdmBlockContextProjection, callback_context) == 0); + assert!(offset_of!(LdmBlockContextProjection, select_block_compressor) == size_of::()); + assert!(size_of::() == 2 * size_of::()); +}; + #[repr(C)] #[derive(Clone, Copy)] struct LdmEntry { @@ -105,6 +122,24 @@ fn use_optimal_parser(strategy: c_int) -> bool { strategy >= ZSTD_BTOPT } +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +enum LdmBlockCompressionPath { + InterspersedSequences, + OptimalParser, +} + +#[inline] +fn ldm_block_compression_path(strategy: c_int) -> Option { + if !(ZSTD_FAST_STRATEGY..=ZSTD_STRATEGY_MAX).contains(&strategy) { + return None; + } + Some(if use_optimal_parser(strategy) { + LdmBlockCompressionPath::OptimalParser + } else { + LdmBlockCompressionPath::InterspersedSequences + }) +} + unsafe extern "C" { fn ZSTD_ldm_rust_gearTable() -> *const u64; fn ZSTD_ldm_rust_prepareBlock(context: *mut c_void, anchor: *const c_void); @@ -974,11 +1009,31 @@ pub unsafe extern "C" fn ZSTD_rust_ldm_blockCompress( src_size: usize, min_match: u32, strategy: c_int, + use_row_match_finder: c_int, ) -> usize { + let Some(path) = ldm_block_compression_path(strategy) else { + return ERROR(ZstdErrorCode::Generic); + }; + if block_context.is_null() { + return ERROR(ZstdErrorCode::Generic); + } + let block_context_projection = + unsafe { &mut *block_context.cast::() }; + if block_context_projection.callback_context.is_null() { + return ERROR(ZstdErrorCode::Generic); + } + let Some(select_block_compressor) = block_context_projection.select_block_compressor else { + return ERROR(ZstdErrorCode::Generic); + }; + let selection_result = unsafe { select_block_compressor(block_context, use_row_match_finder) }; + if ERR_isError(selection_result) { + return selection_result; + } + let raw_seq_store = raw_seq_store.cast::(); let input = src.cast::(); let input_end = input.wrapping_add(src_size); - if use_optimal_parser(strategy) { + if path == LdmBlockCompressionPath::OptimalParser { unsafe { ZSTD_ldm_rust_setLdmSeqStore(block_context, raw_seq_store.cast::()) }; let last_literals = unsafe { ZSTD_ldm_rust_compressLiterals(block_context, seq_store, reps, src, src_size) @@ -1121,6 +1176,28 @@ mod tests { assert!(use_optimal_parser(9)); } + #[test] + fn block_compression_path_validates_strategy_and_preserves_parser_boundary() { + assert_eq!(ldm_block_compression_path(ZSTD_FAST_STRATEGY - 1), None); + assert_eq!( + ldm_block_compression_path(ZSTD_FAST_STRATEGY), + Some(LdmBlockCompressionPath::InterspersedSequences) + ); + assert_eq!( + ldm_block_compression_path(ZSTD_BTOPT - 1), + Some(LdmBlockCompressionPath::InterspersedSequences) + ); + assert_eq!( + ldm_block_compression_path(ZSTD_BTOPT), + Some(LdmBlockCompressionPath::OptimalParser) + ); + assert_eq!( + ldm_block_compression_path(ZSTD_STRATEGY_MAX), + Some(LdmBlockCompressionPath::OptimalParser) + ); + assert_eq!(ldm_block_compression_path(ZSTD_STRATEGY_MAX + 1), None); + } + #[test] fn fast_table_kind_preserves_c_strategy_dispatch() { assert_eq!(