diff --git a/lib/compress/zstd_compress.c b/lib/compress/zstd_compress.c index 49c878679..a596183a6 100644 --- a/lib/compress/zstd_compress.c +++ b/lib/compress/zstd_compress.c @@ -139,7 +139,8 @@ int ZSTD_rust_params_resolveEnableLdm( 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); -U32 ZSTD_rust_params_dedicatedDictSearch_getHashLog(U32 hashLog); +ZSTD_compressionParameters +ZSTD_rust_params_dedicatedDictSearch_getCParams(ZSTD_compressionParameters cParams); ZSTD_compressionParameters ZSTD_rust_params_dedicatedDictSearch_revertCParams(ZSTD_compressionParameters cParams); int ZSTD_rust_params_shouldAttachDict(int strategy, int dedicatedDictSearch, @@ -5654,23 +5655,9 @@ int ZSTD_defaultCLevel(void) { return ZSTD_rust_params_defaultCLevel(); } static ZSTD_compressionParameters ZSTD_dedicatedDictSearch_getCParams(int const compressionLevel, size_t const dictSize) { - ZSTD_compressionParameters cParams = ZSTD_getCParams_internal(compressionLevel, 0, dictSize, ZSTD_cpm_createCDict); - switch (cParams.strategy) { - case ZSTD_fast: - case ZSTD_dfast: - break; - case ZSTD_greedy: - case ZSTD_lazy: - case ZSTD_lazy2: - cParams.hashLog = ZSTD_rust_params_dedicatedDictSearch_getHashLog(cParams.hashLog); - break; - case ZSTD_btlazy2: - case ZSTD_btopt: - case ZSTD_btultra: - case ZSTD_btultra2: - break; - } - return cParams; + ZSTD_compressionParameters const cParams = + ZSTD_getCParams_internal(compressionLevel, 0, dictSize, ZSTD_cpm_createCDict); + return ZSTD_rust_params_dedicatedDictSearch_getCParams(cParams); } /** diff --git a/rust/src/zstd_compress_params.rs b/rust/src/zstd_compress_params.rs index ff789a786..db95b6f9d 100644 --- a/rust/src/zstd_compress_params.rs +++ b/rust/src/zstd_compress_params.rs @@ -448,6 +448,19 @@ fn dedicated_dict_search_get_hash_log(hash_log: u32) -> u32 { hash_log.wrapping_add(ZSTD_LAZY_DDSS_BUCKET_LOG) } +#[inline] +fn dedicated_dict_search_get_cparams( + mut cparams: ZSTD_compressionParameters, +) -> ZSTD_compressionParameters { + match cparams.strategy { + ZSTD_GREEDY | ZSTD_LAZY | ZSTD_LAZY2 => { + cparams.hashLog = dedicated_dict_search_get_hash_log(cparams.hashLog); + } + _ => {} + } + cparams +} + #[inline] fn dedicated_dict_search_revert_hash_log(hash_log: u32) -> u32 { hash_log @@ -538,10 +551,12 @@ pub extern "C" fn ZSTD_rust_params_dedicatedDictSearchIsSupported( c_int::from(dedicated_dict_search_is_supported(cparams)) } -/// C ABI for the hash-log adjustment used by `ZSTD_dedicatedDictSearch_getCParams()`. +/// C ABI for `ZSTD_dedicatedDictSearch_getCParams()`. #[no_mangle] -pub extern "C" fn ZSTD_rust_params_dedicatedDictSearch_getHashLog(hash_log: u32) -> u32 { - dedicated_dict_search_get_hash_log(hash_log) +pub extern "C" fn ZSTD_rust_params_dedicatedDictSearch_getCParams( + cparams: ZSTD_compressionParameters, +) -> ZSTD_compressionParameters { + dedicated_dict_search_get_cparams(cparams) } /// C ABI for `ZSTD_dedicatedDictSearch_revertCParams()`. @@ -1437,7 +1452,7 @@ mod tests { let should_adjust = matches!(strategy, ZSTD_GREEDY | ZSTD_LAZY | ZSTD_LAZY2); let original = 12; let adjusted = if should_adjust { - ZSTD_rust_params_dedicatedDictSearch_getHashLog(original) + dedicated_dict_search_get_hash_log(original) } else { original }; @@ -1451,7 +1466,7 @@ mod tests { assert_eq!(reverted, original); } - assert_eq!(ZSTD_rust_params_dedicatedDictSearch_getHashLog(u32::MAX), 1); + assert_eq!(dedicated_dict_search_get_hash_log(u32::MAX), 1); assert_eq!( dedicated_dict_search_revert_hash_log(ZSTD_HASHLOG_MIN as u32), ZSTD_HASHLOG_MIN as u32 @@ -1468,6 +1483,76 @@ mod tests { assert_eq!(dedicated_dict_search_revert_hash_log(1), u32::MAX); } + #[test] + fn dedicated_dict_search_get_cparams_preserves_fields_and_strategy_policy() { + let strategies = [ + ZSTD_FAST, + ZSTD_DFAST, + ZSTD_GREEDY, + ZSTD_LAZY, + ZSTD_LAZY2, + ZSTD_BTLAZY2, + ZSTD_BTOPT, + ZSTD_BTULTRA, + ZSTD_BTULTRA2, + -1, + 99, + ]; + + for strategy in strategies { + let input = ZSTD_compressionParameters { + windowLog: 21, + chainLog: 19, + hashLog: 12, + searchLog: 5, + minMatch: 4, + targetLength: 48, + strategy, + }; + let output = ZSTD_rust_params_dedicatedDictSearch_getCParams(input); + let expected_hash_log = if matches!(strategy, ZSTD_GREEDY | ZSTD_LAZY | ZSTD_LAZY2) { + 14 + } else { + input.hashLog + }; + + assert_eq!(output.hashLog, expected_hash_log, "strategy {strategy}"); + assert_eq!(output.windowLog, input.windowLog, "strategy {strategy}"); + assert_eq!(output.chainLog, input.chainLog, "strategy {strategy}"); + assert_eq!(output.searchLog, input.searchLog, "strategy {strategy}"); + assert_eq!(output.minMatch, input.minMatch, "strategy {strategy}"); + assert_eq!( + output.targetLength, input.targetLength, + "strategy {strategy}" + ); + assert_eq!(output.strategy, input.strategy, "strategy {strategy}"); + } + + for strategy in [ZSTD_GREEDY, ZSTD_LAZY, ZSTD_LAZY2] { + let mut input = ZSTD_compressionParameters { + hashLog: u32::MAX, + strategy, + ..ZSTD_compressionParameters::default() + }; + assert_eq!( + ZSTD_rust_params_dedicatedDictSearch_getCParams(input).hashLog, + 1 + ); + + input.hashLog = 0; + assert_eq!( + ZSTD_rust_params_dedicatedDictSearch_getCParams(input).hashLog, + 2 + ); + + input.hashLog = 1; + assert_eq!( + ZSTD_rust_params_dedicatedDictSearch_getCParams(input).hashLog, + 3 + ); + } + } + #[test] fn dedicated_dict_search_revert_cparams_preserves_fields_and_strategy_policy() { let strategies = [