From 50e7addbc616d7149c89d4538057f954ee630b39 Mon Sep 17 00:00:00 2001 From: ddidderr Date: Sun, 19 Jul 2026 15:35:14 +0200 Subject: [PATCH] feat(compress): move advanced API policy into Rust Move ZSTD_compress_advanced parameter validation and init-before-compress orchestration into Rust. Keep simpleApiParams initialization and the private advanced compression operation behind C callbacks. Test Plan: - cargo test --manifest-path rust/Cargo.toml --lib - cargo clippy --manifest-path rust/Cargo.toml --all-targets -- -D warnings - make -B -C programs -j1 zstd - make -C tests -j1 test-zstream ZSTREAM_TESTTIME=-T1s - focused compress_advanced unit tests --- lib/compress/zstd_compress.c | 65 +++++++++-- rust/src/zstd_compress.rs | 214 ++++++++++++++++++++++++++++++++++- 2 files changed, 269 insertions(+), 10 deletions(-) diff --git a/lib/compress/zstd_compress.c b/lib/compress/zstd_compress.c index cd33bd24a..f38918b07 100644 --- a/lib/compress/zstd_compress.c +++ b/lib/compress/zstd_compress.c @@ -374,6 +374,31 @@ typedef char ZSTD_rust_copy_cctx_state_layout[ == 5 * sizeof(void*) && sizeof(ZSTD_rust_copyCCtxState) == 6 * sizeof(void*)) ? 1 : -1]; +typedef void (*ZSTD_rust_compressAdvancedInitParams_f)( + void* context, const ZSTD_parameters* params); +typedef size_t (*ZSTD_rust_compressAdvancedInternal_f)( + void* context, void* dst, size_t dstCapacity, + const void* src, size_t srcSize, + const void* dict, size_t dictSize); +typedef struct { + void* callbackContext; + ZSTD_rust_compressAdvancedInitParams_f initParams; + ZSTD_rust_compressAdvancedInternal_f compressInternal; +} ZSTD_rust_compressAdvancedState; +size_t ZSTD_rust_compressAdvanced( + const ZSTD_rust_compressAdvancedState* state, + void* dst, size_t dstCapacity, + const void* src, size_t srcSize, + const void* dict, size_t dictSize, + const ZSTD_parameters* params); +typedef char ZSTD_rust_compress_advanced_state_layout[ + (offsetof(ZSTD_rust_compressAdvancedState, callbackContext) == 0 + && offsetof(ZSTD_rust_compressAdvancedState, initParams) + == sizeof(void*) + && offsetof(ZSTD_rust_compressAdvancedState, compressInternal) + == 2 * sizeof(void*) + && sizeof(ZSTD_rust_compressAdvancedState) == 3 * sizeof(void*)) + ? 1 : -1]; ZSTD_frameProgression ZSTD_rust_frameProgression(U64 consumedSrcSize, size_t buffered, U64 producedCSize); @@ -4869,20 +4894,26 @@ size_t ZSTD_compressEnd(ZSTD_CCtx* cctx, return ZSTD_compressEnd_public(cctx, dst, dstCapacity, src, srcSize); } +static void ZSTD_rust_compressAdvanced_initParams( + void* context, const ZSTD_parameters* params); +static size_t ZSTD_rust_compressAdvanced_compressInternal( + void* context, void* dst, size_t dstCapacity, + const void* src, size_t srcSize, + const void* dict, size_t dictSize); + size_t ZSTD_compress_advanced (ZSTD_CCtx* cctx, void* dst, size_t dstCapacity, const void* src, size_t srcSize, - const void* dict,size_t dictSize, + const void* dict,size_t dictSize, ZSTD_parameters params) { + ZSTD_rust_compressAdvancedState state; DEBUGLOG(4, "ZSTD_compress_advanced"); - FORWARD_IF_ERROR(ZSTD_checkCParams(params.cParams), ""); - ZSTD_CCtxParams_init_internal(&cctx->simpleApiParams, ¶ms, ZSTD_NO_CLEVEL); - return ZSTD_compress_advanced_internal(cctx, - dst, dstCapacity, - src, srcSize, - dict, dictSize, - &cctx->simpleApiParams); + state.callbackContext = cctx; + state.initParams = ZSTD_rust_compressAdvanced_initParams; + state.compressInternal = ZSTD_rust_compressAdvanced_compressInternal; + return ZSTD_rust_compressAdvanced( + &state, dst, dstCapacity, src, srcSize, dict, dictSize, ¶ms); } /* Internal */ @@ -4900,6 +4931,24 @@ size_t ZSTD_compress_advanced_internal( return ZSTD_compressEnd_public(cctx, dst, dstCapacity, src, srcSize); } +static void ZSTD_rust_compressAdvanced_initParams( + void* context, const ZSTD_parameters* params) +{ + ZSTD_CCtxParams_init_internal( + &((ZSTD_CCtx*)context)->simpleApiParams, params, ZSTD_NO_CLEVEL); +} + +static size_t ZSTD_rust_compressAdvanced_compressInternal( + void* context, void* dst, size_t dstCapacity, + const void* src, size_t srcSize, + const void* dict, size_t dictSize) +{ + ZSTD_CCtx* const cctx = (ZSTD_CCtx*)context; + return ZSTD_compress_advanced_internal( + cctx, dst, dstCapacity, src, srcSize, dict, dictSize, + &cctx->simpleApiParams); +} + size_t ZSTD_compress_usingDict(ZSTD_CCtx* cctx, void* dst, size_t dstCapacity, const void* src, size_t srcSize, diff --git a/rust/src/zstd_compress.rs b/rust/src/zstd_compress.rs index 14086042f..6d277b33f 100644 --- a/rust/src/zstd_compress.rs +++ b/rust/src/zstd_compress.rs @@ -24,8 +24,9 @@ use crate::zstd_compress_frame::{ use crate::zstd_compress_literals::min_gain; use crate::zstd_compress_params::{ ZSTD_compressionParameters, ZSTD_frameParameters, ZSTD_parameters, - ZSTD_rust_params_adjustCParams, ZSTD_rust_params_maxNbSeq, ZSTD_rust_params_selectCParams, - ZSTD_RUST_CPM_NO_ATTACH_DICT, ZSTD_RUST_PS_AUTO, ZSTD_RUST_PS_DISABLE, + ZSTD_rust_params_adjustCParams, ZSTD_rust_params_checkCParams, ZSTD_rust_params_maxNbSeq, + ZSTD_rust_params_selectCParams, ZSTD_RUST_CPM_NO_ATTACH_DICT, ZSTD_RUST_PS_AUTO, + ZSTD_RUST_PS_DISABLE, }; use crate::zstd_compress_params_api::{ ZSTD_CCtxParams_setParameter, ZSTD_CCtx_params, ZSTD_rust_isUpdateAuthorized, @@ -1468,6 +1469,84 @@ pub unsafe extern "C" fn ZSTD_rust_copyCCtx(state: *const ZSTD_rust_copyCCtxStat } } +type CompressAdvancedInitParamsFn = unsafe extern "C" fn(*mut c_void, *const ZSTD_parameters); +type CompressAdvancedInternalFn = unsafe extern "C" fn( + *mut c_void, + *mut c_void, + usize, + *const c_void, + usize, + *const c_void, + usize, +) -> usize; + +/// Explicit projection for the public `ZSTD_compress_advanced` wrapper. +/// +/// Rust owns parameter validation and the init-then-compress order. C retains +/// the private `simpleApiParams` initialization and advanced compression call. +#[repr(C)] +pub struct ZSTD_rust_compressAdvancedState { + callback_context: *mut c_void, + init_params: Option, + compress_internal: Option, +} + +const _: () = { + assert!(size_of::() == size_of::()); + assert!(size_of::() == size_of::()); + assert!(offset_of!(ZSTD_rust_compressAdvancedState, callback_context) == 0); + assert!(offset_of!(ZSTD_rust_compressAdvancedState, init_params) == size_of::()); + assert!( + offset_of!(ZSTD_rust_compressAdvancedState, compress_internal) == 2 * size_of::() + ); + assert!(size_of::() == size_of::<[usize; 3]>()); +}; + +/// Validate parameters and invoke the C-owned advanced compression operation. +#[no_mangle] +pub unsafe extern "C" fn ZSTD_rust_compressAdvanced( + state: *const ZSTD_rust_compressAdvancedState, + dst: *mut c_void, + dst_capacity: usize, + src: *const c_void, + src_size: usize, + dict: *const c_void, + dict_size: usize, + params: *const ZSTD_parameters, +) -> usize { + if state.is_null() || params.is_null() { + return ERROR(ZstdErrorCode::Generic); + } + let state = unsafe { &*state }; + let Some(init_params) = state.init_params else { + return ERROR(ZstdErrorCode::Generic); + }; + let Some(compress_internal) = state.compress_internal else { + return ERROR(ZstdErrorCode::Generic); + }; + if state.callback_context.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 { init_params(state.callback_context, params) }; + unsafe { + compress_internal( + state.callback_context, + dst, + dst_capacity, + src, + src_size, + dict, + dict_size, + ) + } +} + type ResetCStreamResetFn = unsafe extern "C" fn(*mut c_void) -> usize; type ResetCStreamSetPledgedSrcSizeFn = unsafe extern "C" fn(*mut c_void, u64) -> usize; @@ -10366,6 +10445,137 @@ mod tests { assert_eq!(context.frame_params.noDictIDFlag, 1); } + #[derive(Default)] + struct CompressAdvancedTestContext { + events: Vec<&'static str>, + init_params: ZSTD_parameters, + dst: *mut c_void, + dst_capacity: usize, + src: *const c_void, + src_size: usize, + dict: *const c_void, + dict_size: usize, + result: usize, + } + + unsafe extern "C" fn compress_advanced_test_init( + context: *mut c_void, + params: *const ZSTD_parameters, + ) { + let context = unsafe { &mut *context.cast::() }; + context.events.push("init"); + context.init_params = unsafe { *params }; + } + + unsafe extern "C" fn compress_advanced_test_internal( + context: *mut c_void, + dst: *mut c_void, + dst_capacity: usize, + src: *const c_void, + src_size: usize, + dict: *const c_void, + dict_size: usize, + ) -> usize { + let context = unsafe { &mut *context.cast::() }; + context.events.push("compress"); + context.dst = dst; + context.dst_capacity = dst_capacity; + context.src = src; + context.src_size = src_size; + context.dict = dict; + context.dict_size = dict_size; + context.result + } + + fn compress_advanced_test_state( + context: &mut CompressAdvancedTestContext, + ) -> ZSTD_rust_compressAdvancedState { + ZSTD_rust_compressAdvancedState { + callback_context: (context as *mut CompressAdvancedTestContext).cast(), + init_params: Some(compress_advanced_test_init), + compress_internal: Some(compress_advanced_test_internal), + } + } + + fn compress_advanced_test_params() -> ZSTD_parameters { + ZSTD_parameters { + cParams: ZSTD_compressionParameters { + windowLog: 10, + chainLog: 11, + hashLog: 12, + searchLog: 13, + minMatch: 4, + targetLength: 5, + strategy: 1, + }, + fParams: ZSTD_frameParameters { + contentSizeFlag: 1, + checksumFlag: 1, + noDictIDFlag: 0, + }, + } + } + + #[test] + fn compress_advanced_validates_then_initializes_and_compresses() { + let mut context = CompressAdvancedTestContext { + result: 17, + ..CompressAdvancedTestContext::default() + }; + let state = compress_advanced_test_state(&mut context); + let params = compress_advanced_test_params(); + let mut dst = [0u8; 8]; + let src = [1u8, 2, 3]; + let dict = [4u8, 5]; + + let result = unsafe { + ZSTD_rust_compressAdvanced( + &state, + dst.as_mut_ptr().cast(), + dst.len(), + src.as_ptr().cast(), + src.len(), + dict.as_ptr().cast(), + dict.len(), + ¶ms, + ) + }; + + assert_eq!(result, context.result); + assert_eq!(context.events, ["init", "compress"]); + assert_eq!(context.init_params, params); + assert_eq!(context.dst, dst.as_mut_ptr().cast()); + assert_eq!(context.dst_capacity, dst.len()); + assert_eq!(context.src, src.as_ptr().cast()); + assert_eq!(context.src_size, src.len()); + assert_eq!(context.dict, dict.as_ptr().cast()); + assert_eq!(context.dict_size, dict.len()); + } + + #[test] + fn compress_advanced_stops_before_callbacks_on_invalid_parameters() { + let mut context = CompressAdvancedTestContext::default(); + let state = compress_advanced_test_state(&mut context); + let mut params = compress_advanced_test_params(); + params.cParams.windowLog = 0; + + let result = unsafe { + ZSTD_rust_compressAdvanced( + &state, + ptr::null_mut(), + 0, + ptr::null(), + 0, + ptr::null(), + 0, + ¶ms, + ) + }; + + assert_eq!(result, ERROR(ZstdErrorCode::ParameterOutOfBound)); + assert!(context.events.is_empty()); + } + #[derive(Default)] struct ResetCStreamTestContext { events: Vec<&'static str>,