diff --git a/lib/compress/zstd_compress.c b/lib/compress/zstd_compress.c index 044ab6d96..452065629 100644 --- a/lib/compress/zstd_compress.c +++ b/lib/compress/zstd_compress.c @@ -1779,14 +1779,24 @@ typedef char ZSTD_rust_init_cdict_state_layout[ && sizeof(ZSTD_rust_initCDictState) == 19 * sizeof(void*)) ? 1 : -1]; -typedef void* (*ZSTD_rust_createCDictAdvancedCreate_f)( +typedef int (*ZSTD_rust_createCDictAdvancedValidateCustomMem_f)(void* context); +typedef size_t (*ZSTD_rust_createCDictAdvancedWorkspaceSize_f)( void* context, size_t dictSize, int dictLoadMethod, const ZSTD_compressionParameters* cParams, int useRowMatchFinder, int enableDedicatedDictSearch); +typedef void* (*ZSTD_rust_createCDictAdvancedAllocate_f)( + void* context, size_t workspaceSize); +typedef void* (*ZSTD_rust_createCDictAdvancedCreate_f)( + void* context, void* workspace, size_t workspaceSize, + 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_createCDictAdvancedFreeWorkspace_f)( + void* context, void* workspace); typedef void (*ZSTD_rust_createCDictAdvancedFree_f)( void* context, void* cdict); typedef struct { @@ -1797,8 +1807,12 @@ typedef struct { const int* useRowMatchFinder; U32 exclusionMask; U32 ldmDefaultWindowLog; + ZSTD_rust_createCDictAdvancedValidateCustomMem_f validateCustomMem; + ZSTD_rust_createCDictAdvancedWorkspaceSize_f workspaceSize; + ZSTD_rust_createCDictAdvancedAllocate_f allocate; ZSTD_rust_createCDictAdvancedCreate_f create; ZSTD_rust_createCDictAdvancedInit_f init; + ZSTD_rust_createCDictAdvancedFreeWorkspace_f freeWorkspace; ZSTD_rust_createCDictAdvancedFree_f free; } ZSTD_rust_createCDictAdvancedState; void* ZSTD_rust_createCDictAdvanced( @@ -1819,14 +1833,22 @@ typedef char ZSTD_rust_create_cdict_advanced_state_layout[ == 5 * sizeof(void*) && offsetof(ZSTD_rust_createCDictAdvancedState, ldmDefaultWindowLog) == 5 * sizeof(void*) + sizeof(U32) - && offsetof(ZSTD_rust_createCDictAdvancedState, create) + && offsetof(ZSTD_rust_createCDictAdvancedState, validateCustomMem) == 5 * sizeof(void*) + 2 * sizeof(U32) - && offsetof(ZSTD_rust_createCDictAdvancedState, init) + && offsetof(ZSTD_rust_createCDictAdvancedState, workspaceSize) == 6 * sizeof(void*) + 2 * sizeof(U32) - && offsetof(ZSTD_rust_createCDictAdvancedState, free) + && offsetof(ZSTD_rust_createCDictAdvancedState, allocate) == 7 * sizeof(void*) + 2 * sizeof(U32) + && offsetof(ZSTD_rust_createCDictAdvancedState, create) + == 8 * sizeof(void*) + 2 * sizeof(U32) + && offsetof(ZSTD_rust_createCDictAdvancedState, init) + == 9 * sizeof(void*) + 2 * sizeof(U32) + && offsetof(ZSTD_rust_createCDictAdvancedState, freeWorkspace) + == 10 * sizeof(void*) + 2 * sizeof(U32) + && offsetof(ZSTD_rust_createCDictAdvancedState, free) + == 11 * sizeof(void*) + 2 * sizeof(U32) && sizeof(ZSTD_rust_createCDictAdvancedState) - == 8 * sizeof(void*) + 2 * sizeof(U32)) + == 12 * sizeof(void*) + 2 * sizeof(U32)) ? 1 : -1]; typedef size_t (*ZSTD_rust_compressBeginResetInternal_f)( @@ -5544,42 +5566,67 @@ static size_t ZSTD_initCDict_internal( (int)dictLoadMethod, (int)dictContentType); } -static ZSTD_CDict* -ZSTD_createCDict_advanced_internal(size_t dictSize, - ZSTD_dictLoadMethod_e dictLoadMethod, - ZSTD_compressionParameters cParams, - ZSTD_ParamSwitch_e useRowMatchFinder, - int enableDedicatedDictSearch, - ZSTD_customMem customMem) +static int ZSTD_rust_createCDictAdvanced_validateCustomMem(void* context) { - if ((!customMem.customAlloc) ^ (!customMem.customFree)) return NULL; - DEBUGLOG(3, "ZSTD_createCDict_advanced_internal (dictSize=%u)", (unsigned)dictSize); + ZSTD_customMem const* const customMem = (const ZSTD_customMem*)context; + return ((!customMem->customAlloc) ^ (!customMem->customFree)) == 0; +} - { size_t const workspaceSize = - ZSTD_cwksp_alloc_size(sizeof(ZSTD_CDict)) + - ZSTD_cwksp_alloc_size(HUF_WORKSPACE_SIZE) + - ZSTD_sizeof_matchState(&cParams, useRowMatchFinder, enableDedicatedDictSearch, /* forCCtx */ 0) + - (dictLoadMethod == ZSTD_dlm_byRef ? 0 - : ZSTD_cwksp_alloc_size(ZSTD_cwksp_align(dictSize, sizeof(void*)))); - void* const workspace = ZSTD_customMalloc(workspaceSize, customMem); - ZSTD_cwksp ws; - ZSTD_CDict* cdict; +static size_t ZSTD_rust_createCDictAdvanced_workspaceSize( + void* context, size_t dictSize, int dictLoadMethod, + const ZSTD_compressionParameters* cParams, + int useRowMatchFinder, int enableDedicatedDictSearch) +{ + (void)context; + if (cParams == NULL) return 0; + return ZSTD_cwksp_alloc_size(sizeof(ZSTD_CDict)) + + ZSTD_cwksp_alloc_size(HUF_WORKSPACE_SIZE) + + ZSTD_sizeof_matchState(cParams, + (ZSTD_ParamSwitch_e)useRowMatchFinder, + enableDedicatedDictSearch, + /* forCCtx */ 0) + + (dictLoadMethod == ZSTD_dlm_byRef ? 0 + : ZSTD_cwksp_alloc_size(ZSTD_cwksp_align(dictSize, sizeof(void*)))); +} - if (!workspace) { - ZSTD_customFree(workspace, customMem); - return NULL; - } +static void* ZSTD_rust_createCDictAdvanced_allocate( + void* context, size_t workspaceSize) +{ + return ZSTD_customMalloc(workspaceSize, *(const ZSTD_customMem*)context); +} - ZSTD_cwksp_init(&ws, workspace, workspaceSize, ZSTD_cwksp_dynamic_alloc); +static void* ZSTD_rust_createCDictAdvanced_create( + void* context, void* workspace, size_t workspaceSize, + size_t dictSize, int dictLoadMethod, + const ZSTD_compressionParameters* cParams, + int useRowMatchFinder, int enableDedicatedDictSearch) +{ + ZSTD_customMem const customMem = *(const ZSTD_customMem*)context; + ZSTD_cwksp ws; + ZSTD_CDict* cdict; - cdict = (ZSTD_CDict*)ZSTD_cwksp_reserve_object(&ws, sizeof(ZSTD_CDict)); - assert(cdict != NULL); - ZSTD_cwksp_move(&cdict->workspace, &ws); - cdict->customMem = customMem; - cdict->compressionLevel = ZSTD_NO_CLEVEL; /* signals advanced API usage */ - cdict->useRowMatchFinder = useRowMatchFinder; - return cdict; - } + (void)dictSize; + (void)dictLoadMethod; + (void)cParams; + (void)enableDedicatedDictSearch; + DEBUGLOG(3, "ZSTD_createCDict_advanced_internal (workspaceSize=%u)", + (unsigned)workspaceSize); + if (workspace == NULL) return NULL; + + ZSTD_cwksp_init(&ws, workspace, workspaceSize, ZSTD_cwksp_dynamic_alloc); + cdict = (ZSTD_CDict*)ZSTD_cwksp_reserve_object(&ws, sizeof(ZSTD_CDict)); + if (cdict == NULL) return NULL; + ZSTD_cwksp_move(&cdict->workspace, &ws); + cdict->customMem = customMem; + cdict->compressionLevel = ZSTD_NO_CLEVEL; /* signals advanced API usage */ + cdict->useRowMatchFinder = (ZSTD_ParamSwitch_e)useRowMatchFinder; + return cdict; +} + +static void ZSTD_rust_createCDictAdvanced_freeWorkspace( + void* context, void* workspace) +{ + ZSTD_customFree(workspace, *(const ZSTD_customMem*)context); } ZSTD_CDict* ZSTD_createCDict_advanced(const void* dictBuffer, size_t dictSize, @@ -5600,18 +5647,6 @@ 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, @@ -5642,7 +5677,6 @@ ZSTD_CDict* ZSTD_createCDict_advanced2( 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; cctxParams = *originalCctxParams; state.callbackContext = &customMem; @@ -5652,8 +5686,12 @@ ZSTD_CDict* ZSTD_createCDict_advanced2( state.useRowMatchFinder = (const int*)&cctxParams.useRowMatchFinder; state.exclusionMask = ZSTD_getCParamsExclusionMask(); state.ldmDefaultWindowLog = ZSTD_LDM_DEFAULT_WINDOW_LOG; + state.validateCustomMem = ZSTD_rust_createCDictAdvanced_validateCustomMem; + state.workspaceSize = ZSTD_rust_createCDictAdvanced_workspaceSize; + state.allocate = ZSTD_rust_createCDictAdvanced_allocate; state.create = ZSTD_rust_createCDictAdvanced_create; state.init = ZSTD_rust_createCDictAdvanced_init; + state.freeWorkspace = ZSTD_rust_createCDictAdvanced_freeWorkspace; state.free = ZSTD_rust_createCDictAdvanced_free; return (ZSTD_CDict*)ZSTD_rust_createCDictAdvanced( &state, dict, dictSize, diff --git a/rust/README.md b/rust/README.md index c6cbc571e..a2039f2dd 100644 --- a/rust/README.md +++ b/rust/README.md @@ -162,10 +162,12 @@ sequencing, legacy public CDict-begin frame-policy construction and unknown-source pledge, and the compressBegin_usingDict family’s unknown-source 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 +CDict advanced private workspace construction, private static-CCtx and static-CDict workspace construction and dictionary-content allocation/loading, -and advanced-CDict private workspace construction and dictionary-content -loading remain in C. Reset policy, +and advanced-CDict dictionary-content loading remain in C. Rust now owns +advanced-CDict custom-memory validation, workspace-size query/allocation, and +allocation/create/init cleanup ordering; C retains the private workspace +size formula, layout, and allocator callbacks. 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 8c157ef99..f0025259e 100644 --- a/rust/src/zstd_compress_dictionary.rs +++ b/rust/src/zstd_compress_dictionary.rs @@ -500,8 +500,20 @@ pub unsafe extern "C" fn ZSTD_rust_createCDict( cdict } +type CreateCDictAdvancedValidateCustomMemFn = unsafe extern "C" fn(*mut c_void) -> c_int; +type CreateCDictAdvancedWorkspaceSizeFn = unsafe extern "C" fn( + *mut c_void, + usize, + c_int, + *const ZSTD_compressionParameters, + c_int, + c_int, +) -> usize; +type CreateCDictAdvancedAllocateFn = unsafe extern "C" fn(*mut c_void, usize) -> *mut c_void; type CreateCDictAdvancedCreateFn = unsafe extern "C" fn( *mut c_void, + *mut c_void, + usize, usize, c_int, *const ZSTD_compressionParameters, @@ -518,14 +530,16 @@ type CreateCDictAdvancedInitFn = unsafe extern "C" fn( *const ZSTD_CCtx_params, ) -> usize; type CreateCDictAdvancedFreeFn = unsafe extern "C" fn(*mut c_void, *mut c_void); +type CreateCDictAdvancedFreeWorkspaceFn = 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. +/// Rust owns context-free parameter preparation, custom-memory validation, +/// allocation, and callback ordering. C retains 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, @@ -535,8 +549,12 @@ pub struct ZSTD_rust_createCDictAdvancedState { use_row_match_finder: *const c_int, exclusion_mask: u32, ldm_default_window_log: u32, + validate_custom_mem: CreateCDictAdvancedValidateCustomMemFn, + workspace_size: CreateCDictAdvancedWorkspaceSizeFn, + allocate: CreateCDictAdvancedAllocateFn, create: CreateCDictAdvancedCreateFn, init: CreateCDictAdvancedInitFn, + free_workspace: CreateCDictAdvancedFreeWorkspaceFn, free: CreateCDictAdvancedFreeFn, } @@ -562,21 +580,37 @@ const _: () = { == 5 * size_of::() + size_of::() ); assert!( - offset_of!(ZSTD_rust_createCDictAdvancedState, create) + offset_of!(ZSTD_rust_createCDictAdvancedState, validate_custom_mem) == size_of::<[usize; 5]>() + size_of::<[u32; 2]>() ); assert!( - offset_of!(ZSTD_rust_createCDictAdvancedState, init) + offset_of!(ZSTD_rust_createCDictAdvancedState, workspace_size) == size_of::<[usize; 6]>() + size_of::<[u32; 2]>() ); assert!( - offset_of!(ZSTD_rust_createCDictAdvancedState, free) + offset_of!(ZSTD_rust_createCDictAdvancedState, allocate) == size_of::<[usize; 7]>() + size_of::<[u32; 2]>() ); assert!( - size_of::() + offset_of!(ZSTD_rust_createCDictAdvancedState, create) == size_of::<[usize; 8]>() + size_of::<[u32; 2]>() ); + assert!( + offset_of!(ZSTD_rust_createCDictAdvancedState, init) + == size_of::<[usize; 9]>() + size_of::<[u32; 2]>() + ); + assert!( + offset_of!(ZSTD_rust_createCDictAdvancedState, free_workspace) + == size_of::<[usize; 10]>() + size_of::<[u32; 2]>() + ); + assert!( + offset_of!(ZSTD_rust_createCDictAdvancedState, free) + == size_of::<[usize; 11]>() + size_of::<[u32; 2]>() + ); + assert!( + size_of::() + == size_of::<[usize; 12]>() + size_of::<[u32; 2]>() + ); }; /// Prepare advanced-CDict parameters and run the C-owned construction path. @@ -613,8 +647,12 @@ pub unsafe extern "C" fn ZSTD_rust_createCDictAdvanced( return ptr::null_mut(); } - let cdict = unsafe { - (state.create)( + if unsafe { (state.validate_custom_mem)(state.callback_context) == 0 } { + return ptr::null_mut(); + } + + let workspace_size = unsafe { + (state.workspace_size)( state.callback_context, dict_size, dict_load_method, @@ -623,7 +661,25 @@ pub unsafe extern "C" fn ZSTD_rust_createCDictAdvanced( *state.enable_dedicated_dict_search, ) }; + let workspace = unsafe { (state.allocate)(state.callback_context, workspace_size) }; + if workspace.is_null() { + return ptr::null_mut(); + } + + let cdict = unsafe { + (state.create)( + state.callback_context, + workspace, + workspace_size, + dict_size, + dict_load_method, + state.cparams, + *state.use_row_match_finder, + *state.enable_dedicated_dict_search, + ) + }; if cdict.is_null() { + unsafe { (state.free_workspace)(state.callback_context, workspace) }; return ptr::null_mut(); } @@ -3203,6 +3259,19 @@ mod tests { #[derive(Default)] struct CreateCDictAdvancedProbe { events: Vec<&'static str>, + custom_mem_valid: c_int, + workspace_size_result: usize, + workspace_size_dict_size: usize, + workspace_size_dict_load_method: c_int, + workspace_size_cparams: *const ZSTD_compressionParameters, + workspace_size_use_row_match_finder: c_int, + workspace_size_enable_dedicated_dict_search: c_int, + allocated_workspace: *mut c_void, + allocated_workspace_size: usize, + create_workspace: *mut c_void, + create_workspace_size: usize, + create_dict_size: usize, + create_dict_load_method: c_int, cparams: *const ZSTD_compressionParameters, use_row_match_finder: c_int, enable_dedicated_dict_search: c_int, @@ -3213,20 +3282,63 @@ mod tests { init_dict_load_method: c_int, init_dict_content_type: c_int, init_cctx_params: *const ZSTD_CCtx_params, + free_workspace: *mut c_void, free_cdict: *mut c_void, init_result: usize, } + unsafe extern "C" fn create_cdict_advanced_test_validate_custom_mem( + context: *mut c_void, + ) -> c_int { + let probe = unsafe { &mut *context.cast::() }; + probe.events.push("validate"); + probe.custom_mem_valid + } + + unsafe extern "C" fn create_cdict_advanced_test_workspace_size( + 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, + ) -> usize { + let probe = unsafe { &mut *context.cast::() }; + probe.events.push("workspace_size"); + probe.workspace_size_dict_size = dict_size; + probe.workspace_size_dict_load_method = dict_load_method; + probe.workspace_size_cparams = cparams; + probe.workspace_size_use_row_match_finder = use_row_match_finder; + probe.workspace_size_enable_dedicated_dict_search = enable_dedicated_dict_search; + probe.workspace_size_result + } + + unsafe extern "C" fn create_cdict_advanced_test_allocate( + context: *mut c_void, + workspace_size: usize, + ) -> *mut c_void { + let probe = unsafe { &mut *context.cast::() }; + probe.events.push("allocate"); + probe.allocated_workspace_size = workspace_size; + probe.allocated_workspace + } + unsafe extern "C" fn create_cdict_advanced_test_create( context: *mut c_void, - _dict_size: usize, - _dict_load_method: c_int, + workspace: *mut c_void, + workspace_size: usize, + 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.create_workspace = workspace; + probe.create_workspace_size = workspace_size; + probe.create_dict_size = dict_size; + probe.create_dict_load_method = dict_load_method; probe.cparams = cparams; probe.use_row_match_finder = use_row_match_finder; probe.enable_dedicated_dict_search = enable_dedicated_dict_search; @@ -3253,6 +3365,15 @@ mod tests { probe.init_result } + unsafe extern "C" fn create_cdict_advanced_test_free_workspace( + context: *mut c_void, + workspace: *mut c_void, + ) { + let probe = unsafe { &mut *context.cast::() }; + probe.events.push("free_workspace"); + probe.free_workspace = workspace; + } + 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"); @@ -3288,8 +3409,12 @@ mod tests { use_row_match_finder, exclusion_mask: 0, ldm_default_window_log: 27, + validate_custom_mem: create_cdict_advanced_test_validate_custom_mem, + workspace_size: create_cdict_advanced_test_workspace_size, + allocate: create_cdict_advanced_test_allocate, create: create_cdict_advanced_test_create, init: create_cdict_advanced_test_init, + free_workspace: create_cdict_advanced_test_free_workspace, free: create_cdict_advanced_test_free, } } @@ -3297,6 +3422,9 @@ mod tests { #[test] fn create_cdict_advanced_publishes_params_and_preserves_create_init_order() { let mut probe = CreateCDictAdvancedProbe { + custom_mem_valid: 1, + workspace_size_result: 123, + allocated_workspace: ptr::dangling_mut(), cdict_result: ptr::dangling_mut(), ..Default::default() }; @@ -3333,7 +3461,26 @@ mod tests { }; assert_eq!(result, probe.cdict_result); - assert_eq!(probe.events, ["create", "init"]); + assert_eq!( + probe.events, + ["validate", "workspace_size", "allocate", "create", "init"] + ); + assert_eq!(probe.workspace_size_dict_size, dict.len()); + assert_eq!(probe.workspace_size_dict_load_method, ZSTD_DLM_BY_REF); + assert_eq!(probe.workspace_size_cparams, &cparams); + assert_eq!( + probe.workspace_size_use_row_match_finder, + use_row_match_finder + ); + assert_eq!( + probe.workspace_size_enable_dedicated_dict_search, + enable_dedicated_dict_search + ); + assert_eq!(probe.allocated_workspace_size, 123); + assert_eq!(probe.create_workspace, probe.allocated_workspace); + assert_eq!(probe.create_workspace_size, 123); + assert_eq!(probe.create_dict_size, dict.len()); + assert_eq!(probe.create_dict_load_method, ZSTD_DLM_BY_REF); assert_eq!(probe.cparams, &cparams); assert_eq!(probe.use_row_match_finder, use_row_match_finder); assert_eq!( @@ -3346,10 +3493,12 @@ mod tests { 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()); + assert!(probe.free_workspace.is_null()); + assert!(probe.free_cdict.is_null()); } #[test] - fn create_cdict_advanced_does_not_init_or_free_after_creation_failure() { + fn create_cdict_advanced_rejects_invalid_custom_memory_before_allocation() { let mut probe = CreateCDictAdvancedProbe::default(); let mut params_storage = create_cdict_advanced_test_params(); let cctx_params = params_storage.as_mut_ptr(); @@ -3369,14 +3518,96 @@ mod tests { }; assert!(result.is_null()); - assert_eq!(probe.events, ["create"]); + assert_eq!(probe.events, ["validate"]); + assert!(probe.create_workspace.is_null()); + assert!(probe.free_workspace.is_null()); + assert!(probe.free_cdict.is_null()); + } + + #[test] + fn create_cdict_advanced_stops_after_allocation_failure() { + let mut probe = CreateCDictAdvancedProbe { + custom_mem_valid: 1, + workspace_size_result: 321, + ..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, ["validate", "workspace_size", "allocate"]); + assert_eq!(probe.allocated_workspace_size, 321); + assert!(probe.create_workspace.is_null()); + assert!(probe.free_workspace.is_null()); + assert!(probe.free_cdict.is_null()); + } + + #[test] + fn create_cdict_advanced_frees_workspace_after_creation_failure() { + let workspace = ptr::dangling_mut::(); + let mut probe = CreateCDictAdvancedProbe { + custom_mem_valid: 1, + workspace_size_result: 321, + allocated_workspace: workspace, + ..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, + [ + "validate", + "workspace_size", + "allocate", + "create", + "free_workspace" + ] + ); + assert_eq!(probe.create_workspace, workspace); + assert_eq!(probe.create_workspace_size, 321); + assert_eq!(probe.create_dict_size, 0); + assert_eq!(probe.free_workspace, workspace); assert!(probe.free_cdict.is_null()); } #[test] fn create_cdict_advanced_frees_after_initialization_failure() { let cdict = ptr::dangling_mut::(); + let workspace = ptr::dangling_mut::(); let mut probe = CreateCDictAdvancedProbe { + custom_mem_valid: 1, + workspace_size_result: 321, + allocated_workspace: workspace, cdict_result: cdict, init_result: ERROR(ZstdErrorCode::MemoryAllocation), ..Default::default() @@ -3399,7 +3630,21 @@ mod tests { }; assert!(result.is_null()); - assert_eq!(probe.events, ["create", "init", "free"]); + assert_eq!( + probe.events, + [ + "validate", + "workspace_size", + "allocate", + "create", + "init", + "free" + ] + ); + assert_eq!(probe.create_workspace, workspace); + assert_eq!(probe.create_workspace_size, 321); + assert_eq!(probe.create_dict_size, 0); + assert!(probe.free_workspace.is_null()); assert_eq!(probe.free_cdict, cdict); }