diff --git a/lib/compress/zstd_compress.c b/lib/compress/zstd_compress.c index 674f2f60d..c1afcf28a 100644 --- a/lib/compress/zstd_compress.c +++ b/lib/compress/zstd_compress.c @@ -133,6 +133,21 @@ typedef char ZSTD_rust_init_cstream_src_size_state_layout[ == 4 * sizeof(void*) && sizeof(ZSTD_rust_initCStreamSrcSizeState) == 5 * sizeof(void*)) ? 1 : -1]; +typedef struct { + void* callbackContext; + ZSTD_rust_initCStreamUsingCDictAdvancedReset_f resetSession; + ZSTD_rust_initCStreamUsingCDictAdvancedRefCDict_f refCDict; + ZSTD_rust_initCStreamSrcSizeSetLevel_f setLevel; +} ZSTD_rust_initCStreamState; +size_t ZSTD_rust_initCStream(const ZSTD_rust_initCStreamState* state, + int compressionLevel); +typedef char ZSTD_rust_init_cstream_state_layout[ + (offsetof(ZSTD_rust_initCStreamState, callbackContext) == 0 + && offsetof(ZSTD_rust_initCStreamState, resetSession) == sizeof(void*) + && offsetof(ZSTD_rust_initCStreamState, refCDict) == 2 * sizeof(void*) + && offsetof(ZSTD_rust_initCStreamState, setLevel) == 3 * sizeof(void*) + && sizeof(ZSTD_rust_initCStreamState) == 4 * sizeof(void*)) + ? 1 : -1]; ZSTD_frameProgression ZSTD_rust_frameProgression(U64 consumedSrcSize, size_t buffered, U64 producedCSize); @@ -5010,11 +5025,13 @@ size_t ZSTD_initCStream_srcSize(ZSTD_CStream* zcs, int compressionLevel, unsigne size_t ZSTD_initCStream(ZSTD_CStream* zcs, int compressionLevel) { + ZSTD_rust_initCStreamState state; DEBUGLOG(4, "ZSTD_initCStream"); - FORWARD_IF_ERROR( ZSTD_CCtx_reset(zcs, ZSTD_reset_session_only) , ""); - FORWARD_IF_ERROR( ZSTD_CCtx_refCDict(zcs, NULL) , ""); - FORWARD_IF_ERROR( ZSTD_CCtx_setParameter(zcs, ZSTD_c_compressionLevel, compressionLevel) , ""); - return 0; + state.callbackContext = zcs; + state.resetSession = ZSTD_rust_initCStreamUsingCDictAdvanced_resetSession; + state.refCDict = ZSTD_rust_initCStreamUsingCDictAdvanced_refCDict; + state.setLevel = ZSTD_rust_initCStreamSrcSize_setLevel; + return ZSTD_rust_initCStream(&state, compressionLevel); } /*====== Compression ======*/ diff --git a/rust/src/zstd_compress.rs b/rust/src/zstd_compress.rs index e4600d89e..9ed9d0072 100644 --- a/rust/src/zstd_compress.rs +++ b/rust/src/zstd_compress.rs @@ -1044,6 +1044,48 @@ pub unsafe extern "C" fn ZSTD_rust_initCStreamSrcSize( 0 } +/// Explicit projection for `ZSTD_initCStream`. +#[repr(C)] +pub struct ZSTD_rust_initCStreamState { + callback_context: *mut c_void, + reset_session: InitCStreamUsingCDictAdvancedResetFn, + ref_cdict: InitCStreamUsingCDictAdvancedRefCDictFn, + set_level: InitCStreamSrcSizeSetLevelFn, +} + +const _: () = { + assert!(offset_of!(ZSTD_rust_initCStreamState, callback_context) == 0); + assert!(offset_of!(ZSTD_rust_initCStreamState, reset_session) == size_of::()); + assert!(offset_of!(ZSTD_rust_initCStreamState, ref_cdict) == 2 * size_of::()); + assert!(offset_of!(ZSTD_rust_initCStreamState, set_level) == 3 * size_of::()); + assert!(size_of::() == 4 * size_of::()); +}; + +#[no_mangle] +pub unsafe extern "C" fn ZSTD_rust_initCStream( + state: *const ZSTD_rust_initCStreamState, + compression_level: c_int, +) -> 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.ref_cdict)(state.callback_context, ptr::null()) }; + if ERR_isError(result) { + return result; + } + let result = unsafe { (state.set_level)(state.callback_context, compression_level) }; + 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; @@ -9162,6 +9204,72 @@ mod tests { assert_eq!(context.events, ["reset", "ref-cdict", "level", "pledged"]); } + fn init_cstream_test_state( + context: &mut InitCStreamUsingCDictAdvancedTestContext, + ) -> ZSTD_rust_initCStreamState { + ZSTD_rust_initCStreamState { + callback_context: (context as *mut InitCStreamUsingCDictAdvancedTestContext).cast(), + reset_session: init_cstream_using_cdict_advanced_test_reset, + ref_cdict: init_cstream_using_cdict_advanced_test_ref_cdict, + set_level: init_cstream_src_size_test_set_level, + } + } + + #[test] + fn init_cstream_preserves_reset_clear_and_level_order() { + let mut context = InitCStreamUsingCDictAdvancedTestContext::default(); + let state = init_cstream_test_state(&mut context); + + let result = unsafe { ZSTD_rust_initCStream(&state, -3) }; + + assert_eq!(result, 0); + assert_eq!(context.events, ["reset", "ref-cdict", "level"]); + assert_eq!(context.cdict, ptr::null()); + assert_eq!(context.compression_level, -3); + } + + #[test] + fn init_cstream_stops_after_reset_error() { + let mut context = InitCStreamUsingCDictAdvancedTestContext { + reset_result: ERROR(ZstdErrorCode::MemoryAllocation), + ..InitCStreamUsingCDictAdvancedTestContext::default() + }; + let state = init_cstream_test_state(&mut context); + + let result = unsafe { ZSTD_rust_initCStream(&state, 4) }; + + assert_eq!(result, ERROR(ZstdErrorCode::MemoryAllocation)); + assert_eq!(context.events, ["reset"]); + } + + #[test] + fn init_cstream_stops_after_cdict_clear_error() { + let mut context = InitCStreamUsingCDictAdvancedTestContext { + ref_result: ERROR(ZstdErrorCode::DictionaryCreationFailed), + ..InitCStreamUsingCDictAdvancedTestContext::default() + }; + let state = init_cstream_test_state(&mut context); + + let result = unsafe { ZSTD_rust_initCStream(&state, 4) }; + + assert_eq!(result, ERROR(ZstdErrorCode::DictionaryCreationFailed)); + assert_eq!(context.events, ["reset", "ref-cdict"]); + } + + #[test] + fn init_cstream_propagates_level_error_last() { + let mut context = InitCStreamUsingCDictAdvancedTestContext { + set_level_result: ERROR(ZstdErrorCode::ParameterOutOfBound), + ..InitCStreamUsingCDictAdvancedTestContext::default() + }; + let state = init_cstream_test_state(&mut context); + + let result = unsafe { ZSTD_rust_initCStream(&state, 4) }; + + assert_eq!(result, ERROR(ZstdErrorCode::ParameterOutOfBound)); + assert_eq!(context.events, ["reset", "ref-cdict", "level"]); + } + #[test] fn pledged_src_size_writes_the_init_stage_value_plus_one() { let mut pledged_src_size_plus_one = 0;