From 5449a4b3c07f64b2b58f7abc8c27f2c579595b5e Mon Sep 17 00:00:00 2001 From: ddidderr Date: Sun, 19 Jul 2026 16:53:31 +0200 Subject: [PATCH] feat(compress): move advanced CDict policy into Rust ZSTD_createCDict_advanced2 previously selected compression parameters, resolved dedicated-dictionary-search fallback, and ordered allocation and initialization entirely in C. That left a large public dictionary boundary outside the Rust rewrite and made the failure ordering implicit in the C wrapper. The Rust params API now owns advanced-CDict parameter preparation, including the dedicated-search override/fallback and row-matchfinder resolution. A Rust dictionary bridge owns the create -> init sequence and frees a created CDict when initialization fails. The C shim retains only the private workspace and custom-memory allocation, dictionary initialization, and teardown callbacks; opaque parameter-field pointers and compile-time layout assertions preserve the existing ABI. Focused probes cover policy publication and all callback failure-ordering paths, and the migration README records the remaining C boundary. Test Plan: - `rustfmt --check` and `cargo fmt --manifest-path rust/Cargo.toml -- --check` -- passed - `cargo test --manifest-path rust/Cargo.toml --release` -- 663 passed - `cargo clippy --manifest-path rust/Cargo.toml --release --all-targets -- -D warnings` -- passed - `make -B -C programs -j1 zstd` -- passed - `make -C tests -j1 test-zstream ZSTREAM_TESTTIME=-T1s` -- 84 tests and both fuzz rounds passed --- lib/compress/zstd_compress.c | 149 ++++++++---- rust/README.md | 8 +- rust/src/zstd_compress_dictionary.rs | 348 +++++++++++++++++++++++++++ rust/src/zstd_compress_params_api.rs | 236 +++++++++++++++++- 4 files changed, 684 insertions(+), 57 deletions(-) diff --git a/lib/compress/zstd_compress.c b/lib/compress/zstd_compress.c index eb78f9f29..044ab6d96 100644 --- a/lib/compress/zstd_compress.c +++ b/lib/compress/zstd_compress.c @@ -1195,8 +1195,6 @@ size_t ZSTD_rust_params_maxNbSeq(size_t blockSize, U32 minMatch, int useSequenceProducer); size_t ZSTD_rust_params_resolveMaxBlockSize(size_t maxBlockSize); size_t ZSTD_rust_params_getBlockSize(size_t maxBlockSize, U32 windowLog); -void ZSTD_rust_params_overrideCParams(ZSTD_compressionParameters* cParams, - const ZSTD_compressionParameters* overrides); int ZSTD_rust_params_resolveExternalSequenceValidation(int mode); int ZSTD_rust_params_rowMatchFinderSupported(int strategy); int ZSTD_rust_params_rowMatchFinderUsed(int strategy, int mode); @@ -1210,9 +1208,6 @@ int ZSTD_rust_params_resolveEnableLdm( int mode, ZSTD_compressionParameters cParams); int ZSTD_rust_params_resolveExternalRepcodeSearch(int mode, int cLevel); int ZSTD_rust_params_cdictIndicesAreTagged(ZSTD_compressionParameters cParams); -int ZSTD_rust_params_dedicatedDictSearchIsSupported(ZSTD_compressionParameters cParams); -ZSTD_compressionParameters -ZSTD_rust_params_dedicatedDictSearch_getCParams(ZSTD_compressionParameters cParams); ZSTD_compressionParameters ZSTD_rust_params_dedicatedDictSearch_revertCParams(ZSTD_compressionParameters cParams); int ZSTD_rust_params_getCParamMode(int cdict_present, int cdict_strategy, @@ -1784,6 +1779,56 @@ typedef char ZSTD_rust_init_cdict_state_layout[ && sizeof(ZSTD_rust_initCDictState) == 19 * sizeof(void*)) ? 1 : -1]; +typedef void* (*ZSTD_rust_createCDictAdvancedCreate_f)( + void* context, size_t dictSize, int dictLoadMethod, + const ZSTD_compressionParameters* cParams, + int useRowMatchFinder, int enableDedicatedDictSearch); +typedef size_t (*ZSTD_rust_createCDictAdvancedInit_f)( + void* context, void* cdict, const void* dict, size_t dictSize, + int dictLoadMethod, int dictContentType, + const ZSTD_CCtx_params* cctxParams); +typedef void (*ZSTD_rust_createCDictAdvancedFree_f)( + void* context, void* cdict); +typedef struct { + void* callbackContext; + ZSTD_CCtx_params* cctxParams; + const ZSTD_compressionParameters* cParams; + const int* enableDedicatedDictSearch; + const int* useRowMatchFinder; + U32 exclusionMask; + U32 ldmDefaultWindowLog; + ZSTD_rust_createCDictAdvancedCreate_f create; + ZSTD_rust_createCDictAdvancedInit_f init; + ZSTD_rust_createCDictAdvancedFree_f free; +} ZSTD_rust_createCDictAdvancedState; +void* ZSTD_rust_createCDictAdvanced( + const ZSTD_rust_createCDictAdvancedState* state, + const void* dict, size_t dictSize, + int dictLoadMethod, int dictContentType); +typedef char ZSTD_rust_create_cdict_advanced_state_layout[ + (offsetof(ZSTD_rust_createCDictAdvancedState, callbackContext) == 0 + && offsetof(ZSTD_rust_createCDictAdvancedState, cctxParams) + == sizeof(void*) + && offsetof(ZSTD_rust_createCDictAdvancedState, cParams) + == 2 * sizeof(void*) + && offsetof(ZSTD_rust_createCDictAdvancedState, enableDedicatedDictSearch) + == 3 * sizeof(void*) + && offsetof(ZSTD_rust_createCDictAdvancedState, useRowMatchFinder) + == 4 * sizeof(void*) + && offsetof(ZSTD_rust_createCDictAdvancedState, exclusionMask) + == 5 * sizeof(void*) + && offsetof(ZSTD_rust_createCDictAdvancedState, ldmDefaultWindowLog) + == 5 * sizeof(void*) + sizeof(U32) + && offsetof(ZSTD_rust_createCDictAdvancedState, create) + == 5 * sizeof(void*) + 2 * sizeof(U32) + && offsetof(ZSTD_rust_createCDictAdvancedState, init) + == 6 * sizeof(void*) + 2 * sizeof(U32) + && offsetof(ZSTD_rust_createCDictAdvancedState, free) + == 7 * sizeof(void*) + 2 * sizeof(U32) + && sizeof(ZSTD_rust_createCDictAdvancedState) + == 8 * sizeof(void*) + 2 * sizeof(U32)) + ? 1 : -1]; + typedef size_t (*ZSTD_rust_compressBeginResetInternal_f)( void* context, const void* params, U64 pledgedSrcSize, size_t loadedDictSize, int zbuff); @@ -2707,9 +2752,6 @@ size_t ZSTD_CCtx_setPledgedSrcSize(ZSTD_CCtx* cctx, unsigned long long pledgedSr &cctx->pledgedSrcSizePlusOne); } -static ZSTD_compressionParameters ZSTD_dedicatedDictSearch_getCParams( - int const compressionLevel, - size_t const dictSize); static void ZSTD_dedicatedDictSearch_revertCParams( ZSTD_compressionParameters* cParams); @@ -5558,6 +5600,36 @@ ZSTD_CDict* ZSTD_createCDict_advanced(const void* dictBuffer, size_t dictSize, &cctxParams, customMem); } +static void* ZSTD_rust_createCDictAdvanced_create( + void* context, size_t dictSize, int dictLoadMethod, + const ZSTD_compressionParameters* cParams, + int useRowMatchFinder, int enableDedicatedDictSearch) +{ + ZSTD_customMem const customMem = *(const ZSTD_customMem*)context; + return ZSTD_createCDict_advanced_internal( + dictSize, (ZSTD_dictLoadMethod_e)dictLoadMethod, *cParams, + (ZSTD_ParamSwitch_e)useRowMatchFinder, + enableDedicatedDictSearch, customMem); +} + +static size_t ZSTD_rust_createCDictAdvanced_init( + void* context, void* cdict, const void* dict, size_t dictSize, + int dictLoadMethod, int dictContentType, + const ZSTD_CCtx_params* cctxParams) +{ + (void)context; + return ZSTD_initCDict_internal( + (ZSTD_CDict*)cdict, dict, dictSize, + (ZSTD_dictLoadMethod_e)dictLoadMethod, + (ZSTD_dictContentType_e)dictContentType, *cctxParams); +} + +static void ZSTD_rust_createCDictAdvanced_free(void* context, void* cdict) +{ + (void)context; + ZSTD_freeCDict((ZSTD_CDict*)cdict); +} + ZSTD_CDict* ZSTD_createCDict_advanced2( const void* dict, size_t dictSize, ZSTD_dictLoadMethod_e dictLoadMethod, @@ -5565,47 +5637,27 @@ ZSTD_CDict* ZSTD_createCDict_advanced2( const ZSTD_CCtx_params* originalCctxParams, ZSTD_customMem customMem) { - ZSTD_CCtx_params cctxParams = *originalCctxParams; - ZSTD_compressionParameters cParams; - ZSTD_CDict* cdict; + ZSTD_CCtx_params cctxParams; + ZSTD_rust_createCDictAdvancedState state; DEBUGLOG(3, "ZSTD_createCDict_advanced2, dictSize=%u, mode=%u", (unsigned)dictSize, (unsigned)dictContentType); + if (originalCctxParams == NULL) return NULL; if (!customMem.customAlloc ^ !customMem.customFree) return NULL; - if (cctxParams.enableDedicatedDictSearch) { - cParams = ZSTD_dedicatedDictSearch_getCParams( - cctxParams.compressionLevel, dictSize); - ZSTD_rust_params_overrideCParams(&cParams, &cctxParams.cParams); - } else { - cParams = ZSTD_getCParamsFromCCtxParams( - &cctxParams, ZSTD_CONTENTSIZE_UNKNOWN, dictSize, ZSTD_cpm_createCDict); - } - - if (!ZSTD_rust_params_dedicatedDictSearchIsSupported(cParams)) { - /* Fall back to non-DDSS params */ - cctxParams.enableDedicatedDictSearch = 0; - cParams = ZSTD_getCParamsFromCCtxParams( - &cctxParams, ZSTD_CONTENTSIZE_UNKNOWN, dictSize, ZSTD_cpm_createCDict); - } - - DEBUGLOG(3, "ZSTD_createCDict_advanced2: DedicatedDictSearch=%u", cctxParams.enableDedicatedDictSearch); - cctxParams.cParams = cParams; - cctxParams.useRowMatchFinder = ZSTD_resolveRowMatchFinderMode(cctxParams.useRowMatchFinder, &cParams); - - cdict = ZSTD_createCDict_advanced_internal(dictSize, - dictLoadMethod, cctxParams.cParams, - cctxParams.useRowMatchFinder, cctxParams.enableDedicatedDictSearch, - customMem); - - if (!cdict || ZSTD_isError( ZSTD_initCDict_internal(cdict, - dict, dictSize, - dictLoadMethod, dictContentType, - cctxParams) )) { - ZSTD_freeCDict(cdict); - return NULL; - } - - return cdict; + cctxParams = *originalCctxParams; + state.callbackContext = &customMem; + state.cctxParams = &cctxParams; + state.cParams = &cctxParams.cParams; + state.enableDedicatedDictSearch = &cctxParams.enableDedicatedDictSearch; + state.useRowMatchFinder = (const int*)&cctxParams.useRowMatchFinder; + state.exclusionMask = ZSTD_getCParamsExclusionMask(); + state.ldmDefaultWindowLog = ZSTD_LDM_DEFAULT_WINDOW_LOG; + state.create = ZSTD_rust_createCDictAdvanced_create; + state.init = ZSTD_rust_createCDictAdvanced_init; + state.free = ZSTD_rust_createCDictAdvanced_free; + return (ZSTD_CDict*)ZSTD_rust_createCDictAdvanced( + &state, dict, dictSize, + (int)dictLoadMethod, (int)dictContentType); } static void* ZSTD_rust_createCDict_create( @@ -7072,13 +7124,6 @@ int ZSTD_maxCLevel(void) { return ZSTD_rust_params_maxCLevel(); } int ZSTD_minCLevel(void) { return ZSTD_rust_params_minCLevel(); } int ZSTD_defaultCLevel(void) { return ZSTD_rust_params_defaultCLevel(); } -static ZSTD_compressionParameters ZSTD_dedicatedDictSearch_getCParams(int const compressionLevel, size_t const dictSize) -{ - ZSTD_compressionParameters const cParams = - ZSTD_getCParams_internal(compressionLevel, 0, dictSize, ZSTD_cpm_createCDict); - return ZSTD_rust_params_dedicatedDictSearch_getCParams(cParams); -} - /** * Reverses the adjustment applied to cparams when enabling dedicated dict * search. This is used to recover the params set to be used in the working diff --git a/rust/README.md b/rust/README.md index c8f225efe..c6cbc571e 100644 --- a/rust/README.md +++ b/rust/README.md @@ -153,8 +153,9 @@ static-CCtx workspace validation and initialization dispatch, and public CDict constructor parameter selection, default-level normalization, and workspace/object teardown ordering, public static-CDict workspace validation and initialization dispatch, public advanced-CDict one-shot validation and -begin/end ordering, public -advanced-compression parameter validation/init policy, public +begin/end ordering, advanced-CDict parameter selection, dedicated-search +fallback, row-matchfinder resolution, and create/init/failure ordering, +public advanced-compression parameter validation/init policy, public usingDict parameter selection, dictionary-presence handling, and default-level normalization, and public usingCDict frame-policy construction and begin/end sequencing, legacy public CDict-begin frame-policy construction and @@ -163,7 +164,8 @@ parameter selection and default-level normalization, and public advanced-begin parameter validation and init-then-begin ordering now run in Rust. CDict advanced allocation/lifecycle machinery, private static-CCtx and static-CDict workspace construction and dictionary-content allocation/loading, -reset policy, +and advanced-CDict private workspace construction and dictionary-content +loading remain in C. Reset policy, private CCtx/matchfinder/workspace operations, and codec/adaptive-policy callbacks remain in C. CDict initialization ordering and scalar publication, shared compression-begin dictionary selection, CDict reset attach-versus-copy diff --git a/rust/src/zstd_compress_dictionary.rs b/rust/src/zstd_compress_dictionary.rs index 3c335ea9f..8c157ef99 100644 --- a/rust/src/zstd_compress_dictionary.rs +++ b/rust/src/zstd_compress_dictionary.rs @@ -500,6 +500,151 @@ pub unsafe extern "C" fn ZSTD_rust_createCDict( cdict } +type CreateCDictAdvancedCreateFn = unsafe extern "C" fn( + *mut c_void, + usize, + c_int, + *const ZSTD_compressionParameters, + c_int, + c_int, +) -> *mut c_void; +type CreateCDictAdvancedInitFn = unsafe extern "C" fn( + *mut c_void, + *mut c_void, + *const c_void, + usize, + c_int, + c_int, + *const ZSTD_CCtx_params, +) -> usize; +type CreateCDictAdvancedFreeFn = unsafe extern "C" fn(*mut c_void, *mut c_void); + +/// Explicit projection for the public `ZSTD_createCDict_advanced2` wrapper. +/// +/// Rust owns the context-free parameter preparation and callback ordering. C +/// retains custom-memory allocation, private workspace construction, CDict +/// initialization, and teardown behind narrow callbacks. The three field +/// pointers keep the private `ZSTD_CCtx_params` layout opaque here while +/// allowing C to publish the fields selected by the Rust parameter leaf. +#[repr(C)] +pub struct ZSTD_rust_createCDictAdvancedState { + callback_context: *mut c_void, + cctx_params: *mut ZSTD_CCtx_params, + cparams: *const ZSTD_compressionParameters, + enable_dedicated_dict_search: *const c_int, + use_row_match_finder: *const c_int, + exclusion_mask: u32, + ldm_default_window_log: u32, + create: CreateCDictAdvancedCreateFn, + init: CreateCDictAdvancedInitFn, + free: CreateCDictAdvancedFreeFn, +} + +const _: () = { + assert!(offset_of!(ZSTD_rust_createCDictAdvancedState, callback_context) == 0); + assert!(offset_of!(ZSTD_rust_createCDictAdvancedState, cctx_params) == size_of::()); + assert!(offset_of!(ZSTD_rust_createCDictAdvancedState, cparams) == 2 * size_of::()); + assert!( + offset_of!( + ZSTD_rust_createCDictAdvancedState, + enable_dedicated_dict_search + ) == 3 * size_of::() + ); + assert!( + offset_of!(ZSTD_rust_createCDictAdvancedState, use_row_match_finder) + == 4 * size_of::() + ); + assert!( + offset_of!(ZSTD_rust_createCDictAdvancedState, exclusion_mask) == 5 * size_of::() + ); + assert!( + offset_of!(ZSTD_rust_createCDictAdvancedState, ldm_default_window_log) + == 5 * size_of::() + size_of::() + ); + assert!( + offset_of!(ZSTD_rust_createCDictAdvancedState, create) + == size_of::<[usize; 5]>() + size_of::<[u32; 2]>() + ); + assert!( + offset_of!(ZSTD_rust_createCDictAdvancedState, init) + == size_of::<[usize; 6]>() + size_of::<[u32; 2]>() + ); + assert!( + offset_of!(ZSTD_rust_createCDictAdvancedState, free) + == size_of::<[usize; 7]>() + size_of::<[u32; 2]>() + ); + assert!( + size_of::() + == size_of::<[usize; 8]>() + size_of::<[u32; 2]>() + ); +}; + +/// Prepare advanced-CDict parameters and run the C-owned construction path. +#[no_mangle] +pub unsafe extern "C" fn ZSTD_rust_createCDictAdvanced( + state: *const ZSTD_rust_createCDictAdvancedState, + dict: *const c_void, + dict_size: usize, + dict_load_method: c_int, + dict_content_type: c_int, +) -> *mut c_void { + if state.is_null() { + return ptr::null_mut(); + } + let state = unsafe { &*state }; + if state.callback_context.is_null() + || state.cctx_params.is_null() + || state.cparams.is_null() + || state.enable_dedicated_dict_search.is_null() + || state.use_row_match_finder.is_null() + { + return ptr::null_mut(); + } + + let prepare_result = unsafe { + crate::zstd_compress_params_api::ZSTD_rust_params_prepareAdvancedCDict( + state.cctx_params, + dict_size, + state.ldm_default_window_log, + state.exclusion_mask, + ) + }; + if ERR_isError(prepare_result) { + return ptr::null_mut(); + } + + let cdict = unsafe { + (state.create)( + state.callback_context, + dict_size, + dict_load_method, + state.cparams, + *state.use_row_match_finder, + *state.enable_dedicated_dict_search, + ) + }; + if cdict.is_null() { + return ptr::null_mut(); + } + + let init_result = unsafe { + (state.init)( + state.callback_context, + cdict, + dict, + dict_size, + dict_load_method, + dict_content_type, + state.cctx_params.cast_const(), + ) + }; + if ERR_isError(init_result) { + unsafe { (state.free)(state.callback_context, cdict) }; + return ptr::null_mut(); + } + cdict +} + type FreeCDictWorkspaceFn = unsafe extern "C" fn(*mut c_void); type FreeCDictObjectFn = unsafe extern "C" fn(*mut c_void); @@ -3055,6 +3200,209 @@ mod tests { assert!(probe.published_cdict.is_null()); } + #[derive(Default)] + struct CreateCDictAdvancedProbe { + events: Vec<&'static str>, + cparams: *const ZSTD_compressionParameters, + use_row_match_finder: c_int, + enable_dedicated_dict_search: c_int, + cdict_result: *mut c_void, + init_cdict: *mut c_void, + init_dict: *const c_void, + init_dict_size: usize, + init_dict_load_method: c_int, + init_dict_content_type: c_int, + init_cctx_params: *const ZSTD_CCtx_params, + free_cdict: *mut c_void, + init_result: usize, + } + + unsafe extern "C" fn create_cdict_advanced_test_create( + context: *mut c_void, + _dict_size: usize, + _dict_load_method: c_int, + cparams: *const ZSTD_compressionParameters, + use_row_match_finder: c_int, + enable_dedicated_dict_search: c_int, + ) -> *mut c_void { + let probe = unsafe { &mut *context.cast::() }; + probe.events.push("create"); + probe.cparams = cparams; + probe.use_row_match_finder = use_row_match_finder; + probe.enable_dedicated_dict_search = enable_dedicated_dict_search; + probe.cdict_result + } + + unsafe extern "C" fn create_cdict_advanced_test_init( + context: *mut c_void, + cdict: *mut c_void, + dict: *const c_void, + dict_size: usize, + dict_load_method: c_int, + dict_content_type: c_int, + cctx_params: *const ZSTD_CCtx_params, + ) -> usize { + let probe = unsafe { &mut *context.cast::() }; + probe.events.push("init"); + probe.init_cdict = cdict; + probe.init_dict = dict; + probe.init_dict_size = dict_size; + probe.init_dict_load_method = dict_load_method; + probe.init_dict_content_type = dict_content_type; + probe.init_cctx_params = cctx_params; + probe.init_result + } + + unsafe extern "C" fn create_cdict_advanced_test_free(context: *mut c_void, cdict: *mut c_void) { + let probe = unsafe { &mut *context.cast::() }; + probe.events.push("free"); + probe.free_cdict = cdict; + } + + fn create_cdict_advanced_test_params() -> MaybeUninit { + let mut storage = MaybeUninit::::zeroed(); + unsafe { + assert_eq!( + crate::zstd_compress_params_api::ZSTD_CCtxParams_init( + storage.as_mut_ptr(), + ZSTD_rust_params_defaultCLevel(), + ), + 0 + ); + } + storage + } + + fn create_cdict_advanced_test_state( + probe: &mut CreateCDictAdvancedProbe, + cctx_params: *mut ZSTD_CCtx_params, + cparams: &ZSTD_compressionParameters, + enable_dedicated_dict_search: &c_int, + use_row_match_finder: &c_int, + ) -> ZSTD_rust_createCDictAdvancedState { + ZSTD_rust_createCDictAdvancedState { + callback_context: (probe as *mut CreateCDictAdvancedProbe).cast(), + cctx_params, + cparams, + enable_dedicated_dict_search, + use_row_match_finder, + exclusion_mask: 0, + ldm_default_window_log: 27, + create: create_cdict_advanced_test_create, + init: create_cdict_advanced_test_init, + free: create_cdict_advanced_test_free, + } + } + + #[test] + fn create_cdict_advanced_publishes_params_and_preserves_create_init_order() { + let mut probe = CreateCDictAdvancedProbe { + cdict_result: ptr::dangling_mut(), + ..Default::default() + }; + let mut params_storage = create_cdict_advanced_test_params(); + let cctx_params = params_storage.as_mut_ptr(); + let cparams = ZSTD_compressionParameters { + windowLog: 20, + chainLog: 19, + hashLog: 18, + searchLog: 5, + minMatch: 4, + targetLength: 16, + strategy: 3, + }; + let enable_dedicated_dict_search = 1; + let use_row_match_finder = 2; + let state = create_cdict_advanced_test_state( + &mut probe, + cctx_params, + &cparams, + &enable_dedicated_dict_search, + &use_row_match_finder, + ); + let dict = [1u8, 2, 3, 4]; + + let result = unsafe { + ZSTD_rust_createCDictAdvanced( + &state, + dict.as_ptr().cast(), + dict.len(), + ZSTD_DLM_BY_REF, + ZSTD_DCT_RAW_CONTENT, + ) + }; + + assert_eq!(result, probe.cdict_result); + assert_eq!(probe.events, ["create", "init"]); + assert_eq!(probe.cparams, &cparams); + assert_eq!(probe.use_row_match_finder, use_row_match_finder); + assert_eq!( + probe.enable_dedicated_dict_search, + enable_dedicated_dict_search + ); + assert_eq!(probe.init_cdict, probe.cdict_result); + assert_eq!(probe.init_dict, dict.as_ptr().cast()); + assert_eq!(probe.init_dict_size, dict.len()); + assert_eq!(probe.init_dict_load_method, ZSTD_DLM_BY_REF); + assert_eq!(probe.init_dict_content_type, ZSTD_DCT_RAW_CONTENT); + assert_eq!(probe.init_cctx_params, cctx_params.cast_const()); + } + + #[test] + fn create_cdict_advanced_does_not_init_or_free_after_creation_failure() { + let mut probe = CreateCDictAdvancedProbe::default(); + let mut params_storage = create_cdict_advanced_test_params(); + let cctx_params = params_storage.as_mut_ptr(); + let cparams = ZSTD_compressionParameters::default(); + let enable_dedicated_dict_search = 0; + let use_row_match_finder = 0; + let state = create_cdict_advanced_test_state( + &mut probe, + cctx_params, + &cparams, + &enable_dedicated_dict_search, + &use_row_match_finder, + ); + + let result = unsafe { + ZSTD_rust_createCDictAdvanced(&state, ptr::null(), 0, ZSTD_DLM_BY_REF, ZSTD_DCT_AUTO) + }; + + assert!(result.is_null()); + assert_eq!(probe.events, ["create"]); + assert!(probe.free_cdict.is_null()); + } + + #[test] + fn create_cdict_advanced_frees_after_initialization_failure() { + let cdict = ptr::dangling_mut::(); + let mut probe = CreateCDictAdvancedProbe { + cdict_result: cdict, + init_result: ERROR(ZstdErrorCode::MemoryAllocation), + ..Default::default() + }; + let mut params_storage = create_cdict_advanced_test_params(); + let cctx_params = params_storage.as_mut_ptr(); + let cparams = ZSTD_compressionParameters::default(); + let enable_dedicated_dict_search = 0; + let use_row_match_finder = 0; + let state = create_cdict_advanced_test_state( + &mut probe, + cctx_params, + &cparams, + &enable_dedicated_dict_search, + &use_row_match_finder, + ); + + let result = unsafe { + ZSTD_rust_createCDictAdvanced(&state, ptr::null(), 0, ZSTD_DLM_BY_REF, ZSTD_DCT_AUTO) + }; + + assert!(result.is_null()); + assert_eq!(probe.events, ["create", "init", "free"]); + assert_eq!(probe.free_cdict, cdict); + } + #[derive(Default)] struct FreeCDictProbe { events: Vec<&'static str>, diff --git a/rust/src/zstd_compress_params_api.rs b/rust/src/zstd_compress_params_api.rs index 68754c6d2..886f0c471 100644 --- a/rust/src/zstd_compress_params_api.rs +++ b/rust/src/zstd_compress_params_api.rs @@ -14,8 +14,12 @@ use crate::errors::{ZstdErrorCode, ERROR}; use crate::zstd_compress_params::{ ZSTD_bounds, ZSTD_compressionParameters, ZSTD_frameParameters, ZSTD_parameters, - ZSTD_rust_params_checkCParams, ZSTD_rust_params_getBounds, - ZSTD_rust_params_resolveMaxBlockSize, ZSTD_RUST_PS_DISABLE, ZSTD_RUST_PS_ENABLE, + ZSTD_rust_params_checkCParams, ZSTD_rust_params_dedicatedDictSearchIsSupported, + ZSTD_rust_params_dedicatedDictSearch_getCParams, ZSTD_rust_params_getBounds, + ZSTD_rust_params_getCParamsFromCCtxParams, ZSTD_rust_params_getCParamsInternal, + ZSTD_rust_params_overrideCParams, ZSTD_rust_params_resolveMaxBlockSize, + ZSTD_rust_params_resolveRowMatchFinderMode, ZSTD_CONTENTSIZE_UNKNOWN, + ZSTD_RUST_CPM_CREATE_CDICT, ZSTD_RUST_PS_DISABLE, ZSTD_RUST_PS_ENABLE, }; use std::mem::size_of; use std::os::raw::{c_int, c_void}; @@ -508,6 +512,72 @@ pub unsafe extern "C" fn ZSTD_rust_CCtxParams_setZstdParams( unsafe { set_zstd_params_impl(cctx_params, zstd_params) } } +/// Prepares the parameter subset used by `ZSTD_createCDict_advanced2()`. +/// +/// The C dictionary allocator remains responsible for allocation and +/// initialization. This leaf owns the context-free parameter selection, +/// dedicated-dictionary-search fallback, and final row-match-finder +/// resolution while keeping the `ZSTD_CCtx_params` fields private here. +#[no_mangle] +pub unsafe extern "C" fn ZSTD_rust_params_prepareAdvancedCDict( + params: *mut ZSTD_CCtx_params, + dictSize: usize, + ldmDefaultWindowLog: u32, + exclusionMask: u32, +) -> usize { + if params.is_null() { + return ERROR(ZstdErrorCode::Generic); + } + + let params = unsafe { &mut *params }; + let mut cparams = if params.enableDedicatedDictSearch != 0 { + let mut cparams = ZSTD_rust_params_getCParamsInternal( + params.compressionLevel, + 0, + dictSize, + ZSTD_RUST_CPM_CREATE_CDICT, + exclusionMask, + ); + cparams = ZSTD_rust_params_dedicatedDictSearch_getCParams(cparams); + unsafe { ZSTD_rust_params_overrideCParams(&mut cparams, ¶ms.cParams) }; + cparams + } else { + ZSTD_rust_params_getCParamsFromCCtxParams( + params.compressionLevel, + params.srcSizeHint, + ZSTD_CONTENTSIZE_UNKNOWN, + dictSize, + ZSTD_RUST_CPM_CREATE_CDICT, + params.ldmParams.enableLdm, + ldmDefaultWindowLog, + params.cParams, + params.useRowMatchFinder, + exclusionMask, + ) + }; + + if ZSTD_rust_params_dedicatedDictSearchIsSupported(cparams) == 0 { + params.enableDedicatedDictSearch = 0; + cparams = ZSTD_rust_params_getCParamsFromCCtxParams( + params.compressionLevel, + params.srcSizeHint, + ZSTD_CONTENTSIZE_UNKNOWN, + dictSize, + ZSTD_RUST_CPM_CREATE_CDICT, + params.ldmParams.enableLdm, + ldmDefaultWindowLog, + params.cParams, + params.useRowMatchFinder, + exclusionMask, + ); + } + + params.cParams = cparams; + params.useRowMatchFinder = + ZSTD_rust_params_resolveRowMatchFinderMode(params.useRowMatchFinder, cparams); + 0 +} + unsafe fn custom_calloc(size: usize, custom_mem: ZSTD_customMem) -> *mut c_void { if let Some(custom_alloc) = custom_mem.customAlloc { let allocation = unsafe { custom_alloc(custom_mem.opaque, size) }; @@ -1520,6 +1590,168 @@ mod tests { } } + #[test] + fn prepare_advanced_cdict_rejects_null_params() { + assert_eq!( + unsafe { ZSTD_rust_params_prepareAdvancedCDict(ptr::null_mut(), 32 * 1024, 27, 0) }, + ERROR(ZstdErrorCode::Generic) + ); + } + + #[test] + fn prepare_advanced_cdict_publishes_normal_policy() { + let mut storage = MaybeUninit::::zeroed(); + let params = storage.as_mut_ptr(); + let dict_size = 32 * 1024; + let ldm_default_window_log = 27; + let exclusion_mask = 0; + + unsafe { + assert_eq!(ZSTD_CCtxParams_init(params, DEFAULT_CLEVEL), 0); + (*params).srcSizeHint = 16 * 1024; + (*params).ldmParams.enableLdm = ZSTD_RUST_PS_ENABLE; + (*params).cParams = ZSTD_compressionParameters { + windowLog: 27, + chainLog: 20, + hashLog: 25, + searchLog: 5, + minMatch: 4, + targetLength: 16, + strategy: STRATEGY_GREEDY, + }; + (*params).useRowMatchFinder = PS_AUTO; + + let expected = ZSTD_rust_params_getCParamsFromCCtxParams( + (*params).compressionLevel, + (*params).srcSizeHint, + ZSTD_CONTENTSIZE_UNKNOWN, + dict_size, + ZSTD_RUST_CPM_CREATE_CDICT, + (*params).ldmParams.enableLdm, + ldm_default_window_log, + (*params).cParams, + (*params).useRowMatchFinder, + exclusion_mask, + ); + let expected_row = + ZSTD_rust_params_resolveRowMatchFinderMode((*params).useRowMatchFinder, expected); + + assert_eq!( + ZSTD_rust_params_prepareAdvancedCDict( + params, + dict_size, + ldm_default_window_log, + exclusion_mask, + ), + 0 + ); + assert_eq!((*params).enableDedicatedDictSearch, 0); + assert_eq!((*params).cParams, expected); + assert_eq!((*params).useRowMatchFinder, expected_row); + } + } + + #[test] + fn prepare_advanced_cdict_applies_dedicated_search_overrides() { + let mut storage = MaybeUninit::::zeroed(); + let params = storage.as_mut_ptr(); + let dict_size = 32 * 1024; + let ldm_default_window_log = 27; + let exclusion_mask = 0; + + unsafe { + assert_eq!(ZSTD_CCtxParams_init(params, DEFAULT_CLEVEL), 0); + (*params).enableDedicatedDictSearch = 1; + (*params).cParams = ZSTD_compressionParameters { + windowLog: 20, + chainLog: 20, + hashLog: 25, + searchLog: 5, + minMatch: 4, + targetLength: 32, + strategy: STRATEGY_GREEDY, + }; + (*params).useRowMatchFinder = PS_DISABLE; + + let mut expected = ZSTD_rust_params_getCParamsInternal( + (*params).compressionLevel, + 0, + dict_size, + ZSTD_RUST_CPM_CREATE_CDICT, + exclusion_mask, + ); + expected = ZSTD_rust_params_dedicatedDictSearch_getCParams(expected); + ZSTD_rust_params_overrideCParams(&mut expected, &(*params).cParams); + assert_eq!(ZSTD_rust_params_dedicatedDictSearchIsSupported(expected), 1); + + assert_eq!( + ZSTD_rust_params_prepareAdvancedCDict( + params, + dict_size, + ldm_default_window_log, + exclusion_mask, + ), + 0 + ); + assert_eq!((*params).enableDedicatedDictSearch, 1); + assert_eq!((*params).cParams, expected); + assert_eq!((*params).useRowMatchFinder, PS_DISABLE); + } + } + + #[test] + fn prepare_advanced_cdict_falls_back_from_unsupported_dedicated_search() { + let mut storage = MaybeUninit::::zeroed(); + let params = storage.as_mut_ptr(); + let dict_size = 32 * 1024; + let ldm_default_window_log = 27; + let exclusion_mask = 0; + + unsafe { + assert_eq!(ZSTD_CCtxParams_init(params, DEFAULT_CLEVEL), 0); + (*params).srcSizeHint = 16 * 1024; + (*params).enableDedicatedDictSearch = 1; + (*params).cParams = ZSTD_compressionParameters { + windowLog: 20, + chainLog: 20, + hashLog: 25, + searchLog: 5, + minMatch: 4, + targetLength: 32, + strategy: STRATEGY_BTOPT, + }; + (*params).useRowMatchFinder = PS_AUTO; + + let expected = ZSTD_rust_params_getCParamsFromCCtxParams( + (*params).compressionLevel, + (*params).srcSizeHint, + ZSTD_CONTENTSIZE_UNKNOWN, + dict_size, + ZSTD_RUST_CPM_CREATE_CDICT, + (*params).ldmParams.enableLdm, + ldm_default_window_log, + (*params).cParams, + (*params).useRowMatchFinder, + exclusion_mask, + ); + let expected_row = + ZSTD_rust_params_resolveRowMatchFinderMode((*params).useRowMatchFinder, expected); + + assert_eq!( + ZSTD_rust_params_prepareAdvancedCDict( + params, + dict_size, + ldm_default_window_log, + exclusion_mask, + ), + 0 + ); + assert_eq!((*params).enableDedicatedDictSearch, 0); + assert_eq!((*params).cParams, expected); + assert_eq!((*params).useRowMatchFinder, expected_row); + } + } + #[test] fn allocation_lifecycle_preserves_default_and_custom_memory_contracts() { unsafe extern "C" fn counting_alloc(opaque: *mut c_void, size: usize) -> *mut c_void {