From 19c3414095285eebd79169b41b9f0847f1a32107 Mon Sep 17 00:00:00 2001 From: ddidderr Date: Tue, 21 Jul 2026 09:13:50 +0200 Subject: [PATCH] refactor(compress): move advanced one-shot order to Rust ZSTD_compress_advanced_internal() still encoded the one-shot begin-then-end sequence in C after public advanced parameter validation had moved behind Rust. That left the public compression boundary split between Rust policy and a C orchestration body, and made the begin-error ordering implicit in the C path. Add an explicit C-layout state with opaque begin and end callbacks. Rust now owns the ordering, forwards the source size as the pledged size, propagates a begin error before invoking the end callback, and returns the end result. The callbacks retain the private ZSTD_CCtx and ZSTD_CCtx_params layouts, so the change moves policy and sequencing without exposing window, workspace, matchfinder, or context internals to Rust. The existing advanced and using-dictionary wrappers continue to use the same C private operations. Test Plan: - `git diff --cached --check` -- passed. - `rustfmt --check --edition 2021 rust/src/zstd_compress.rs` -- passed. - Not run: Cargo, make, native tests, builds, or heavy verification per request. --- lib/compress/zstd_compress.c | 62 ++++++++++- rust/src/zstd_compress.rs | 206 +++++++++++++++++++++++++++++++++++ 2 files changed, 264 insertions(+), 4 deletions(-) diff --git a/lib/compress/zstd_compress.c b/lib/compress/zstd_compress.c index f3351a214..491b8bae6 100644 --- a/lib/compress/zstd_compress.c +++ b/lib/compress/zstd_compress.c @@ -678,6 +678,32 @@ typedef char ZSTD_rust_compress_advanced_state_layout[ == 2 * sizeof(void*) && sizeof(ZSTD_rust_compressAdvancedState) == 3 * sizeof(void*)) ? 1 : -1]; +typedef size_t (*ZSTD_rust_compressAdvancedInternalBegin_f)( + void* context, const void* dict, size_t dictSize, + const void* params, U64 pledgedSrcSize); +typedef size_t (*ZSTD_rust_compressAdvancedInternalEnd_f)( + void* context, void* dst, size_t dstCapacity, + const void* src, size_t srcSize); +typedef struct { + void* callbackContext; + ZSTD_rust_compressAdvancedInternalBegin_f begin; + ZSTD_rust_compressAdvancedInternalEnd_f end; +} ZSTD_rust_compressAdvancedInternalState; +size_t ZSTD_rust_compressAdvancedInternal( + const ZSTD_rust_compressAdvancedInternalState* state, + void* dst, size_t dstCapacity, + const void* src, size_t srcSize, + const void* dict, size_t dictSize, + const void* params); +typedef char ZSTD_rust_compress_advanced_internal_state_layout[ + (offsetof(ZSTD_rust_compressAdvancedInternalState, callbackContext) == 0 + && offsetof(ZSTD_rust_compressAdvancedInternalState, begin) + == sizeof(void*) + && offsetof(ZSTD_rust_compressAdvancedInternalState, end) + == 2 * sizeof(void*) + && sizeof(ZSTD_rust_compressAdvancedInternalState) + == 3 * sizeof(void*)) + ? 1 : -1]; typedef void (*ZSTD_rust_compressUsingDictInitParams_f)( void* context, const ZSTD_parameters* params, int compressionLevel); typedef struct { @@ -6950,6 +6976,12 @@ 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); +static size_t ZSTD_rust_compressAdvancedInternal_begin( + void* context, const void* dict, size_t dictSize, + const void* params, U64 pledgedSrcSize); +static size_t ZSTD_rust_compressAdvancedInternal_end( + void* context, void* dst, size_t dstCapacity, + const void* src, size_t srcSize); size_t ZSTD_compress_advanced (ZSTD_CCtx* cctx, void* dst, size_t dstCapacity, @@ -6974,11 +7006,33 @@ size_t ZSTD_compress_advanced_internal( const void* dict,size_t dictSize, const ZSTD_CCtx_params* params) { + ZSTD_rust_compressAdvancedInternalState state; DEBUGLOG(4, "ZSTD_compress_advanced_internal (srcSize:%u)", (unsigned)srcSize); - FORWARD_IF_ERROR( ZSTD_compressBegin_internal(cctx, - dict, dictSize, ZSTD_dct_auto, ZSTD_dtlm_fast, NULL, - params, srcSize, ZSTDb_not_buffered) , ""); - return ZSTD_compressEnd_public(cctx, dst, dstCapacity, src, srcSize); + state.callbackContext = cctx; + state.begin = ZSTD_rust_compressAdvancedInternal_begin; + state.end = ZSTD_rust_compressAdvancedInternal_end; + return ZSTD_rust_compressAdvancedInternal( + &state, dst, dstCapacity, src, srcSize, + dict, dictSize, params); +} + +static size_t ZSTD_rust_compressAdvancedInternal_begin( + void* context, const void* dict, size_t dictSize, + const void* params, U64 pledgedSrcSize) +{ + return ZSTD_compressBegin_internal( + (ZSTD_CCtx*)context, dict, dictSize, + ZSTD_dct_auto, ZSTD_dtlm_fast, NULL, + (const ZSTD_CCtx_params*)params, + pledgedSrcSize, ZSTDb_not_buffered); +} + +static size_t ZSTD_rust_compressAdvancedInternal_end( + void* context, void* dst, size_t dstCapacity, + const void* src, size_t srcSize) +{ + return ZSTD_compressEnd_public( + (ZSTD_CCtx*)context, dst, dstCapacity, src, srcSize); } static void ZSTD_rust_compressAdvanced_initParams( diff --git a/rust/src/zstd_compress.rs b/rust/src/zstd_compress.rs index 1b566a77d..3770c2551 100644 --- a/rust/src/zstd_compress.rs +++ b/rust/src/zstd_compress.rs @@ -3424,6 +3424,69 @@ pub unsafe extern "C" fn ZSTD_rust_compressAdvanced( } } +type CompressAdvancedInternalBeginFn = + unsafe extern "C" fn(*mut c_void, *const c_void, usize, *const c_void, u64) -> usize; +type CompressAdvancedInternalEndFn = + unsafe extern "C" fn(*mut c_void, *mut c_void, usize, *const c_void, usize) -> usize; + +/// Explicit projection for the private-operation boundary behind the public +/// advanced compression wrappers. +/// +/// Rust owns the begin-then-end ordering and stops before the end callback if +/// the private C begin operation fails. The C callbacks retain the +/// configuration-dependent `ZSTD_CCtx` and `ZSTD_CCtx_params` layouts. +#[repr(C)] +pub struct ZSTD_rust_compressAdvancedInternalState { + callback_context: *mut c_void, + begin: CompressAdvancedInternalBeginFn, + end: CompressAdvancedInternalEndFn, +} + +const _: () = { + assert!(size_of::() == size_of::()); + assert!(size_of::() == size_of::()); + assert!(offset_of!(ZSTD_rust_compressAdvancedInternalState, callback_context) == 0); + assert!(offset_of!(ZSTD_rust_compressAdvancedInternalState, begin) == size_of::()); + assert!(offset_of!(ZSTD_rust_compressAdvancedInternalState, end) == 2 * size_of::()); + assert!(size_of::() == size_of::<[usize; 3]>()); +}; + +/// Run the one-shot advanced compression order while keeping private C state +/// behind the begin and end callbacks. +#[no_mangle] +pub unsafe extern "C" fn ZSTD_rust_compressAdvancedInternal( + state: *const ZSTD_rust_compressAdvancedInternalState, + dst: *mut c_void, + dst_capacity: usize, + src: *const c_void, + src_size: usize, + dict: *const c_void, + dict_size: usize, + params: *const c_void, +) -> usize { + if state.is_null() || params.is_null() { + return ERROR(ZstdErrorCode::Generic); + } + let state = unsafe { &*state }; + if state.callback_context.is_null() { + return ERROR(ZstdErrorCode::Generic); + } + + let begin_result = unsafe { + (state.begin)( + state.callback_context, + dict, + dict_size, + params, + src_size as u64, + ) + }; + if ERR_isError(begin_result) { + return begin_result; + } + unsafe { (state.end)(state.callback_context, dst, dst_capacity, src, src_size) } +} + 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; @@ -18024,6 +18087,149 @@ mod tests { assert!(context.events.is_empty()); } + #[derive(Default)] + struct CompressAdvancedInternalTestContext { + events: Vec<&'static str>, + begin_result: usize, + end_result: usize, + begin_dict: *const c_void, + begin_dict_size: usize, + begin_params: *const c_void, + pledged_src_size: u64, + end_dst: *mut c_void, + end_dst_capacity: usize, + end_src: *const c_void, + end_src_size: usize, + } + + unsafe extern "C" fn compress_advanced_internal_test_begin( + context: *mut c_void, + dict: *const c_void, + dict_size: usize, + params: *const c_void, + pledged_src_size: u64, + ) -> usize { + let context = unsafe { &mut *context.cast::() }; + context.events.push("begin"); + context.begin_dict = dict; + context.begin_dict_size = dict_size; + context.begin_params = params; + context.pledged_src_size = pledged_src_size; + context.begin_result + } + + unsafe extern "C" fn compress_advanced_internal_test_end( + context: *mut c_void, + dst: *mut c_void, + dst_capacity: usize, + src: *const c_void, + src_size: usize, + ) -> usize { + let context = unsafe { &mut *context.cast::() }; + context.events.push("end"); + context.end_dst = dst; + context.end_dst_capacity = dst_capacity; + context.end_src = src; + context.end_src_size = src_size; + context.end_result + } + + fn compress_advanced_internal_test_state( + context: &mut CompressAdvancedInternalTestContext, + ) -> ZSTD_rust_compressAdvancedInternalState { + ZSTD_rust_compressAdvancedInternalState { + callback_context: (context as *mut CompressAdvancedInternalTestContext).cast(), + begin: compress_advanced_internal_test_begin, + end: compress_advanced_internal_test_end, + } + } + + #[test] + fn compress_advanced_internal_preserves_begin_then_end_order() { + let mut context = CompressAdvancedInternalTestContext { + end_result: 31, + ..CompressAdvancedInternalTestContext::default() + }; + let state = compress_advanced_internal_test_state(&mut context); + let mut dst = [0u8; 8]; + let src = [1u8, 2, 3]; + let dict = [4u8, 5]; + let params = [6u8; 4]; + + let result = unsafe { + ZSTD_rust_compressAdvancedInternal( + &state, + dst.as_mut_ptr().cast(), + dst.len(), + src.as_ptr().cast(), + src.len(), + dict.as_ptr().cast(), + dict.len(), + params.as_ptr().cast(), + ) + }; + + assert_eq!(result, context.end_result); + assert_eq!(context.events, ["begin", "end"]); + assert_eq!(context.begin_dict, dict.as_ptr().cast()); + assert_eq!(context.begin_dict_size, dict.len()); + assert_eq!(context.begin_params, params.as_ptr().cast()); + assert_eq!(context.pledged_src_size, src.len() as u64); + assert_eq!(context.end_dst, dst.as_mut_ptr().cast()); + assert_eq!(context.end_dst_capacity, dst.len()); + assert_eq!(context.end_src, src.as_ptr().cast()); + assert_eq!(context.end_src_size, src.len()); + } + + #[test] + fn compress_advanced_internal_stops_before_end_after_begin_error() { + let mut context = CompressAdvancedInternalTestContext { + begin_result: ERROR(ZstdErrorCode::StageWrong), + end_result: 37, + ..CompressAdvancedInternalTestContext::default() + }; + let state = compress_advanced_internal_test_state(&mut context); + let params = [0u8; 1]; + + let result = unsafe { + ZSTD_rust_compressAdvancedInternal( + &state, + ptr::null_mut(), + 0, + ptr::null(), + 17, + ptr::null(), + 0, + params.as_ptr().cast(), + ) + }; + + assert_eq!(result, ERROR(ZstdErrorCode::StageWrong)); + assert_eq!(context.events, ["begin"]); + } + + #[test] + fn compress_advanced_internal_rejects_missing_params_before_callbacks() { + let mut context = CompressAdvancedInternalTestContext::default(); + let state = compress_advanced_internal_test_state(&mut context); + + let result = unsafe { + ZSTD_rust_compressAdvancedInternal( + &state, + ptr::null_mut(), + 0, + ptr::null(), + 0, + ptr::null(), + 0, + ptr::null(), + ) + }; + + assert_eq!(result, ERROR(ZstdErrorCode::Generic)); + assert!(context.events.is_empty()); + } + #[derive(Default)] struct CompressBeginAdvancedTestContext { events: Vec<&'static str>,