diff --git a/lib/compress/zstd_compress.c b/lib/compress/zstd_compress.c index b2f025dc6..537c60b4f 100644 --- a/lib/compress/zstd_compress.c +++ b/lib/compress/zstd_compress.c @@ -422,6 +422,31 @@ typedef char ZSTD_rust_compress_using_dict_state_layout[ == 3 * sizeof(void*) && sizeof(ZSTD_rust_compressUsingDictState) == 4 * sizeof(void*)) ? 1 : -1]; +typedef void (*ZSTD_rust_compressBeginAdvancedInitParams_f)( + void* cctxParams, const ZSTD_parameters* params); +typedef size_t (*ZSTD_rust_compressBeginAdvancedBegin_f)( + void* context, const void* dict, size_t dictSize, + const void* cctxParams, U64 pledgedSrcSize); +typedef struct { + void* cctx; + void* cctxParams; + ZSTD_rust_compressBeginAdvancedInitParams_f initParams; + ZSTD_rust_compressBeginAdvancedBegin_f begin; +} ZSTD_rust_compressBeginAdvancedState; +size_t ZSTD_rust_compressBeginAdvanced( + const ZSTD_rust_compressBeginAdvancedState* state, + const void* dict, size_t dictSize, + const ZSTD_parameters* params, U64 pledgedSrcSize); +typedef char ZSTD_rust_compress_begin_advanced_state_layout[ + (offsetof(ZSTD_rust_compressBeginAdvancedState, cctx) == 0 + && offsetof(ZSTD_rust_compressBeginAdvancedState, cctxParams) + == sizeof(void*) + && offsetof(ZSTD_rust_compressBeginAdvancedState, initParams) + == 2 * sizeof(void*) + && offsetof(ZSTD_rust_compressBeginAdvancedState, begin) + == 3 * sizeof(void*) + && sizeof(ZSTD_rust_compressBeginAdvancedState) == 4 * sizeof(void*)) + ? 1 : -1]; ZSTD_frameProgression ZSTD_rust_frameProgression(U64 consumedSrcSize, size_t buffered, U64 producedCSize); @@ -4841,6 +4866,24 @@ static size_t ZSTD_rust_compressBeginUsingDict_begin( pledgedSrcSize, ZSTDb_not_buffered); } +static void ZSTD_rust_compressBeginAdvanced_initParams( + void* cctxParams, const ZSTD_parameters* params) +{ + ZSTD_CCtxParams_init_internal( + (ZSTD_CCtx_params*)cctxParams, params, ZSTD_NO_CLEVEL); +} + +static size_t ZSTD_rust_compressBeginAdvanced_begin( + void* context, const void* dict, size_t dictSize, + const void* cctxParams, U64 pledgedSrcSize) +{ + return ZSTD_compressBegin_internal( + (ZSTD_CCtx*)context, dict, dictSize, + ZSTD_dct_auto, ZSTD_dtlm_fast, NULL, + (const ZSTD_CCtx_params*)cctxParams, + pledgedSrcSize, ZSTDb_not_buffered); +} + size_t ZSTD_compressBegin_advanced_internal(ZSTD_CCtx* cctx, const void* dict, size_t dictSize, ZSTD_dictContentType_e dictContentType, @@ -4865,12 +4908,14 @@ size_t ZSTD_compressBegin_advanced(ZSTD_CCtx* cctx, const void* dict, size_t dictSize, ZSTD_parameters params, unsigned long long pledgedSrcSize) { + ZSTD_rust_compressBeginAdvancedState state; ZSTD_CCtx_params cctxParams; - ZSTD_CCtxParams_init_internal(&cctxParams, ¶ms, ZSTD_NO_CLEVEL); - return ZSTD_compressBegin_advanced_internal(cctx, - dict, dictSize, ZSTD_dct_auto, ZSTD_dtlm_fast, - NULL /*cdict*/, - &cctxParams, pledgedSrcSize); + state.cctx = cctx; + state.cctxParams = &cctxParams; + state.initParams = ZSTD_rust_compressBeginAdvanced_initParams; + state.begin = ZSTD_rust_compressBeginAdvanced_begin; + return ZSTD_rust_compressBeginAdvanced( + &state, dict, dictSize, ¶ms, pledgedSrcSize); } static size_t diff --git a/rust/src/zstd_compress.rs b/rust/src/zstd_compress.rs index 714ea6bb3..4371d0fda 100644 --- a/rust/src/zstd_compress.rs +++ b/rust/src/zstd_compress.rs @@ -1547,6 +1547,66 @@ pub unsafe extern "C" fn ZSTD_rust_compressAdvanced( } } +type CompressBeginAdvancedInitParamsFn = unsafe extern "C" fn(*mut c_void, *const ZSTD_parameters); +type CompressBeginAdvancedBeginFn = + unsafe extern "C" fn(*mut c_void, *const c_void, usize, *const c_void, u64) -> usize; + +/// Explicit projection for the public `ZSTD_compressBegin_advanced` wrapper. +/// +/// Rust owns public parameter validation and init-then-begin ordering. C +/// retains private `ZSTD_CCtx_params` initialization and the final begin +/// operation behind callbacks. +#[repr(C)] +pub struct ZSTD_rust_compressBeginAdvancedState { + cctx: *mut c_void, + cctx_params: *mut c_void, + init_params: CompressBeginAdvancedInitParamsFn, + begin: CompressBeginAdvancedBeginFn, +} + +const _: () = { + assert!(offset_of!(ZSTD_rust_compressBeginAdvancedState, cctx) == 0); + assert!(offset_of!(ZSTD_rust_compressBeginAdvancedState, cctx_params) == size_of::()); + assert!( + offset_of!(ZSTD_rust_compressBeginAdvancedState, init_params) == 2 * size_of::() + ); + assert!(offset_of!(ZSTD_rust_compressBeginAdvancedState, begin) == 3 * size_of::()); + assert!(size_of::() == size_of::<[usize; 4]>()); +}; + +/// Validate parameters and begin a frame through the C-owned operation. +#[no_mangle] +pub unsafe extern "C" fn ZSTD_rust_compressBeginAdvanced( + state: *const ZSTD_rust_compressBeginAdvancedState, + dict: *const c_void, + dict_size: usize, + params: *const ZSTD_parameters, + pledged_src_size: u64, +) -> usize { + if state.is_null() || params.is_null() { + return ERROR(ZstdErrorCode::Generic); + } + let state = unsafe { &*state }; + if state.cctx.is_null() || state.cctx_params.is_null() { + return ERROR(ZstdErrorCode::Generic); + } + let params = unsafe { &*params }; + let check_result = ZSTD_rust_params_checkCParams(params.cParams); + if ERR_isError(check_result) { + return check_result; + } + unsafe { + (state.init_params)(state.cctx_params, params); + (state.begin)( + state.cctx, + dict, + dict_size, + state.cctx_params, + pledged_src_size, + ) + } +} + type CompressUsingDictInitParamsFn = unsafe extern "C" fn(*mut c_void, *const ZSTD_parameters, c_int); @@ -10657,6 +10717,92 @@ mod tests { assert!(context.events.is_empty()); } + #[derive(Default)] + struct CompressBeginAdvancedTestContext { + events: Vec<&'static str>, + params: ZSTD_parameters, + dict: *const c_void, + dict_size: usize, + cctx_params: *const c_void, + pledged_src_size: u64, + result: usize, + } + + unsafe extern "C" fn compress_begin_advanced_test_init( + context: *mut c_void, + params: *const ZSTD_parameters, + ) { + let context = unsafe { &mut *context.cast::() }; + context.events.push("init"); + context.params = unsafe { *params }; + } + + unsafe extern "C" fn compress_begin_advanced_test_begin( + context: *mut c_void, + dict: *const c_void, + dict_size: usize, + cctx_params: *const c_void, + pledged_src_size: u64, + ) -> usize { + let context = unsafe { &mut *context.cast::() }; + context.events.push("begin"); + context.dict = dict; + context.dict_size = dict_size; + context.cctx_params = cctx_params; + context.pledged_src_size = pledged_src_size; + context.result + } + + fn compress_begin_advanced_test_state( + context: &mut CompressBeginAdvancedTestContext, + ) -> ZSTD_rust_compressBeginAdvancedState { + ZSTD_rust_compressBeginAdvancedState { + cctx: (context as *mut CompressBeginAdvancedTestContext).cast(), + cctx_params: (context as *mut CompressBeginAdvancedTestContext).cast(), + init_params: compress_begin_advanced_test_init, + begin: compress_begin_advanced_test_begin, + } + } + + #[test] + fn compress_begin_advanced_validates_and_preserves_init_begin_order() { + let mut context = CompressBeginAdvancedTestContext { + result: 29, + ..CompressBeginAdvancedTestContext::default() + }; + let state = compress_begin_advanced_test_state(&mut context); + let params = compress_advanced_test_params(); + let dict = [1u8, 2, 3]; + + let result = unsafe { + ZSTD_rust_compressBeginAdvanced(&state, dict.as_ptr().cast(), dict.len(), ¶ms, 123) + }; + + assert_eq!(result, context.result); + assert_eq!(context.events, ["init", "begin"]); + assert_eq!(context.params, params); + assert_eq!(context.dict, dict.as_ptr().cast()); + assert_eq!(context.dict_size, dict.len()); + assert_eq!( + context.cctx_params, + (&context as *const CompressBeginAdvancedTestContext).cast() + ); + assert_eq!(context.pledged_src_size, 123); + } + + #[test] + fn compress_begin_advanced_stops_before_callbacks_on_invalid_parameters() { + let mut context = CompressBeginAdvancedTestContext::default(); + let state = compress_begin_advanced_test_state(&mut context); + let mut params = compress_advanced_test_params(); + params.cParams.windowLog = 0; + + let result = unsafe { ZSTD_rust_compressBeginAdvanced(&state, ptr::null(), 0, ¶ms, 0) }; + + assert_eq!(result, ERROR(ZstdErrorCode::ParameterOutOfBound)); + assert!(context.events.is_empty()); + } + #[derive(Default)] struct CompressUsingDictTestContext { events: Vec<&'static str>,