diff --git a/lib/compress/zstd_compress.c b/lib/compress/zstd_compress.c index 62c9a753b..543d87793 100644 --- a/lib/compress/zstd_compress.c +++ b/lib/compress/zstd_compress.c @@ -1958,12 +1958,101 @@ U32 ZSTD_rust_resolveRepcodeToRawOffset(const U32 rep[ZSTD_REP_NUM], U32 offBase, U32 ll0); size_t ZSTD_rust_loadCEntropy(ZSTD_compressedBlockState_t* bs, void* workspace, const void* dict, size_t dictSize); -/* Rust owns dictionary-ingestion policy. The content loader remains a narrow - * callback because it needs configuration-sensitive C-private layouts. */ +/* Rust owns dictionary-content orchestration and policy. The callbacks below + * keep configuration-sensitive match-state, LDM, workspace, and matchfinder + * operations private to C. */ typedef size_t (*ZSTD_rust_loadDictionaryContent_f)( void* matchState, void* ldmState, void* workspaceState, const void* params, const void* src, size_t srcSize, int dtlm, int tfp); +typedef void (*ZSTD_rust_loadDictionaryContent_assert_f)(void* context); +typedef void (*ZSTD_rust_loadDictionaryContent_assertWindowEmpty_f)( + void* context, int ldm); +typedef void (*ZSTD_rust_loadDictionaryContent_windowUpdate_f)( + void* context, int ldm, const void* src, size_t srcSize); +typedef void (*ZSTD_rust_loadDictionaryContent_setLdmLoadedDictEnd_f)( + void* context, const void* iend, int forceWindow); +typedef void (*ZSTD_rust_loadDictionaryContent_publishMatchState_f)( + void* context, const void* ip, const void* iend, + int forceWindow, int deterministicRefPrefix); +typedef void (*ZSTD_rust_loadDictionaryContent_fillLdm_f)( + void* context, const void* ip, const void* iend); +typedef void (*ZSTD_rust_loadDictionaryContent_overflowCorrect_f)( + void* context, const void* ip, const void* iend); +typedef void (*ZSTD_rust_loadDictionaryContent_fillTable_f)( + void* context, const void* iend, int dtlm, int tfp); +typedef void (*ZSTD_rust_loadDictionaryContent_loadMatch_f)( + void* context, const void* ip); +typedef void (*ZSTD_rust_loadDictionaryContent_loadTree_f)( + void* context, const void* ip, const void* iend); +typedef void (*ZSTD_rust_loadDictionaryContent_publishFinalIndex_f)( + void* context, const void* iend); + +typedef struct { + void* callbackContext; + size_t currentMax; + size_t windowStartIndex; + size_t shortCacheTagBits; + size_t chunkSizeMax; + size_t hashReadSize; + size_t hashLog; + size_t chainLog; + size_t cdictIndicesTagged; + size_t ldmEnabled; + size_t hasLdmState; + size_t forceWindow; + size_t deterministicRefPrefix; + size_t strategy; + size_t useRowMatchFinder; + size_t dedicatedDictSearch; + size_t dtlm; + size_t tfp; + ZSTD_rust_loadDictionaryContent_assert_f assertCParams; + ZSTD_rust_loadDictionaryContent_assertWindowEmpty_f assertWindowEmpty; + ZSTD_rust_loadDictionaryContent_windowUpdate_f windowUpdate; + ZSTD_rust_loadDictionaryContent_setLdmLoadedDictEnd_f setLdmLoadedDictEnd; + ZSTD_rust_loadDictionaryContent_publishMatchState_f publishMatchState; + ZSTD_rust_loadDictionaryContent_fillLdm_f fillLdm; + ZSTD_rust_loadDictionaryContent_overflowCorrect_f overflowCorrect; + ZSTD_rust_loadDictionaryContent_fillTable_f fillHashTable; + ZSTD_rust_loadDictionaryContent_fillTable_f fillDoubleHashTable; + ZSTD_rust_loadDictionaryContent_loadMatch_f loadDedicated; + ZSTD_rust_loadDictionaryContent_loadMatch_f loadRow; + ZSTD_rust_loadDictionaryContent_loadMatch_f loadChain; + ZSTD_rust_loadDictionaryContent_loadTree_f loadTree; + ZSTD_rust_loadDictionaryContent_publishFinalIndex_f publishFinalIndex; + ZSTD_rust_loadDictionaryContent_assert_f assertLazyConfiguration; + ZSTD_rust_loadDictionaryContent_assert_f assertInvalidStrategy; +} ZSTD_rust_loadDictionaryContentState; + +#define ZSTD_RUST_LOAD_DICT_ASSERT(name, condition) \ + typedef char name[(condition) ? 1 : -1] +ZSTD_RUST_LOAD_DICT_ASSERT(ZSTD_rust_load_dict_context_offset, + offsetof(ZSTD_rust_loadDictionaryContentState, + callbackContext) == 0); +ZSTD_RUST_LOAD_DICT_ASSERT(ZSTD_rust_load_dict_scalar_offset, + offsetof(ZSTD_rust_loadDictionaryContentState, + currentMax) == sizeof(void*)); +ZSTD_RUST_LOAD_DICT_ASSERT(ZSTD_rust_load_dict_callback_offset, + offsetof(ZSTD_rust_loadDictionaryContentState, + assertCParams) == 18 * sizeof(void*)); +ZSTD_RUST_LOAD_DICT_ASSERT(ZSTD_rust_load_dict_final_callback_offset, + offsetof(ZSTD_rust_loadDictionaryContentState, + publishFinalIndex) == 31 * sizeof(void*)); +ZSTD_RUST_LOAD_DICT_ASSERT(ZSTD_rust_load_dict_lazy_callback_offset, + offsetof(ZSTD_rust_loadDictionaryContentState, + assertLazyConfiguration) == 32 * sizeof(void*)); +ZSTD_RUST_LOAD_DICT_ASSERT(ZSTD_rust_load_dict_invalid_callback_offset, + offsetof(ZSTD_rust_loadDictionaryContentState, + assertInvalidStrategy) == 33 * sizeof(void*)); +ZSTD_RUST_LOAD_DICT_ASSERT(ZSTD_rust_load_dict_state_size, + sizeof(ZSTD_rust_loadDictionaryContentState) + == 34 * sizeof(void*)); +#undef ZSTD_RUST_LOAD_DICT_ASSERT + +size_t ZSTD_rust_loadDictionaryContent( + const ZSTD_rust_loadDictionaryContentState* state, + const void* src, size_t srcSize); size_t ZSTD_rust_compressInsertDictionary( ZSTD_compressedBlockState_t* bs, void* matchState, void* ldmState, void* workspaceState, @@ -5698,174 +5787,267 @@ size_t ZSTD_compressBlock(ZSTD_CCtx* cctx, void* dst, size_t dstCapacity, const return ZSTD_compressBlock_deprecated(cctx, dst, dstCapacity, src, srcSize); } -/*! ZSTD_loadDictionaryContent() : - * @return : 0, or an error code - */ -static size_t -ZSTD_loadDictionaryContent(ZSTD_MatchState_t* ms, - ldmState_t* ls, - ZSTD_cwksp* ws, - ZSTD_CCtx_params const* params, - const void* src, size_t srcSize, - ZSTD_dictTableLoadMethod_e dtlm, - ZSTD_tableFillPurpose_e tfp) +typedef struct { + ZSTD_MatchState_t* matchState; + ldmState_t* ldmState; + ZSTD_cwksp* workspace; + const ZSTD_CCtx_params* params; +} ZSTD_loadDictionaryContent_context; + +static void ZSTD_loadDictionaryContent_assertCParams(void* opaque) { - const BYTE* ip = (const BYTE*) src; - const BYTE* const iend = ip + srcSize; - int const loadLdmDict = params->ldmParams.enableLdm == ZSTD_ps_enable && ls != NULL; + ZSTD_loadDictionaryContent_context const* const context = + (const ZSTD_loadDictionaryContent_context*)opaque; + ZSTD_assertEqualCParams(context->params->cParams, + context->matchState->cParams); +} - /* Assert that the ms params match the params we're being given */ - ZSTD_assertEqualCParams(params->cParams, ms->cParams); - - { /* Ensure large dictionaries can't cause index overflow */ - - /* Allow the dictionary to set indices up to exactly ZSTD_CURRENT_MAX. - * Dictionaries right at the edge will immediately trigger overflow - * correction, but I don't want to insert extra constraints here. - */ - U32 maxDictSize = ZSTD_CURRENT_MAX - ZSTD_WINDOW_START_INDEX; - - int const CDictTaggedIndices = ZSTD_CDictIndicesAreTagged(¶ms->cParams); - if (CDictTaggedIndices && tfp == ZSTD_tfp_forCDict) { - /* Some dictionary matchfinders in zstd use "short cache", - * which treats the lower ZSTD_SHORT_CACHE_TAG_BITS of each - * CDict hashtable entry as a tag rather than as part of an index. - * When short cache is used, we need to truncate the dictionary - * so that its indices don't overlap with the tag. */ - U32 const shortCacheMaxDictSize = (1u << (32 - ZSTD_SHORT_CACHE_TAG_BITS)) - ZSTD_WINDOW_START_INDEX; - maxDictSize = MIN(maxDictSize, shortCacheMaxDictSize); - assert(!loadLdmDict); - } - - /* If the dictionary is too large, only load the suffix of the dictionary. */ - if (srcSize > maxDictSize) { - ip = iend - maxDictSize; - src = ip; - srcSize = maxDictSize; - } +static void ZSTD_loadDictionaryContent_assertWindowEmpty(void* opaque, int ldm) +{ + ZSTD_loadDictionaryContent_context const* const context = + (const ZSTD_loadDictionaryContent_context*)opaque; + if (ldm) { + assert(context->ldmState != NULL); + assert(ZSTD_window_isEmpty(context->ldmState->window)); + } else { + assert(ZSTD_window_isEmpty(context->matchState->window)); } +} - if (srcSize > ZSTD_CHUNKSIZE_MAX) { - /* We must have cleared our windows when our source is this large. */ - assert(ZSTD_window_isEmpty(ms->window)); - if (loadLdmDict) assert(ZSTD_window_isEmpty(ls->window)); +static void ZSTD_loadDictionaryContent_windowUpdate( + void* opaque, int ldm, const void* src, size_t srcSize) +{ + ZSTD_loadDictionaryContent_context const* const context = + (const ZSTD_loadDictionaryContent_context*)opaque; + if (ldm) { + assert(context->ldmState != NULL); + ZSTD_window_update(&context->ldmState->window, src, srcSize, + /* forceNonContiguous */ 0); + } else { + ZSTD_window_update(&context->matchState->window, src, srcSize, + /* forceNonContiguous */ 0); } - ZSTD_window_update(&ms->window, src, srcSize, /* forceNonContiguous */ 0); +} - DEBUGLOG(4, "ZSTD_loadDictionaryContent: useRowMatchFinder=%d", (int)params->useRowMatchFinder); +static void ZSTD_loadDictionaryContent_setLdmLoadedDictEnd( + void* opaque, const void* iend, int forceWindow) +{ + ZSTD_loadDictionaryContent_context const* const context = + (const ZSTD_loadDictionaryContent_context*)opaque; + assert(context->ldmState != NULL); + context->ldmState->loadedDictEnd = forceWindow + ? 0 + : (U32)((const BYTE*)iend - context->ldmState->window.base); +} - if (loadLdmDict) { /* Load the entire dict into LDM matchfinders. */ - DEBUGLOG(4, "ZSTD_loadDictionaryContent: Trigger loadLdmDict"); - ZSTD_window_update(&ls->window, src, srcSize, /* forceNonContiguous */ 0); - ls->loadedDictEnd = params->forceWindow ? 0 : (U32)(iend - ls->window.base); - ZSTD_ldm_fillHashTable(ls, ip, iend, ¶ms->ldmParams); - DEBUGLOG(4, "ZSTD_loadDictionaryContent: ZSTD_ldm_fillHashTable completes"); - } +static void ZSTD_loadDictionaryContent_publishMatchState( + void* opaque, const void* ip, const void* iend, + int forceWindow, int deterministicRefPrefix) +{ + ZSTD_loadDictionaryContent_context const* const context = + (const ZSTD_loadDictionaryContent_context*)opaque; + context->matchState->nextToUpdate = + (U32)((const BYTE*)ip - context->matchState->window.base); + context->matchState->loadedDictEnd = forceWindow + ? 0 + : (U32)((const BYTE*)iend - context->matchState->window.base); + context->matchState->forceNonContiguous = deterministicRefPrefix; +} - /* If the dict is larger than we can reasonably index in our tables, only load the suffix. */ - { U32 maxDictSize = 1U << MIN(MAX(params->cParams.hashLog + 3, params->cParams.chainLog + 1), 31); - if (srcSize > maxDictSize) { - ip = iend - maxDictSize; - src = ip; - srcSize = maxDictSize; - } - } +static void ZSTD_loadDictionaryContent_fillLdm( + void* opaque, const void* ip, const void* iend) +{ + ZSTD_loadDictionaryContent_context const* const context = + (const ZSTD_loadDictionaryContent_context*)opaque; + assert(context->ldmState != NULL); + ZSTD_ldm_fillHashTable(context->ldmState, (const BYTE*)ip, (const BYTE*)iend, + &context->params->ldmParams); + DEBUGLOG(4, "ZSTD_loadDictionaryContent: ZSTD_ldm_fillHashTable completes"); +} - ms->nextToUpdate = (U32)(ip - ms->window.base); - ms->loadedDictEnd = params->forceWindow ? 0 : (U32)(iend - ms->window.base); - ms->forceNonContiguous = params->deterministicRefPrefix; +static void ZSTD_loadDictionaryContent_overflowCorrect( + void* opaque, const void* ip, const void* iend) +{ + ZSTD_loadDictionaryContent_context const* const context = + (const ZSTD_loadDictionaryContent_context*)opaque; + ZSTD_overflowCorrectIfNeeded(context->matchState, context->workspace, + context->params, ip, iend); +} - if (srcSize <= HASH_READ_SIZE) return 0; +static void ZSTD_loadDictionaryContent_fillHashTable( + void* opaque, const void* iend, int dtlm, int tfp) +{ + ZSTD_loadDictionaryContent_context const* const context = + (const ZSTD_loadDictionaryContent_context*)opaque; + ZSTD_fillHashTable(context->matchState, (const BYTE*)iend, + (ZSTD_dictTableLoadMethod_e)dtlm, + (ZSTD_tableFillPurpose_e)tfp); +} - ZSTD_overflowCorrectIfNeeded(ms, ws, params, ip, iend); - - switch(params->cParams.strategy) - { - case ZSTD_fast: - ZSTD_fillHashTable(ms, iend, dtlm, tfp); - break; - case ZSTD_dfast: +static void ZSTD_loadDictionaryContent_fillDoubleHashTable( + void* opaque, const void* iend, int dtlm, int tfp) +{ + ZSTD_loadDictionaryContent_context const* const context = + (const ZSTD_loadDictionaryContent_context*)opaque; #ifndef ZSTD_EXCLUDE_DFAST_BLOCK_COMPRESSOR - ZSTD_fillDoubleHashTable(ms, iend, dtlm, tfp); + ZSTD_fillDoubleHashTable(context->matchState, (const BYTE*)iend, + (ZSTD_dictTableLoadMethod_e)dtlm, + (ZSTD_tableFillPurpose_e)tfp); #else - assert(0); /* shouldn't be called: cparams should've been adjusted. */ + (void)context; (void)iend; (void)dtlm; (void)tfp; + assert(0); /* shouldn't be called: cparams should've been adjusted. */ #endif - break; +} - case ZSTD_greedy: - case ZSTD_lazy: - case ZSTD_lazy2: +static void ZSTD_loadDictionaryContent_loadDedicated(void* opaque, const void* ip) +{ + ZSTD_loadDictionaryContent_context const* const context = + (const ZSTD_loadDictionaryContent_context*)opaque; #if !defined(ZSTD_EXCLUDE_GREEDY_BLOCK_COMPRESSOR) \ || !defined(ZSTD_EXCLUDE_LAZY_BLOCK_COMPRESSOR) \ || !defined(ZSTD_EXCLUDE_LAZY2_BLOCK_COMPRESSOR) - assert(srcSize >= HASH_READ_SIZE); - if (ms->dedicatedDictSearch) { - assert(ms->chainTable != NULL); - ZSTD_dedicatedDictSearch_lazy_loadDictionary(ms, iend-HASH_READ_SIZE); - } else { - assert(params->useRowMatchFinder != ZSTD_ps_auto); - if (params->useRowMatchFinder == ZSTD_ps_enable) { - size_t const tagTableSize = ((size_t)1 << params->cParams.hashLog); - ZSTD_memset(ms->tagTable, 0, tagTableSize); - ZSTD_row_update(ms, iend-HASH_READ_SIZE); - DEBUGLOG(4, "Using row-based hash table for lazy dict"); - } else { - ZSTD_insertAndFindFirstIndex(ms, iend-HASH_READ_SIZE); - DEBUGLOG(4, "Using chain-based hash table for lazy dict"); - } - } + assert(context->matchState->chainTable != NULL); + ZSTD_dedicatedDictSearch_lazy_loadDictionary(context->matchState, + (const BYTE*)ip); #else - assert(0); /* shouldn't be called: cparams should've been adjusted. */ + (void)context; (void)ip; + assert(0); /* shouldn't be called: cparams should've been adjusted. */ #endif - break; +} - case ZSTD_btlazy2: /* we want the dictionary table fully sorted */ - case ZSTD_btopt: - case ZSTD_btultra: - case ZSTD_btultra2: +static void ZSTD_loadDictionaryContent_loadRow(void* opaque, const void* ip) +{ + ZSTD_loadDictionaryContent_context const* const context = + (const ZSTD_loadDictionaryContent_context*)opaque; +#if !defined(ZSTD_EXCLUDE_GREEDY_BLOCK_COMPRESSOR) \ + || !defined(ZSTD_EXCLUDE_LAZY_BLOCK_COMPRESSOR) \ + || !defined(ZSTD_EXCLUDE_LAZY2_BLOCK_COMPRESSOR) + size_t const tagTableSize = ((size_t)1 << context->params->cParams.hashLog); + ZSTD_memset(context->matchState->tagTable, 0, tagTableSize); + ZSTD_row_update(context->matchState, (const BYTE*)ip); + DEBUGLOG(4, "Using row-based hash table for lazy dict"); +#else + (void)context; (void)ip; + assert(0); /* shouldn't be called: cparams should've been adjusted. */ +#endif +} + +static void ZSTD_loadDictionaryContent_loadChain(void* opaque, const void* ip) +{ + ZSTD_loadDictionaryContent_context const* const context = + (const ZSTD_loadDictionaryContent_context*)opaque; +#if !defined(ZSTD_EXCLUDE_GREEDY_BLOCK_COMPRESSOR) \ + || !defined(ZSTD_EXCLUDE_LAZY_BLOCK_COMPRESSOR) \ + || !defined(ZSTD_EXCLUDE_LAZY2_BLOCK_COMPRESSOR) + ZSTD_insertAndFindFirstIndex(context->matchState, (const BYTE*)ip); + DEBUGLOG(4, "Using chain-based hash table for lazy dict"); +#else + (void)context; (void)ip; + assert(0); /* shouldn't be called: cparams should've been adjusted. */ +#endif +} + +static void ZSTD_loadDictionaryContent_loadTree( + void* opaque, const void* ip, const void* iend) +{ + ZSTD_loadDictionaryContent_context const* const context = + (const ZSTD_loadDictionaryContent_context*)opaque; #if !defined(ZSTD_EXCLUDE_BTLAZY2_BLOCK_COMPRESSOR) \ || !defined(ZSTD_EXCLUDE_BTOPT_BLOCK_COMPRESSOR) \ || !defined(ZSTD_EXCLUDE_BTULTRA_BLOCK_COMPRESSOR) - assert(srcSize >= HASH_READ_SIZE); - DEBUGLOG(4, "Fill %u bytes into the Binary Tree", (unsigned)srcSize); - ZSTD_updateTree(ms, iend-HASH_READ_SIZE, iend); + DEBUGLOG(4, "Fill %u bytes into the Binary Tree", + (unsigned)((const BYTE*)iend - (const BYTE*)ip)); + ZSTD_updateTree(context->matchState, (const BYTE*)ip, (const BYTE*)iend); #else - assert(0); /* shouldn't be called: cparams should've been adjusted. */ + (void)context; (void)ip; (void)iend; + assert(0); /* shouldn't be called: cparams should've been adjusted. */ #endif - break; +} - default: - assert(0); /* not possible : not a valid strategy id */ - } +static void ZSTD_loadDictionaryContent_publishFinalIndex(void* opaque, const void* iend) +{ + ZSTD_loadDictionaryContent_context const* const context = + (const ZSTD_loadDictionaryContent_context*)opaque; + context->matchState->nextToUpdate = + (U32)((const BYTE*)iend - context->matchState->window.base); +} - ms->nextToUpdate = (U32)(iend - ms->window.base); - return 0; +static void ZSTD_loadDictionaryContent_assertLazyConfiguration(void* opaque) +{ + ZSTD_loadDictionaryContent_context const* const context = + (const ZSTD_loadDictionaryContent_context*)opaque; + assert(context->params->useRowMatchFinder != ZSTD_ps_auto); +} + +static void ZSTD_loadDictionaryContent_assertInvalidStrategy(void* opaque) +{ + (void)opaque; + assert(0); /* not possible: not a valid strategy id */ } -/* Dictionary entropy-header loading lives in Rust; C keeps the public internal - * symbol and the dictionary-content orchestration around it. */ +/* Dictionary entropy-header loading and content orchestration live in Rust; + * this keeps the original internal symbol available to C callers. */ size_t ZSTD_loadCEntropy(ZSTD_compressedBlockState_t* bs, void* workspace, const void* const dict, size_t dictSize) { return ZSTD_rust_loadCEntropy(bs, workspace, dict, dictSize); } -/* Keep the match-state/content operation private to C. Rust passes only - * opaque pointers here after it has selected the dictionary path. */ +/* Package only scalar policy inputs and opaque private-operation callbacks for + * the Rust dictionary-content orchestrator. */ static size_t ZSTD_loadDictionaryContent_callback( void* matchState, void* ldmState, void* workspaceState, const void* params, const void* src, size_t srcSize, int dtlm, int tfp) { - return ZSTD_loadDictionaryContent( - (ZSTD_MatchState_t*)matchState, - (ldmState_t*)ldmState, - (ZSTD_cwksp*)workspaceState, - (const ZSTD_CCtx_params*)params, - src, srcSize, - (ZSTD_dictTableLoadMethod_e)dtlm, - (ZSTD_tableFillPurpose_e)tfp); + ZSTD_loadDictionaryContent_context context; + ZSTD_rust_loadDictionaryContentState state; + ZSTD_MatchState_t* const ms = (ZSTD_MatchState_t*)matchState; + ldmState_t* const ls = (ldmState_t*)ldmState; + ZSTD_cwksp* const ws = (ZSTD_cwksp*)workspaceState; + ZSTD_CCtx_params const* const cctxParams = (const ZSTD_CCtx_params*)params; + + context.matchState = ms; + context.ldmState = ls; + context.workspace = ws; + context.params = cctxParams; + + state.callbackContext = &context; + state.currentMax = ZSTD_CURRENT_MAX; + state.windowStartIndex = ZSTD_WINDOW_START_INDEX; + state.shortCacheTagBits = ZSTD_SHORT_CACHE_TAG_BITS; + state.chunkSizeMax = ZSTD_CHUNKSIZE_MAX; + state.hashReadSize = HASH_READ_SIZE; + state.hashLog = cctxParams->cParams.hashLog; + state.chainLog = cctxParams->cParams.chainLog; + state.cdictIndicesTagged = ZSTD_CDictIndicesAreTagged(&cctxParams->cParams); + state.ldmEnabled = cctxParams->ldmParams.enableLdm == ZSTD_ps_enable; + state.hasLdmState = ls != NULL; + state.forceWindow = cctxParams->forceWindow; + state.deterministicRefPrefix = cctxParams->deterministicRefPrefix; + state.strategy = cctxParams->cParams.strategy; + state.useRowMatchFinder = cctxParams->useRowMatchFinder; + state.dedicatedDictSearch = ms->dedicatedDictSearch; + state.dtlm = (size_t)dtlm; + state.tfp = (size_t)tfp; + state.assertCParams = ZSTD_loadDictionaryContent_assertCParams; + state.assertWindowEmpty = ZSTD_loadDictionaryContent_assertWindowEmpty; + state.windowUpdate = ZSTD_loadDictionaryContent_windowUpdate; + state.setLdmLoadedDictEnd = ZSTD_loadDictionaryContent_setLdmLoadedDictEnd; + state.publishMatchState = ZSTD_loadDictionaryContent_publishMatchState; + state.fillLdm = ZSTD_loadDictionaryContent_fillLdm; + state.overflowCorrect = ZSTD_loadDictionaryContent_overflowCorrect; + state.fillHashTable = ZSTD_loadDictionaryContent_fillHashTable; + state.fillDoubleHashTable = ZSTD_loadDictionaryContent_fillDoubleHashTable; + state.loadDedicated = ZSTD_loadDictionaryContent_loadDedicated; + state.loadRow = ZSTD_loadDictionaryContent_loadRow; + state.loadChain = ZSTD_loadDictionaryContent_loadChain; + state.loadTree = ZSTD_loadDictionaryContent_loadTree; + state.publishFinalIndex = ZSTD_loadDictionaryContent_publishFinalIndex; + state.assertLazyConfiguration = ZSTD_loadDictionaryContent_assertLazyConfiguration; + state.assertInvalidStrategy = ZSTD_loadDictionaryContent_assertInvalidStrategy; + return ZSTD_rust_loadDictionaryContent(&state, src, srcSize); } /** ZSTD_compress_insertDictionary() : diff --git a/rust/src/zstd_compress_dictionary.rs b/rust/src/zstd_compress_dictionary.rs index 862bf20ef..ac5d10115 100644 --- a/rust/src/zstd_compress_dictionary.rs +++ b/rust/src/zstd_compress_dictionary.rs @@ -6,8 +6,8 @@ //! //! This is the Rust leaf for `ZSTD_dictNCountRepeat()` and //! `ZSTD_loadCEntropy()` in `lib/compress/zstd_compress.c`. C retains the -//! original symbol through a thin compatibility shim and keeps dictionary -//! content loading and compression-context ownership on its side. +//! original symbols through thin compatibility shims and keeps the +//! configuration-sensitive dictionary-content operations on its side. use crate::bits::ZSTD_highbit32; use crate::common::{LL_FSE_LOG, MAX_LL, MAX_ML, MAX_OFF, ML_FSE_LOG, OFF_FSE_LOG}; @@ -61,6 +61,328 @@ pub type LoadDictionaryContentFn = unsafe extern "C" fn( tfp: c_int, ) -> usize; +type LoadDictionaryContentAssertFn = unsafe extern "C" fn(context: *mut c_void); +type LoadDictionaryContentAssertWindowEmptyFn = + unsafe extern "C" fn(context: *mut c_void, ldm: c_int); +type LoadDictionaryContentWindowUpdateFn = + unsafe extern "C" fn(context: *mut c_void, ldm: c_int, src: *const c_void, src_size: usize); +type LoadDictionaryContentSetLdmLoadedDictEndFn = + unsafe extern "C" fn(context: *mut c_void, iend: *const c_void, force_window: c_int); +type LoadDictionaryContentPublishMatchStateFn = unsafe extern "C" fn( + context: *mut c_void, + ip: *const c_void, + iend: *const c_void, + force_window: c_int, + deterministic_ref_prefix: c_int, +); +type LoadDictionaryContentFillLdmFn = + unsafe extern "C" fn(context: *mut c_void, ip: *const c_void, iend: *const c_void); +type LoadDictionaryContentOverflowCorrectFn = + unsafe extern "C" fn(context: *mut c_void, ip: *const c_void, iend: *const c_void); +type LoadDictionaryContentFillTableFn = + unsafe extern "C" fn(context: *mut c_void, iend: *const c_void, dtlm: c_int, tfp: c_int); +type LoadDictionaryContentLoadMatchFn = + unsafe extern "C" fn(context: *mut c_void, ip: *const c_void); +type LoadDictionaryContentLoadTreeFn = + unsafe extern "C" fn(context: *mut c_void, ip: *const c_void, iend: *const c_void); +type LoadDictionaryContentPublishFinalIndexFn = + unsafe extern "C" fn(context: *mut c_void, iend: *const c_void); + +/// Rust-owned policy projection for `ZSTD_loadDictionaryContent()`. +/// +/// All configuration-dependent C layouts stay behind callbacks. The scalar +/// fields are copied by the C adapter so Rust owns suffix selection, window +/// and LDM ordering, index publication, and strategy dispatch without seeing +/// `ZSTD_MatchState_t`, `ldmState_t`, or workspace internals. +#[repr(C)] +pub struct ZSTD_rust_loadDictionaryContentState { + callback_context: *mut c_void, + current_max: usize, + window_start_index: usize, + short_cache_tag_bits: usize, + chunk_size_max: usize, + hash_read_size: usize, + hash_log: usize, + chain_log: usize, + cdict_indices_tagged: usize, + ldm_enabled: usize, + has_ldm_state: usize, + force_window: usize, + deterministic_ref_prefix: usize, + strategy: usize, + use_row_match_finder: usize, + dedicated_dict_search: usize, + dtlm: usize, + tfp: usize, + assert_c_params: LoadDictionaryContentAssertFn, + assert_window_empty: LoadDictionaryContentAssertWindowEmptyFn, + window_update: LoadDictionaryContentWindowUpdateFn, + set_ldm_loaded_dict_end: LoadDictionaryContentSetLdmLoadedDictEndFn, + publish_match_state: LoadDictionaryContentPublishMatchStateFn, + fill_ldm: LoadDictionaryContentFillLdmFn, + overflow_correct: LoadDictionaryContentOverflowCorrectFn, + fill_hash_table: LoadDictionaryContentFillTableFn, + fill_double_hash_table: LoadDictionaryContentFillTableFn, + load_dedicated: LoadDictionaryContentLoadMatchFn, + load_row: LoadDictionaryContentLoadMatchFn, + load_chain: LoadDictionaryContentLoadMatchFn, + load_tree: LoadDictionaryContentLoadTreeFn, + publish_final_index: LoadDictionaryContentPublishFinalIndexFn, + assert_lazy_configuration: LoadDictionaryContentAssertFn, + assert_invalid_strategy: LoadDictionaryContentAssertFn, +} + +const _: () = { + assert!(size_of::() == size_of::()); + assert!(size_of::() == size_of::()); + assert!(size_of::() == size_of::()); + assert!(size_of::() == size_of::()); + assert!(size_of::() == size_of::()); + assert!(size_of::() == size_of::()); + assert!(size_of::() == size_of::()); + assert!(size_of::() == size_of::()); + assert!(size_of::() == size_of::()); + assert!(size_of::() == size_of::()); + assert!(size_of::() == size_of::()); + assert!(offset_of!(ZSTD_rust_loadDictionaryContentState, callback_context) == 0); + assert!(offset_of!(ZSTD_rust_loadDictionaryContentState, current_max) == size_of::()); + assert!( + offset_of!(ZSTD_rust_loadDictionaryContentState, assert_c_params) + == 18 * size_of::() + ); + assert!( + offset_of!(ZSTD_rust_loadDictionaryContentState, publish_final_index) + == 31 * size_of::() + ); + assert!( + offset_of!( + ZSTD_rust_loadDictionaryContentState, + assert_lazy_configuration + ) == 32 * size_of::() + ); + assert!( + offset_of!( + ZSTD_rust_loadDictionaryContentState, + assert_invalid_strategy + ) == 33 * size_of::() + ); + assert!(size_of::() == 34 * size_of::()); +}; + +const ZSTD_TFP_FOR_CDICT: usize = 1; +const ZSTD_PS_ENABLE: usize = 1; +const ZSTD_FAST_STRATEGY: usize = 1; +const ZSTD_DFAST_STRATEGY: usize = 2; +const ZSTD_GREEDY_STRATEGY: usize = 3; +const ZSTD_LAZY_STRATEGY: usize = 4; +const ZSTD_LAZY2_STRATEGY: usize = 5; +const ZSTD_BTLAZY2_STRATEGY: usize = 6; +const ZSTD_BTOPT_STRATEGY: usize = 7; +const ZSTD_BTULTRA_STRATEGY: usize = 8; +const ZSTD_BTULTRA2_STRATEGY: usize = 9; + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +enum DictionaryTablePolicy { + Fast, + DoubleFast, + Dedicated, + Row, + Chain, + Tree, + Invalid, +} + +#[inline] +fn dictionary_suffix(src_size: usize, max_size: usize) -> (usize, usize) { + if src_size > max_size { + (src_size - max_size, max_size) + } else { + (0, src_size) + } +} + +#[inline] +fn dictionary_initial_max_size( + current_max: usize, + window_start_index: usize, + short_cache_tag_bits: usize, + cdict_indices_tagged: usize, + tfp: usize, + load_ldm_dict: bool, +) -> usize { + let mut max_size = current_max - window_start_index; + if cdict_indices_tagged != 0 && tfp == ZSTD_TFP_FOR_CDICT { + let short_cache_max_size = + (1usize << (32 - short_cache_tag_bits) as u32) - window_start_index; + max_size = max_size.min(short_cache_max_size); + debug_assert!(!load_ldm_dict); + } + max_size +} + +#[inline] +fn dictionary_table_max_size(hash_log: usize, chain_log: usize) -> usize { + let max_shift = hash_log + .saturating_add(3) + .max(chain_log.saturating_add(1)) + .min(31); + 1usize.checked_shl(max_shift as u32).unwrap() +} + +#[inline] +fn dictionary_table_policy( + strategy: usize, + dedicated_dict_search: usize, + use_row_match_finder: usize, +) -> DictionaryTablePolicy { + match strategy { + ZSTD_FAST_STRATEGY => DictionaryTablePolicy::Fast, + ZSTD_DFAST_STRATEGY => DictionaryTablePolicy::DoubleFast, + ZSTD_GREEDY_STRATEGY | ZSTD_LAZY_STRATEGY | ZSTD_LAZY2_STRATEGY => { + if dedicated_dict_search != 0 { + DictionaryTablePolicy::Dedicated + } else if use_row_match_finder == ZSTD_PS_ENABLE { + DictionaryTablePolicy::Row + } else { + DictionaryTablePolicy::Chain + } + } + ZSTD_BTLAZY2_STRATEGY + | ZSTD_BTOPT_STRATEGY + | ZSTD_BTULTRA_STRATEGY + | ZSTD_BTULTRA2_STRATEGY => DictionaryTablePolicy::Tree, + _ => DictionaryTablePolicy::Invalid, + } +} + +unsafe fn load_dictionary_content( + state: &ZSTD_rust_loadDictionaryContentState, + src: *const c_void, + src_size: usize, +) -> usize { + unsafe { (state.assert_c_params)(state.callback_context) }; + + let load_ldm_dict = state.ldm_enabled != 0 && state.has_ldm_state != 0; + let src_bytes = src.cast::(); + let iend = unsafe { src_bytes.add(src_size) }; + + let (first_offset, mut loaded_src_size) = dictionary_suffix( + src_size, + dictionary_initial_max_size( + state.current_max, + state.window_start_index, + state.short_cache_tag_bits, + state.cdict_indices_tagged, + state.tfp, + load_ldm_dict, + ), + ); + let mut ip = unsafe { src_bytes.add(first_offset) }; + + if loaded_src_size > state.chunk_size_max { + unsafe { (state.assert_window_empty)(state.callback_context, 0) }; + if load_ldm_dict { + unsafe { (state.assert_window_empty)(state.callback_context, 1) }; + } + } + unsafe { (state.window_update)(state.callback_context, 0, ip.cast(), loaded_src_size) }; + + if load_ldm_dict { + unsafe { + (state.window_update)(state.callback_context, 1, ip.cast(), loaded_src_size); + (state.set_ldm_loaded_dict_end)( + state.callback_context, + iend.cast(), + (state.force_window != 0) as c_int, + ); + (state.fill_ldm)(state.callback_context, ip.cast(), iend.cast()); + } + } + + let (_, table_src_size) = dictionary_suffix( + loaded_src_size, + dictionary_table_max_size(state.hash_log, state.chain_log), + ); + if table_src_size != loaded_src_size { + ip = unsafe { iend.sub(table_src_size) }; + loaded_src_size = table_src_size; + } + + unsafe { + (state.publish_match_state)( + state.callback_context, + ip.cast(), + iend.cast(), + (state.force_window != 0) as c_int, + (state.deterministic_ref_prefix != 0) as c_int, + ) + }; + + if loaded_src_size <= state.hash_read_size { + return 0; + } + + unsafe { (state.overflow_correct)(state.callback_context, ip.cast(), iend.cast()) }; + + let hash_ip = unsafe { iend.sub(state.hash_read_size) }; + match dictionary_table_policy( + state.strategy, + state.dedicated_dict_search, + state.use_row_match_finder, + ) { + DictionaryTablePolicy::Fast => unsafe { + (state.fill_hash_table)( + state.callback_context, + iend.cast(), + state.dtlm as c_int, + state.tfp as c_int, + ) + }, + DictionaryTablePolicy::DoubleFast => unsafe { + (state.fill_double_hash_table)( + state.callback_context, + iend.cast(), + state.dtlm as c_int, + state.tfp as c_int, + ) + }, + DictionaryTablePolicy::Dedicated => unsafe { + (state.load_dedicated)(state.callback_context, hash_ip.cast()) + }, + DictionaryTablePolicy::Row => unsafe { + (state.assert_lazy_configuration)(state.callback_context); + (state.load_row)(state.callback_context, hash_ip.cast()) + }, + DictionaryTablePolicy::Chain => unsafe { + (state.assert_lazy_configuration)(state.callback_context); + (state.load_chain)(state.callback_context, hash_ip.cast()) + }, + DictionaryTablePolicy::Tree => unsafe { + (state.load_tree)(state.callback_context, hash_ip.cast(), iend.cast()) + }, + DictionaryTablePolicy::Invalid => unsafe { + (state.assert_invalid_strategy)(state.callback_context) + }, + } + + unsafe { (state.publish_final_index)(state.callback_context, iend.cast()) }; + 0 +} + +/// Orchestrate raw/full dictionary content loading while C retains private +/// state mutation and matchfinder operations behind narrow callbacks. +#[no_mangle] +pub unsafe extern "C" fn ZSTD_rust_loadDictionaryContent( + state: *const ZSTD_rust_loadDictionaryContentState, + src: *const c_void, + src_size: usize, +) -> usize { + if state.is_null() || src.is_null() { + return ERROR(ZstdErrorCode::Generic); + } + unsafe { load_dictionary_content(&*state, src, src_size) } +} + /// These callbacks are deliberately opaque: C retains the private CCtx, /// local/prefix dictionary layouts, allocation, and CDict lifetime rules. /// Rust owns the stage checks and the order in which clear/assign operations @@ -5251,4 +5573,39 @@ mod tests { assert_eq!(dict_n_count_repeat(&normalized, 3, 3), FSE_REPEAT_CHECK); assert_eq!(dict_n_count_repeat(&normalized, 2, 2), FSE_REPEAT_CHECK); } + + #[test] + fn dictionary_content_policy_preserves_both_suffix_limits() { + assert_eq!(dictionary_suffix(100, 20), (80, 20)); + assert_eq!(dictionary_suffix(20, 20), (0, 20)); + assert_eq!( + dictionary_initial_max_size(100, 2, 28, 1, ZSTD_TFP_FOR_CDICT, false), + 14 + ); + assert_eq!(dictionary_table_max_size(4, 5), 128); + } + + #[test] + fn dictionary_content_policy_selects_matchfinder_path() { + assert_eq!( + dictionary_table_policy(ZSTD_FAST_STRATEGY, 0, 0), + DictionaryTablePolicy::Fast + ); + assert_eq!( + dictionary_table_policy(ZSTD_GREEDY_STRATEGY, 1, ZSTD_PS_ENABLE), + DictionaryTablePolicy::Dedicated + ); + assert_eq!( + dictionary_table_policy(ZSTD_LAZY_STRATEGY, 0, ZSTD_PS_ENABLE), + DictionaryTablePolicy::Row + ); + assert_eq!( + dictionary_table_policy(ZSTD_LAZY2_STRATEGY, 0, 2), + DictionaryTablePolicy::Chain + ); + assert_eq!( + dictionary_table_policy(ZSTD_BTULTRA2_STRATEGY, 0, 0), + DictionaryTablePolicy::Tree + ); + } }