From b0c03f41f5ee10e8d769c8bc505f2b65f1da3f12 Mon Sep 17 00:00:00 2001 From: ddidderr Date: Sun, 19 Jul 2026 11:58:35 +0200 Subject: [PATCH] feat(compress): move cdict stream init policy into Rust Move ZSTD_initCStream_usingCDict_advanced's reset, pledge, frame-parameter, and CDict callback ordering into the Rust projection. Keep private CCtx and CDict mutations in C callbacks and pass frame parameters as scalars across the ABI boundary. Test Plan: - ulimit -v 41943040; CARGO_BUILD_JOBS=1 cargo test --manifest-path rust/Cargo.toml init_cstream_using_cdict_advanced -- --test-threads=1 - ulimit -v 41943040; make -B -C programs -j1 zstd - git diff --cached --check --- lib/compress/zstd_compress.c | 72 +++++++++++- rust/src/zstd_compress.rs | 221 +++++++++++++++++++++++++++++++++++ 2 files changed, 288 insertions(+), 5 deletions(-) diff --git a/lib/compress/zstd_compress.c b/lib/compress/zstd_compress.c index e25bb2c18..6fd43af7e 100644 --- a/lib/compress/zstd_compress.c +++ b/lib/compress/zstd_compress.c @@ -68,6 +68,37 @@ typedef char ZSTD_rust_reset_cstream_state_layout[ == 2 * sizeof(void*) && sizeof(ZSTD_rust_resetCStreamState) == 3 * sizeof(void*)) ? 1 : -1]; +typedef size_t (*ZSTD_rust_initCStreamUsingCDictAdvancedReset_f)(void* context); +typedef size_t (*ZSTD_rust_initCStreamUsingCDictAdvancedSetPledgedSrcSize_f)( + void* context, unsigned long long pledgedSrcSize); +typedef void (*ZSTD_rust_initCStreamUsingCDictAdvancedSetFrameParams_f)( + void* context, unsigned contentSizeFlag, unsigned checksumFlag, + unsigned noDictIDFlag); +typedef size_t (*ZSTD_rust_initCStreamUsingCDictAdvancedRefCDict_f)( + void* context, const void* cdict); +typedef struct { + void* callbackContext; + ZSTD_rust_initCStreamUsingCDictAdvancedReset_f resetSession; + ZSTD_rust_initCStreamUsingCDictAdvancedSetPledgedSrcSize_f setPledgedSrcSize; + ZSTD_rust_initCStreamUsingCDictAdvancedSetFrameParams_f setFrameParams; + ZSTD_rust_initCStreamUsingCDictAdvancedRefCDict_f refCDict; +} ZSTD_rust_initCStreamUsingCDictAdvancedState; +size_t ZSTD_rust_initCStreamUsingCDictAdvanced( + const ZSTD_rust_initCStreamUsingCDictAdvancedState* state, + unsigned long long pledgedSrcSize, unsigned contentSizeFlag, + unsigned checksumFlag, unsigned noDictIDFlag, const void* cdict); +typedef char ZSTD_rust_init_cstream_using_cdict_advanced_state_layout[ + (offsetof(ZSTD_rust_initCStreamUsingCDictAdvancedState, callbackContext) == 0 + && offsetof(ZSTD_rust_initCStreamUsingCDictAdvancedState, resetSession) + == sizeof(void*) + && offsetof(ZSTD_rust_initCStreamUsingCDictAdvancedState, setPledgedSrcSize) + == 2 * sizeof(void*) + && offsetof(ZSTD_rust_initCStreamUsingCDictAdvancedState, setFrameParams) + == 3 * sizeof(void*) + && offsetof(ZSTD_rust_initCStreamUsingCDictAdvancedState, refCDict) + == 4 * sizeof(void*) + && sizeof(ZSTD_rust_initCStreamUsingCDictAdvancedState) == 5 * sizeof(void*)) + ? 1 : -1]; ZSTD_frameProgression ZSTD_rust_frameProgression(U64 consumedSrcSize, size_t buffered, U64 producedCSize); @@ -4833,17 +4864,48 @@ size_t ZSTD_initCStream_internal(ZSTD_CStream* zcs, /* ZSTD_initCStream_usingCDict_advanced() : * same as ZSTD_initCStream_usingCDict(), with control over frame parameters */ +static size_t ZSTD_rust_initCStreamUsingCDictAdvanced_resetSession(void* context) +{ + return ZSTD_CCtx_reset((ZSTD_CCtx*)context, ZSTD_reset_session_only); +} + +static size_t ZSTD_rust_initCStreamUsingCDictAdvanced_setPledgedSrcSize( + void* context, unsigned long long pledgedSrcSize) +{ + return ZSTD_CCtx_setPledgedSrcSize((ZSTD_CCtx*)context, pledgedSrcSize); +} + +static void ZSTD_rust_initCStreamUsingCDictAdvanced_setFrameParams( + void* context, unsigned contentSizeFlag, unsigned checksumFlag, + unsigned noDictIDFlag) +{ + ZSTD_CCtx* const zcs = (ZSTD_CCtx*)context; + zcs->requestedParams.fParams.contentSizeFlag = contentSizeFlag; + zcs->requestedParams.fParams.checksumFlag = checksumFlag; + zcs->requestedParams.fParams.noDictIDFlag = noDictIDFlag; +} + +static size_t ZSTD_rust_initCStreamUsingCDictAdvanced_refCDict( + void* context, const void* cdict) +{ + return ZSTD_CCtx_refCDict((ZSTD_CCtx*)context, (const ZSTD_CDict*)cdict); +} + size_t ZSTD_initCStream_usingCDict_advanced(ZSTD_CStream* zcs, const ZSTD_CDict* cdict, ZSTD_frameParameters fParams, unsigned long long pledgedSrcSize) { + ZSTD_rust_initCStreamUsingCDictAdvancedState state; DEBUGLOG(4, "ZSTD_initCStream_usingCDict_advanced"); - FORWARD_IF_ERROR( ZSTD_CCtx_reset(zcs, ZSTD_reset_session_only) , ""); - FORWARD_IF_ERROR( ZSTD_CCtx_setPledgedSrcSize(zcs, pledgedSrcSize) , ""); - zcs->requestedParams.fParams = fParams; - FORWARD_IF_ERROR( ZSTD_CCtx_refCDict(zcs, cdict) , ""); - return 0; + state.callbackContext = zcs; + state.resetSession = ZSTD_rust_initCStreamUsingCDictAdvanced_resetSession; + state.setPledgedSrcSize = ZSTD_rust_initCStreamUsingCDictAdvanced_setPledgedSrcSize; + state.setFrameParams = ZSTD_rust_initCStreamUsingCDictAdvanced_setFrameParams; + state.refCDict = ZSTD_rust_initCStreamUsingCDictAdvanced_refCDict; + return ZSTD_rust_initCStreamUsingCDictAdvanced( + &state, pledgedSrcSize, fParams.contentSizeFlag, + fParams.checksumFlag, fParams.noDictIDFlag, cdict); } /* note : cdict must outlive compression session */ diff --git a/rust/src/zstd_compress.rs b/rust/src/zstd_compress.rs index a8334d59f..495581352 100644 --- a/rust/src/zstd_compress.rs +++ b/rust/src/zstd_compress.rs @@ -857,6 +857,98 @@ pub unsafe extern "C" fn ZSTD_rust_resetCStream( 0 } +type InitCStreamUsingCDictAdvancedResetFn = unsafe extern "C" fn(*mut c_void) -> usize; +type InitCStreamUsingCDictAdvancedSetPledgedSrcSizeFn = + unsafe extern "C" fn(*mut c_void, u64) -> usize; +type InitCStreamUsingCDictAdvancedSetFrameParamsFn = + unsafe extern "C" fn(*mut c_void, c_uint, c_uint, c_uint); +type InitCStreamUsingCDictAdvancedRefCDictFn = + unsafe extern "C" fn(*mut c_void, *const c_void) -> usize; + +/// Explicit projection for `ZSTD_initCStream_usingCDict_advanced`. +/// +/// Rust owns the public wrapper's callback order and error propagation while +/// C retains the private context and dictionary/parameter mutations. +#[repr(C)] +pub struct ZSTD_rust_initCStreamUsingCDictAdvancedState { + callback_context: *mut c_void, + reset_session: InitCStreamUsingCDictAdvancedResetFn, + set_pledged_src_size: InitCStreamUsingCDictAdvancedSetPledgedSrcSizeFn, + set_frame_params: InitCStreamUsingCDictAdvancedSetFrameParamsFn, + ref_cdict: InitCStreamUsingCDictAdvancedRefCDictFn, +} + +const _: () = { + assert!(size_of::() == size_of::()); + assert!(size_of::() == size_of::()); + assert!(size_of::() == size_of::()); + assert!(size_of::() == size_of::()); + assert!( + offset_of!( + ZSTD_rust_initCStreamUsingCDictAdvancedState, + callback_context + ) == 0 + ); + assert!( + offset_of!(ZSTD_rust_initCStreamUsingCDictAdvancedState, reset_session) + == size_of::() + ); + assert!( + offset_of!( + ZSTD_rust_initCStreamUsingCDictAdvancedState, + set_pledged_src_size + ) == 2 * size_of::() + ); + assert!( + offset_of!( + ZSTD_rust_initCStreamUsingCDictAdvancedState, + set_frame_params + ) == 3 * size_of::() + ); + assert!( + offset_of!(ZSTD_rust_initCStreamUsingCDictAdvancedState, ref_cdict) + == 4 * size_of::() + ); + assert!(size_of::() == 5 * size_of::()); +}; + +#[no_mangle] +pub unsafe extern "C" fn ZSTD_rust_initCStreamUsingCDictAdvanced( + state: *const ZSTD_rust_initCStreamUsingCDictAdvancedState, + pledged_src_size: u64, + content_size_flag: c_uint, + checksum_flag: c_uint, + no_dict_id_flag: c_uint, + cdict: *const c_void, +) -> usize { + if state.is_null() { + return ERROR(ZstdErrorCode::Generic); + } + let state = unsafe { &*state }; + + let result = unsafe { (state.reset_session)(state.callback_context) }; + if ERR_isError(result) { + return result; + } + let result = unsafe { (state.set_pledged_src_size)(state.callback_context, pledged_src_size) }; + if ERR_isError(result) { + return result; + } + unsafe { + (state.set_frame_params)( + state.callback_context, + content_size_flag, + checksum_flag, + no_dict_id_flag, + ) + }; + let result = unsafe { (state.ref_cdict)(state.callback_context, cdict) }; + if ERR_isError(result) { + return result; + } + 0 +} + type CompressStreamBlockFn = unsafe extern "C" fn(*mut c_void, *mut c_void, usize, *const c_void, usize) -> usize; type CompressStreamResetFn = unsafe extern "C" fn(*mut c_void) -> usize; @@ -8701,6 +8793,135 @@ mod tests { assert_eq!(context.pledged_src_size, 123); } + #[derive(Default)] + struct InitCStreamUsingCDictAdvancedTestContext { + events: Vec<&'static str>, + reset_result: usize, + pledged_result: usize, + ref_result: usize, + pledged_src_size: u64, + frame_params: [c_uint; 3], + cdict: *const c_void, + } + + unsafe fn init_cstream_using_cdict_advanced_test_context( + context: *mut c_void, + ) -> &'static mut InitCStreamUsingCDictAdvancedTestContext { + unsafe { &mut *context.cast::() } + } + + unsafe extern "C" fn init_cstream_using_cdict_advanced_test_reset( + context: *mut c_void, + ) -> usize { + let context = unsafe { init_cstream_using_cdict_advanced_test_context(context) }; + context.events.push("reset"); + context.reset_result + } + + unsafe extern "C" fn init_cstream_using_cdict_advanced_test_set_pledged( + context: *mut c_void, + pledged_src_size: u64, + ) -> usize { + let context = unsafe { init_cstream_using_cdict_advanced_test_context(context) }; + context.events.push("pledged"); + context.pledged_src_size = pledged_src_size; + context.pledged_result + } + + unsafe extern "C" fn init_cstream_using_cdict_advanced_test_set_frame_params( + context: *mut c_void, + content_size_flag: c_uint, + checksum_flag: c_uint, + no_dict_id_flag: c_uint, + ) { + let context = unsafe { init_cstream_using_cdict_advanced_test_context(context) }; + context.events.push("frame"); + context.frame_params = [content_size_flag, checksum_flag, no_dict_id_flag]; + } + + unsafe extern "C" fn init_cstream_using_cdict_advanced_test_ref_cdict( + context: *mut c_void, + cdict: *const c_void, + ) -> usize { + let context = unsafe { init_cstream_using_cdict_advanced_test_context(context) }; + context.events.push("ref-cdict"); + context.cdict = cdict; + context.ref_result + } + + fn init_cstream_using_cdict_advanced_test_state( + context: &mut InitCStreamUsingCDictAdvancedTestContext, + ) -> ZSTD_rust_initCStreamUsingCDictAdvancedState { + ZSTD_rust_initCStreamUsingCDictAdvancedState { + callback_context: (context as *mut InitCStreamUsingCDictAdvancedTestContext).cast(), + reset_session: init_cstream_using_cdict_advanced_test_reset, + set_pledged_src_size: init_cstream_using_cdict_advanced_test_set_pledged, + set_frame_params: init_cstream_using_cdict_advanced_test_set_frame_params, + ref_cdict: init_cstream_using_cdict_advanced_test_ref_cdict, + } + } + + #[test] + fn init_cstream_using_cdict_advanced_preserves_order_and_scalars() { + let cdict = ptr::dangling::(); + let mut context = InitCStreamUsingCDictAdvancedTestContext::default(); + let state = init_cstream_using_cdict_advanced_test_state(&mut context); + + let result = unsafe { ZSTD_rust_initCStreamUsingCDictAdvanced(&state, 77, 1, 2, 3, cdict) }; + + assert_eq!(result, 0); + assert_eq!(context.events, ["reset", "pledged", "frame", "ref-cdict"]); + assert_eq!(context.pledged_src_size, 77); + assert_eq!(context.frame_params, [1, 2, 3]); + assert_eq!(context.cdict, cdict); + } + + #[test] + fn init_cstream_using_cdict_advanced_stops_after_reset_error() { + let mut context = InitCStreamUsingCDictAdvancedTestContext { + reset_result: ERROR(ZstdErrorCode::MemoryAllocation), + ..InitCStreamUsingCDictAdvancedTestContext::default() + }; + let state = init_cstream_using_cdict_advanced_test_state(&mut context); + + let result = + unsafe { ZSTD_rust_initCStreamUsingCDictAdvanced(&state, 77, 1, 2, 3, ptr::null()) }; + + assert_eq!(result, ERROR(ZstdErrorCode::MemoryAllocation)); + assert_eq!(context.events, ["reset"]); + } + + #[test] + fn init_cstream_using_cdict_advanced_stops_after_pledged_size_error() { + let mut context = InitCStreamUsingCDictAdvancedTestContext { + pledged_result: ERROR(ZstdErrorCode::StageWrong), + ..InitCStreamUsingCDictAdvancedTestContext::default() + }; + let state = init_cstream_using_cdict_advanced_test_state(&mut context); + + let result = + unsafe { ZSTD_rust_initCStreamUsingCDictAdvanced(&state, 77, 1, 2, 3, ptr::null()) }; + + assert_eq!(result, ERROR(ZstdErrorCode::StageWrong)); + assert_eq!(context.events, ["reset", "pledged"]); + } + + #[test] + fn init_cstream_using_cdict_advanced_propagates_ref_error_after_frame_params() { + let mut context = InitCStreamUsingCDictAdvancedTestContext { + ref_result: ERROR(ZstdErrorCode::DictionaryCreationFailed), + ..InitCStreamUsingCDictAdvancedTestContext::default() + }; + let state = init_cstream_using_cdict_advanced_test_state(&mut context); + + let result = + unsafe { ZSTD_rust_initCStreamUsingCDictAdvanced(&state, 77, 1, 2, 3, ptr::null()) }; + + assert_eq!(result, ERROR(ZstdErrorCode::DictionaryCreationFailed)); + assert_eq!(context.events, ["reset", "pledged", "frame", "ref-cdict"]); + assert_eq!(context.frame_params, [1, 2, 3]); + } + #[test] fn pledged_src_size_writes_the_init_stage_value_plus_one() { let mut pledged_src_size_plus_one = 0;