diff --git a/lib/compress/zstd_compress.c b/lib/compress/zstd_compress.c index b7782611d..c7e38e550 100644 --- a/lib/compress/zstd_compress.c +++ b/lib/compress/zstd_compress.c @@ -1751,6 +1751,25 @@ typedef char ZSTD_rust_next_input_size_hint_mt_or_st_state_layout[ ? 1 : -1]; size_t ZSTD_rust_nextInputSizeHintMTorST( const ZSTD_rust_nextInputSizeHintMTorSTState* state); +typedef size_t (*ZSTD_rust_estimateCDictSizeAdvanced_f)( + void* context, size_t dictSize, ZSTD_compressionParameters cParams, + int dictLoadMethod); +typedef struct { + void* callbackContext; + const U32* exclusionMask; + ZSTD_rust_estimateCDictSizeAdvanced_f estimateAdvanced; +} ZSTD_rust_estimateCDictSizeState; +typedef char zstd_rust_estimate_cdict_size_state_layout[ + (offsetof(ZSTD_rust_estimateCDictSizeState, callbackContext) == 0 + && offsetof(ZSTD_rust_estimateCDictSizeState, exclusionMask) + == sizeof(void*) + && offsetof(ZSTD_rust_estimateCDictSizeState, estimateAdvanced) + == 2 * sizeof(void*) + && sizeof(ZSTD_rust_estimateCDictSizeState) == 3 * sizeof(void*)) + ? 1 : -1]; +size_t ZSTD_rust_estimateCDictSize( + const ZSTD_rust_estimateCDictSizeState* state, + size_t dictSize, int compressionLevel); size_t ZSTD_rust_sizeofCDict(size_t objectSize, size_t workspaceSize); size_t ZSTD_rust_sizeofLocalDict(int dictBufferPresent, size_t dictSize, size_t cdictSize); @@ -7104,10 +7123,23 @@ size_t ZSTD_estimateCDictSize_advanced( dictSize, cParams, (int)dictLoadMethod, &sizing); } +static size_t ZSTD_rust_estimateCDictSize_advanced( + void* context, size_t dictSize, ZSTD_compressionParameters cParams, + int dictLoadMethod) +{ + (void)context; + return ZSTD_estimateCDictSize_advanced( + dictSize, cParams, (ZSTD_dictLoadMethod_e)dictLoadMethod); +} + size_t ZSTD_estimateCDictSize(size_t dictSize, int compressionLevel) { - ZSTD_compressionParameters const cParams = ZSTD_getCParams_internal(compressionLevel, ZSTD_CONTENTSIZE_UNKNOWN, dictSize, ZSTD_cpm_createCDict); - return ZSTD_estimateCDictSize_advanced(dictSize, cParams, ZSTD_dlm_byCopy); + ZSTD_rust_estimateCDictSizeState state; + U32 const exclusionMask = ZSTD_getCParamsExclusionMask(); + state.callbackContext = NULL; + state.exclusionMask = &exclusionMask; + state.estimateAdvanced = ZSTD_rust_estimateCDictSize_advanced; + return ZSTD_rust_estimateCDictSize(&state, dictSize, compressionLevel); } size_t ZSTD_sizeof_CDict(const ZSTD_CDict* cdict) diff --git a/rust/src/zstd_compress_dictionary.rs b/rust/src/zstd_compress_dictionary.rs index 9b1d58643..b5577b045 100644 --- a/rust/src/zstd_compress_dictionary.rs +++ b/rust/src/zstd_compress_dictionary.rs @@ -19,7 +19,8 @@ use crate::zstd_compress_params::{ ZSTD_compressionParameters, ZSTD_frameParameters, ZSTD_parameters, ZSTD_rustCDictSizing, ZSTD_rust_params_allocateChainTable, ZSTD_rust_params_defaultCLevel, ZSTD_rust_params_estimateCDictWorkspaceSize, ZSTD_rust_params_getCParams, - ZSTD_rust_params_getParamsInternal, ZSTD_rust_params_resolveRowMatchFinderMode, + ZSTD_rust_params_getCParamsInternal, ZSTD_rust_params_getParamsInternal, + ZSTD_rust_params_resolveRowMatchFinderMode, ZSTD_rust_params_rowMatchFinderUsed, ZSTD_rust_params_shouldAttachDict, ZSTD_CONTENTSIZE_UNKNOWN, ZSTD_RUST_CPM_CREATE_CDICT, ZSTD_RUST_CPM_NO_ATTACH_DICT, ZSTD_RUST_PS_AUTO, @@ -55,6 +56,66 @@ const fn dictionary_content_reservation_size(dict_size: usize) -> usize { dict_size.wrapping_add(mask) & !mask } +type EstimateCDictSizeAdvancedFn = unsafe extern "C" fn( + *mut c_void, + usize, + ZSTD_compressionParameters, + c_int, +) -> usize; + +/// Projection for the scalar policy wrapper of `ZSTD_estimateCDictSize()`. +/// +/// Rust selects the unknown-source/create-CDict parameters and the public +/// by-copy construction mode. C keeps the advanced estimator's private +/// layout-size inputs and workspace arithmetic behind the opaque callback. +#[repr(C)] +pub struct ZSTD_rust_estimateCDictSizeState { + callback_context: *mut c_void, + exclusion_mask: *const c_uint, + estimate_advanced: EstimateCDictSizeAdvancedFn, +} + +const _: () = { + assert!(offset_of!(ZSTD_rust_estimateCDictSizeState, callback_context) == 0); + assert!(offset_of!(ZSTD_rust_estimateCDictSizeState, exclusion_mask) == size_of::()); + assert!( + offset_of!(ZSTD_rust_estimateCDictSizeState, estimate_advanced) + == 2 * size_of::() + ); + assert!(size_of::() == size_of::<[usize; 3]>()); +}; + +/// Select public CDict-size parameters before entering the C-owned estimator. +#[no_mangle] +pub unsafe extern "C" fn ZSTD_rust_estimateCDictSize( + state: *const ZSTD_rust_estimateCDictSizeState, + dict_size: usize, + compression_level: c_int, +) -> usize { + if state.is_null() { + return 0; + } + let state = unsafe { &*state }; + if state.exclusion_mask.is_null() { + return 0; + } + let cparams = ZSTD_rust_params_getCParamsInternal( + compression_level, + ZSTD_CONTENTSIZE_UNKNOWN, + dict_size, + ZSTD_RUST_CPM_CREATE_CDICT, + unsafe { *state.exclusion_mask }, + ); + unsafe { + (state.estimate_advanced)( + state.callback_context, + dict_size, + cparams, + ZSTD_DLM_BY_COPY, + ) + } +} + /// C keeps match-state/content insertion private because it depends on the /// configuration-sensitive `ZSTD_MatchState_t`, `ldmState_t`, workspace, and /// parameter layouts. Rust owns the dictionary dispatch and calls this narrow @@ -419,7 +480,6 @@ fn dictionary_corrupted() -> usize { const ZSTD_CCTX_INIT_STAGE: c_int = 0; const ZSTD_DLM_BY_REF: c_int = 1; -#[cfg(test)] const ZSTD_DLM_BY_COPY: c_int = 0; #[inline] @@ -2576,6 +2636,89 @@ mod tests { use crate::mem::MEM_writeLE32; use std::mem::{size_of, MaybeUninit}; + #[derive(Default)] + struct EstimateCDictSizeProbe { + events: Vec<&'static str>, + callback_context: *mut c_void, + dict_size: usize, + cparams: ZSTD_compressionParameters, + dict_load_method: c_int, + result: usize, + } + + unsafe extern "C" fn estimate_cdict_size_test_advanced( + context: *mut c_void, + dict_size: usize, + cparams: ZSTD_compressionParameters, + dict_load_method: c_int, + ) -> usize { + let probe = unsafe { &mut *context.cast::() }; + probe.events.push("advanced"); + probe.callback_context = context; + probe.dict_size = dict_size; + probe.cparams = cparams; + probe.dict_load_method = dict_load_method; + probe.result + } + + fn estimate_cdict_size_test_state( + probe: &mut EstimateCDictSizeProbe, + exclusion_mask: &c_uint, + ) -> ZSTD_rust_estimateCDictSizeState { + ZSTD_rust_estimateCDictSizeState { + callback_context: (probe as *mut EstimateCDictSizeProbe).cast(), + exclusion_mask, + estimate_advanced: estimate_cdict_size_test_advanced, + } + } + + #[test] + fn estimate_cdict_size_selects_create_policy_before_advanced_callback() { + let mut probe = EstimateCDictSizeProbe { + result: 97, + ..Default::default() + }; + let dict_size = 4096; + let compression_level = 7; + let exclusion_mask = 0; + let state = estimate_cdict_size_test_state(&mut probe, &exclusion_mask); + + let result = unsafe { + ZSTD_rust_estimateCDictSize(&state, dict_size, compression_level) + }; + + assert_eq!(result, probe.result); + assert_eq!(probe.events, ["advanced"]); + assert_eq!(probe.callback_context, state.callback_context); + assert_eq!(probe.dict_size, dict_size); + assert_eq!( + probe.cparams, + ZSTD_rust_params_getCParamsInternal( + compression_level, + ZSTD_CONTENTSIZE_UNKNOWN, + dict_size, + ZSTD_RUST_CPM_CREATE_CDICT, + exclusion_mask, + ) + ); + assert_eq!(probe.dict_load_method, ZSTD_DLM_BY_COPY); + } + + #[test] + fn estimate_cdict_size_rejects_missing_policy_inputs_before_callback() { + let mut probe = EstimateCDictSizeProbe::default(); + let state = ZSTD_rust_estimateCDictSizeState { + callback_context: (&mut probe as *mut EstimateCDictSizeProbe).cast(), + exclusion_mask: ptr::null(), + estimate_advanced: estimate_cdict_size_test_advanced, + }; + + let result = unsafe { ZSTD_rust_estimateCDictSize(&state, 1, 3) }; + + assert_eq!(result, 0); + assert!(probe.events.is_empty()); + } + const DICT_CONTENT_SIZE: usize = 16; #[derive(Clone, Copy)]