From af1fd3e0fc03a2f32f7ef926bfbed79a2616a399 Mon Sep 17 00:00:00 2001 From: ddidderr Date: Sun, 19 Jul 2026 12:37:52 +0200 Subject: [PATCH] feat(compress): move advanced stream init policy into Rust Move ZSTD_initCStream_advanced's public initialization policy into the Rust projection. Rust owns legacy pledged-size normalization and the reset, pledge, parameter-check, parameter-copy, and dictionary-load ordering while C retains the private CCtx parameter mutation callbacks. Test Plan: - ulimit -v 41943040; CARGO_BUILD_JOBS=1 cargo test --manifest-path rust/Cargo.toml init_cstream_advanced -- --test-threads=1 - ulimit -v 41943040; CARGO_BUILD_JOBS=1 cargo clippy --manifest-path rust/Cargo.toml --all-targets -- -D warnings - ulimit -v 41943040; make -B -C programs -j1 zstd - git diff --cached --check --- lib/compress/zstd_compress.c | 70 +++++++++-- rust/src/zstd_compress.rs | 230 ++++++++++++++++++++++++++++++++++- 2 files changed, 285 insertions(+), 15 deletions(-) diff --git a/lib/compress/zstd_compress.c b/lib/compress/zstd_compress.c index b4522ef1f..338228ce2 100644 --- a/lib/compress/zstd_compress.c +++ b/lib/compress/zstd_compress.c @@ -167,6 +167,35 @@ typedef char ZSTD_rust_init_cstream_using_dict_state_layout[ == 3 * sizeof(void*) && sizeof(ZSTD_rust_initCStreamUsingDictState) == 4 * sizeof(void*)) ? 1 : -1]; +typedef size_t (*ZSTD_rust_initCStreamAdvancedCheckCParams_f)( + void* context, ZSTD_compressionParameters cParams); +typedef void (*ZSTD_rust_initCStreamAdvancedSetZstdParams_f)( + void* context, const ZSTD_parameters* params); +typedef struct { + void* callbackContext; + ZSTD_rust_initCStreamUsingCDictAdvancedReset_f resetSession; + ZSTD_rust_initCStreamUsingCDictAdvancedSetPledgedSrcSize_f setPledgedSrcSize; + ZSTD_rust_initCStreamAdvancedCheckCParams_f checkCParams; + ZSTD_rust_initCStreamAdvancedSetZstdParams_f setZstdParams; + ZSTD_rust_initCStreamUsingDictLoadDictionary_f loadDictionary; +} ZSTD_rust_initCStreamAdvancedState; +size_t ZSTD_rust_initCStreamAdvanced( + const ZSTD_rust_initCStreamAdvancedState* state, + const ZSTD_parameters* params, unsigned long long pss, + const void* dict, size_t dictSize); +typedef char ZSTD_rust_init_cstream_advanced_state_layout[ + (offsetof(ZSTD_rust_initCStreamAdvancedState, callbackContext) == 0 + && offsetof(ZSTD_rust_initCStreamAdvancedState, resetSession) == sizeof(void*) + && offsetof(ZSTD_rust_initCStreamAdvancedState, setPledgedSrcSize) + == 2 * sizeof(void*) + && offsetof(ZSTD_rust_initCStreamAdvancedState, checkCParams) + == 3 * sizeof(void*) + && offsetof(ZSTD_rust_initCStreamAdvancedState, setZstdParams) + == 4 * sizeof(void*) + && offsetof(ZSTD_rust_initCStreamAdvancedState, loadDictionary) + == 5 * sizeof(void*) + && sizeof(ZSTD_rust_initCStreamAdvancedState) == 6 * sizeof(void*)) + ? 1 : -1]; ZSTD_frameProgression ZSTD_rust_frameProgression(U64 consumedSrcSize, size_t buffered, U64 producedCSize); @@ -4992,22 +5021,43 @@ size_t ZSTD_initCStream_usingCDict(ZSTD_CStream* zcs, const ZSTD_CDict* cdict) * pledgedSrcSize must be exact. * if srcSize is not known at init time, use value ZSTD_CONTENTSIZE_UNKNOWN. * dict is loaded with default parameters ZSTD_dct_auto and ZSTD_dlm_byCopy. */ +static size_t ZSTD_rust_initCStreamAdvanced_checkCParams( + void* context, ZSTD_compressionParameters cParams) +{ + (void)context; + return ZSTD_checkCParams(cParams); +} + +static void ZSTD_rust_initCStreamAdvanced_setZstdParams( + void* context, const ZSTD_parameters* params) +{ + ZSTD_CCtx* const zcs = (ZSTD_CCtx*)context; + ZSTD_CCtxParams_setZstdParams(&zcs->requestedParams, params); +} + +static size_t ZSTD_rust_initCStreamUsingDict_loadDictionary( + void* context, const void* dict, size_t dictSize) +{ + return ZSTD_CCtx_loadDictionary((ZSTD_CCtx*)context, dict, dictSize); +} + size_t ZSTD_initCStream_advanced(ZSTD_CStream* zcs, const void* dict, size_t dictSize, ZSTD_parameters params, unsigned long long pss) { + ZSTD_rust_initCStreamAdvancedState state; /* for compatibility with older programs relying on this behavior. * Users should now specify ZSTD_CONTENTSIZE_UNKNOWN. * This line will be removed in the future. */ - U64 const pledgedSrcSize = (pss==0 && params.fParams.contentSizeFlag==0) ? ZSTD_CONTENTSIZE_UNKNOWN : pss; DEBUGLOG(4, "ZSTD_initCStream_advanced"); - FORWARD_IF_ERROR( ZSTD_CCtx_reset(zcs, ZSTD_reset_session_only) , ""); - FORWARD_IF_ERROR( ZSTD_CCtx_setPledgedSrcSize(zcs, pledgedSrcSize) , ""); - FORWARD_IF_ERROR( ZSTD_checkCParams(params.cParams) , ""); - ZSTD_CCtxParams_setZstdParams(&zcs->requestedParams, ¶ms); - FORWARD_IF_ERROR( ZSTD_CCtx_loadDictionary(zcs, dict, dictSize) , ""); - return 0; + state.callbackContext = zcs; + state.resetSession = ZSTD_rust_initCStreamUsingCDictAdvanced_resetSession; + state.setPledgedSrcSize = ZSTD_rust_initCStreamUsingCDictAdvanced_setPledgedSrcSize; + state.checkCParams = ZSTD_rust_initCStreamAdvanced_checkCParams; + state.setZstdParams = ZSTD_rust_initCStreamAdvanced_setZstdParams; + state.loadDictionary = ZSTD_rust_initCStreamUsingDict_loadDictionary; + return ZSTD_rust_initCStreamAdvanced(&state, ¶ms, pss, dict, dictSize); } static size_t ZSTD_rust_initCStreamSrcSize_setLevel( @@ -5017,12 +5067,6 @@ static size_t ZSTD_rust_initCStreamSrcSize_setLevel( (ZSTD_CCtx*)context, ZSTD_c_compressionLevel, compressionLevel); } -static size_t ZSTD_rust_initCStreamUsingDict_loadDictionary( - void* context, const void* dict, size_t dictSize) -{ - return ZSTD_CCtx_loadDictionary((ZSTD_CCtx*)context, dict, dictSize); -} - size_t ZSTD_initCStream_usingDict(ZSTD_CStream* zcs, const void* dict, size_t dictSize, int compressionLevel) { ZSTD_rust_initCStreamUsingDictState state; diff --git a/rust/src/zstd_compress.rs b/rust/src/zstd_compress.rs index cf48d57f4..d23992c9e 100644 --- a/rust/src/zstd_compress.rs +++ b/rust/src/zstd_compress.rs @@ -23,8 +23,9 @@ use crate::zstd_compress_frame::{ }; use crate::zstd_compress_literals::min_gain; use crate::zstd_compress_params::{ - 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_compressionParameters, 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, }; use crate::zstd_compress_sequences::SeqDef; use crate::zstd_compress_stats::{ @@ -1135,6 +1136,79 @@ pub unsafe extern "C" fn ZSTD_rust_initCStreamUsingDict( 0 } +type InitCStreamAdvancedCheckCParamsFn = + unsafe extern "C" fn(*mut c_void, ZSTD_compressionParameters) -> usize; +type InitCStreamAdvancedSetZstdParamsFn = unsafe extern "C" fn(*mut c_void, *const ZSTD_parameters); + +/// Explicit projection for `ZSTD_initCStream_advanced`. +#[repr(C)] +pub struct ZSTD_rust_initCStreamAdvancedState { + callback_context: *mut c_void, + reset_session: InitCStreamUsingCDictAdvancedResetFn, + set_pledged_src_size: InitCStreamUsingCDictAdvancedSetPledgedSrcSizeFn, + check_c_params: InitCStreamAdvancedCheckCParamsFn, + set_zstd_params: InitCStreamAdvancedSetZstdParamsFn, + load_dictionary: InitCStreamUsingDictLoadDictionaryFn, +} + +const _: () = { + assert!(offset_of!(ZSTD_rust_initCStreamAdvancedState, callback_context) == 0); + assert!(offset_of!(ZSTD_rust_initCStreamAdvancedState, reset_session) == size_of::()); + assert!( + offset_of!(ZSTD_rust_initCStreamAdvancedState, set_pledged_src_size) + == 2 * size_of::() + ); + assert!( + offset_of!(ZSTD_rust_initCStreamAdvancedState, check_c_params) == 3 * size_of::() + ); + assert!( + offset_of!(ZSTD_rust_initCStreamAdvancedState, set_zstd_params) == 4 * size_of::() + ); + assert!( + offset_of!(ZSTD_rust_initCStreamAdvancedState, load_dictionary) == 5 * size_of::() + ); + assert!(size_of::() == 6 * size_of::()); +}; + +#[no_mangle] +pub unsafe extern "C" fn ZSTD_rust_initCStreamAdvanced( + state: *const ZSTD_rust_initCStreamAdvancedState, + params: *const ZSTD_parameters, + pss: u64, + dict: *const c_void, + dict_size: usize, +) -> usize { + if state.is_null() || params.is_null() { + return ERROR(ZstdErrorCode::Generic); + } + let state = unsafe { &*state }; + let params = unsafe { &*params }; + let pledged_src_size = if pss == 0 && params.fParams.contentSizeFlag == 0 { + ZSTD_CONTENTSIZE_UNKNOWN + } else { + pss + }; + + 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; + } + let result = unsafe { (state.check_c_params)(state.callback_context, params.cParams) }; + if ERR_isError(result) { + return result; + } + unsafe { (state.set_zstd_params)(state.callback_context, params) }; + let result = unsafe { (state.load_dictionary)(state.callback_context, dict, dict_size) }; + 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; @@ -5443,6 +5517,7 @@ pub unsafe extern "C" fn ZSTD_compressStream2( mod tests { use super::*; use crate::errors::ERR_getErrorCode; + use crate::zstd_compress_params::ZSTD_frameParameters; use std::io::Write; use std::process::{Command, Stdio}; @@ -8986,12 +9061,14 @@ mod tests { pledged_result: usize, set_level_result: usize, load_dict_result: usize, + check_params_result: usize, ref_result: usize, pledged_src_size: u64, compression_level: c_int, frame_params: [c_uint; 3], cdict: *const c_void, dict_size: usize, + zstd_params: ZSTD_parameters, } unsafe fn init_cstream_using_cdict_advanced_test_context( @@ -9040,6 +9117,24 @@ mod tests { context.load_dict_result } + unsafe extern "C" fn init_cstream_advanced_test_check_c_params( + context: *mut c_void, + _c_params: ZSTD_compressionParameters, + ) -> usize { + let context = unsafe { init_cstream_using_cdict_advanced_test_context(context) }; + context.events.push("check"); + context.check_params_result + } + + unsafe extern "C" fn init_cstream_advanced_test_set_zstd_params( + context: *mut c_void, + params: *const ZSTD_parameters, + ) { + let context = unsafe { init_cstream_using_cdict_advanced_test_context(context) }; + context.events.push("params"); + context.zstd_params = unsafe { *params }; + } + unsafe extern "C" fn init_cstream_using_cdict_advanced_test_set_frame_params( context: *mut c_void, content_size_flag: c_uint, @@ -9401,6 +9496,137 @@ mod tests { assert_eq!(context.events, ["reset", "level", "load-dict"]); } + fn init_cstream_advanced_test_state( + context: &mut InitCStreamUsingCDictAdvancedTestContext, + ) -> ZSTD_rust_initCStreamAdvancedState { + ZSTD_rust_initCStreamAdvancedState { + 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, + check_c_params: init_cstream_advanced_test_check_c_params, + set_zstd_params: init_cstream_advanced_test_set_zstd_params, + load_dictionary: init_cstream_using_dict_test_load_dictionary, + } + } + + fn init_cstream_advanced_test_params(content_size_flag: c_int) -> 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: content_size_flag, + checksumFlag: 1, + noDictIDFlag: 2, + }, + } + } + + #[test] + fn init_cstream_advanced_preserves_order_and_normalizes_unknown_pledge() { + let params = init_cstream_advanced_test_params(0); + let dict = ptr::dangling::(); + let mut context = InitCStreamUsingCDictAdvancedTestContext::default(); + let state = init_cstream_advanced_test_state(&mut context); + + let result = unsafe { ZSTD_rust_initCStreamAdvanced(&state, ¶ms, 0, dict, 123) }; + + assert_eq!(result, 0); + assert_eq!( + context.events, + ["reset", "pledged", "check", "params", "load-dict"] + ); + assert_eq!(context.pledged_src_size, ZSTD_CONTENTSIZE_UNKNOWN); + assert_eq!(context.zstd_params, params); + assert_eq!(context.cdict, dict); + assert_eq!(context.dict_size, 123); + } + + #[test] + fn init_cstream_advanced_keeps_zero_pledge_for_known_empty_frame() { + let params = init_cstream_advanced_test_params(1); + let mut context = InitCStreamUsingCDictAdvancedTestContext::default(); + let state = init_cstream_advanced_test_state(&mut context); + + let result = unsafe { ZSTD_rust_initCStreamAdvanced(&state, ¶ms, 0, ptr::null(), 0) }; + + assert_eq!(result, 0); + assert_eq!(context.pledged_src_size, 0); + assert_eq!( + context.events, + ["reset", "pledged", "check", "params", "load-dict"] + ); + } + + #[test] + fn init_cstream_advanced_stops_after_reset_error() { + let params = init_cstream_advanced_test_params(0); + let mut context = InitCStreamUsingCDictAdvancedTestContext { + reset_result: ERROR(ZstdErrorCode::MemoryAllocation), + ..InitCStreamUsingCDictAdvancedTestContext::default() + }; + let state = init_cstream_advanced_test_state(&mut context); + + let result = unsafe { ZSTD_rust_initCStreamAdvanced(&state, ¶ms, 77, ptr::null(), 0) }; + + assert_eq!(result, ERROR(ZstdErrorCode::MemoryAllocation)); + assert_eq!(context.events, ["reset"]); + } + + #[test] + fn init_cstream_advanced_stops_after_pledged_size_error() { + let params = init_cstream_advanced_test_params(0); + let mut context = InitCStreamUsingCDictAdvancedTestContext { + pledged_result: ERROR(ZstdErrorCode::StageWrong), + ..InitCStreamUsingCDictAdvancedTestContext::default() + }; + let state = init_cstream_advanced_test_state(&mut context); + + let result = unsafe { ZSTD_rust_initCStreamAdvanced(&state, ¶ms, 77, ptr::null(), 0) }; + + assert_eq!(result, ERROR(ZstdErrorCode::StageWrong)); + assert_eq!(context.events, ["reset", "pledged"]); + } + + #[test] + fn init_cstream_advanced_stops_after_parameter_error() { + let params = init_cstream_advanced_test_params(0); + let mut context = InitCStreamUsingCDictAdvancedTestContext { + check_params_result: ERROR(ZstdErrorCode::ParameterOutOfBound), + ..InitCStreamUsingCDictAdvancedTestContext::default() + }; + let state = init_cstream_advanced_test_state(&mut context); + + let result = unsafe { ZSTD_rust_initCStreamAdvanced(&state, ¶ms, 77, ptr::null(), 0) }; + + assert_eq!(result, ERROR(ZstdErrorCode::ParameterOutOfBound)); + assert_eq!(context.events, ["reset", "pledged", "check"]); + } + + #[test] + fn init_cstream_advanced_propagates_dictionary_error_last() { + let params = init_cstream_advanced_test_params(0); + let mut context = InitCStreamUsingCDictAdvancedTestContext { + load_dict_result: ERROR(ZstdErrorCode::DictionaryCreationFailed), + ..InitCStreamUsingCDictAdvancedTestContext::default() + }; + let state = init_cstream_advanced_test_state(&mut context); + + let result = unsafe { ZSTD_rust_initCStreamAdvanced(&state, ¶ms, 77, ptr::null(), 0) }; + + assert_eq!(result, ERROR(ZstdErrorCode::DictionaryCreationFailed)); + assert_eq!( + context.events, + ["reset", "pledged", "check", "params", "load-dict"] + ); + } + #[test] fn pledged_src_size_writes_the_init_stage_value_plus_one() { let mut pledged_src_size_plus_one = 0;