diff --git a/lib/compress/zstd_compress.c b/lib/compress/zstd_compress.c index ae4666096..5f437ba56 100644 --- a/lib/compress/zstd_compress.c +++ b/lib/compress/zstd_compress.c @@ -61,6 +61,37 @@ size_t ZSTD_rust_prepareCCtxForSimpleCompression(void* cctx, int ZSTD_rust_compressCCtxStrategy(size_t srcSize, int compressionLevel); size_t ZSTD_rust_resetCCtxForSimpleCompressionSession(void* cctx); void ZSTD_rust_markSimpleCompression2Complete(void* cctx); +typedef size_t (*ZSTD_rust_compress2Reset_f)(void* context); +typedef void (*ZSTD_rust_compress2SetBufferModes_f)( + void* context, int inBufferMode, int outBufferMode); +typedef size_t (*ZSTD_rust_compress2StreamEnd_f)( + void* context, void* dst, size_t dstCapacity, size_t* dstPos, + const void* src, size_t srcSize, size_t* srcPos); +typedef struct { + void* callbackContext; + ZSTD_rust_compress2Reset_f resetSession; + ZSTD_rust_compress2SetBufferModes_f setBufferModes; + ZSTD_rust_compress2StreamEnd_f compressStreamEnd; + int originalInBufferMode; + int originalOutBufferMode; +} ZSTD_rust_compress2State; +size_t ZSTD_rust_compress2(const ZSTD_rust_compress2State* state, + void* dst, size_t dstCapacity, + const void* src, size_t srcSize); +typedef char ZSTD_rust_compress2_state_layout[ + (offsetof(ZSTD_rust_compress2State, callbackContext) == 0 + && offsetof(ZSTD_rust_compress2State, resetSession) == sizeof(void*) + && offsetof(ZSTD_rust_compress2State, setBufferModes) + == 2 * sizeof(void*) + && offsetof(ZSTD_rust_compress2State, compressStreamEnd) + == 3 * sizeof(void*) + && offsetof(ZSTD_rust_compress2State, originalInBufferMode) + == 4 * sizeof(void*) + && offsetof(ZSTD_rust_compress2State, originalOutBufferMode) + == 4 * sizeof(void*) + sizeof(int) + && sizeof(ZSTD_rust_compress2State) + == 4 * sizeof(void*) + 2 * sizeof(int)) + ? 1 : -1]; size_t ZSTD_compress2_c(ZSTD_CCtx* cctx, void* dst, size_t dstCapacity, const void* src, size_t srcSize); @@ -5392,35 +5423,41 @@ size_t ZSTD_compressStream2_simpleArgs ( } } +static size_t ZSTD_rust_compress2_resetSession(void* context) +{ + return ZSTD_CCtx_reset((ZSTD_CCtx*)context, ZSTD_reset_session_only); +} + +static void ZSTD_rust_compress2_setBufferModes( + void* context, int inBufferMode, int outBufferMode) +{ + ZSTD_CCtx* const cctx = (ZSTD_CCtx*)context; + cctx->requestedParams.inBufferMode = (ZSTD_bufferMode_e)inBufferMode; + cctx->requestedParams.outBufferMode = (ZSTD_bufferMode_e)outBufferMode; +} + +static size_t ZSTD_rust_compress2_streamEnd( + void* context, void* dst, size_t dstCapacity, size_t* dstPos, + const void* src, size_t srcSize, size_t* srcPos) +{ + return ZSTD_compressStream2_simpleArgs( + (ZSTD_CCtx*)context, dst, dstCapacity, dstPos, + src, srcSize, srcPos, ZSTD_e_end); +} + size_t ZSTD_compress2_c(ZSTD_CCtx* cctx, void* dst, size_t dstCapacity, const void* src, size_t srcSize) { - ZSTD_bufferMode_e const originalInBufferMode = cctx->requestedParams.inBufferMode; - ZSTD_bufferMode_e const originalOutBufferMode = cctx->requestedParams.outBufferMode; + ZSTD_rust_compress2State state; + state.callbackContext = cctx; + state.resetSession = ZSTD_rust_compress2_resetSession; + state.setBufferModes = ZSTD_rust_compress2_setBufferModes; + state.compressStreamEnd = ZSTD_rust_compress2_streamEnd; + state.originalInBufferMode = (int)cctx->requestedParams.inBufferMode; + state.originalOutBufferMode = (int)cctx->requestedParams.outBufferMode; DEBUGLOG(4, "ZSTD_compress2 (srcSize=%u)", (unsigned)srcSize); - ZSTD_CCtx_reset(cctx, ZSTD_reset_session_only); - /* Enable stable input/output buffers. */ - cctx->requestedParams.inBufferMode = ZSTD_bm_stable; - cctx->requestedParams.outBufferMode = ZSTD_bm_stable; - { size_t oPos = 0; - size_t iPos = 0; - size_t const result = ZSTD_compressStream2_simpleArgs(cctx, - dst, dstCapacity, &oPos, - src, srcSize, &iPos, - ZSTD_e_end); - /* Reset to the original values. */ - cctx->requestedParams.inBufferMode = originalInBufferMode; - cctx->requestedParams.outBufferMode = originalOutBufferMode; - - FORWARD_IF_ERROR(result, "ZSTD_compressStream2_simpleArgs failed"); - if (result != 0) { /* compression not completed, due to lack of output space */ - assert(oPos == dstCapacity); - RETURN_ERROR(dstSize_tooSmall, ""); - } - assert(iPos == srcSize); /* all input is expected consumed */ - return oPos; - } + return ZSTD_rust_compress2(&state, dst, dstCapacity, src, srcSize); } /* The explicit-delimiter adapter is also used by the external sequence diff --git a/rust/src/zstd_compress.rs b/rust/src/zstd_compress.rs index 63dd42494..64efc9745 100644 --- a/rust/src/zstd_compress.rs +++ b/rust/src/zstd_compress.rs @@ -586,6 +586,115 @@ pub unsafe extern "C" fn ZSTD_rust_compressContinue( } } +type Compress2ResetFn = unsafe extern "C" fn(*mut c_void) -> usize; +type Compress2SetBufferModesFn = unsafe extern "C" fn(*mut c_void, c_int, c_int); +type Compress2StreamEndFn = unsafe extern "C" fn( + *mut c_void, + *mut c_void, + usize, + *mut usize, + *const c_void, + usize, + *mut usize, +) -> usize; + +/// Explicit projection for the C fallback behind `ZSTD_compress2`. +/// +/// Rust owns the reset/mode-switch/stream-call ordering and result policy. +/// The opaque callback context remains in C, where callbacks retain access to +/// the private `ZSTD_CCtx` layout and the simple-arguments stream adapter. +#[repr(C)] +pub struct ZSTD_rust_compress2State { + callback_context: *mut c_void, + reset_session: Compress2ResetFn, + set_buffer_modes: Compress2SetBufferModesFn, + compress_stream_end: Compress2StreamEndFn, + original_in_buffer_mode: c_int, + original_out_buffer_mode: c_int, +} + +const _: () = { + assert!(offset_of!(ZSTD_rust_compress2State, callback_context) == 0); + assert!(offset_of!(ZSTD_rust_compress2State, reset_session) == size_of::()); + assert!(offset_of!(ZSTD_rust_compress2State, set_buffer_modes) == 2 * size_of::()); + assert!(offset_of!(ZSTD_rust_compress2State, compress_stream_end) == 3 * size_of::()); + assert!( + offset_of!(ZSTD_rust_compress2State, original_in_buffer_mode) == 4 * size_of::() + ); + assert!( + offset_of!(ZSTD_rust_compress2State, original_out_buffer_mode) + == 4 * size_of::() + size_of::() + ); + assert!( + size_of::() == 4 * size_of::() + 2 * size_of::() + ); +}; + +unsafe fn compress2_body_with( + state: &ZSTD_rust_compress2State, + dst: *mut c_void, + dst_capacity: usize, + src: *const c_void, + src_size: usize, +) -> usize { + let reset_result = unsafe { (state.reset_session)(state.callback_context) }; + unsafe { + (state.set_buffer_modes)(state.callback_context, ZSTD_BM_STABLE, ZSTD_BM_STABLE); + } + + let mut output_pos = 0; + let mut input_pos = 0; + let result = if ERR_isError(reset_result) { + reset_result + } else { + unsafe { + (state.compress_stream_end)( + state.callback_context, + dst, + dst_capacity, + &mut output_pos, + src, + src_size, + &mut input_pos, + ) + } + }; + + unsafe { + (state.set_buffer_modes)( + state.callback_context, + state.original_in_buffer_mode, + state.original_out_buffer_mode, + ); + } + + if ERR_isError(result) { + return result; + } + if result != 0 { + debug_assert_eq!(output_pos, dst_capacity); + return ERROR(ZstdErrorCode::DstSizeTooSmall); + } + debug_assert_eq!(input_pos, src_size); + output_pos +} + +/// Drive the `ZSTD_compress2_c()` fallback without crossing the C context +/// layout. +#[no_mangle] +pub unsafe extern "C" fn ZSTD_rust_compress2( + state: *const ZSTD_rust_compress2State, + dst: *mut c_void, + dst_capacity: usize, + src: *const c_void, + src_size: usize, +) -> usize { + if state.is_null() { + return ERROR(ZstdErrorCode::Generic); + } + unsafe { compress2_body_with(&*state, dst, dst_capacity, src, src_size) } +} + 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; @@ -4903,6 +5012,135 @@ mod tests { const ZSTD_BTOPT: c_int = 7; const ZSTD_BTULTRA2: c_int = 9; + #[derive(Default)] + struct Compress2TestContext { + events: Vec<&'static str>, + in_buffer_mode: c_int, + out_buffer_mode: c_int, + reset_result: usize, + stream_result: usize, + output_pos: usize, + input_pos: usize, + } + + unsafe fn compress2_test_context(context: *mut c_void) -> &'static mut Compress2TestContext { + unsafe { &mut *context.cast::() } + } + + unsafe extern "C" fn compress2_test_reset(context: *mut c_void) -> usize { + let context = unsafe { compress2_test_context(context) }; + context.events.push("reset"); + context.reset_result + } + + unsafe extern "C" fn compress2_test_set_buffer_modes( + context: *mut c_void, + in_buffer_mode: c_int, + out_buffer_mode: c_int, + ) { + let context = unsafe { compress2_test_context(context) }; + context.events.push( + if in_buffer_mode == ZSTD_BM_STABLE && out_buffer_mode == ZSTD_BM_STABLE { + "stable" + } else { + "restore" + }, + ); + context.in_buffer_mode = in_buffer_mode; + context.out_buffer_mode = out_buffer_mode; + } + + unsafe extern "C" fn compress2_test_stream_end( + context: *mut c_void, + _dst: *mut c_void, + _dst_capacity: usize, + dst_pos: *mut usize, + _src: *const c_void, + _src_size: usize, + src_pos: *mut usize, + ) -> usize { + let context = unsafe { compress2_test_context(context) }; + context.events.push("stream-end"); + unsafe { + *dst_pos = context.output_pos; + *src_pos = context.input_pos; + } + context.stream_result + } + + fn compress2_test_state( + context: &mut Compress2TestContext, + original_in_buffer_mode: c_int, + original_out_buffer_mode: c_int, + ) -> ZSTD_rust_compress2State { + ZSTD_rust_compress2State { + callback_context: (context as *mut Compress2TestContext).cast(), + reset_session: compress2_test_reset, + set_buffer_modes: compress2_test_set_buffer_modes, + compress_stream_end: compress2_test_stream_end, + original_in_buffer_mode, + original_out_buffer_mode, + } + } + + #[test] + fn compress2_fallback_restores_modes_after_codec_error() { + let mut context = Compress2TestContext { + stream_result: ERROR(ZstdErrorCode::MemoryAllocation), + output_pos: 4, + input_pos: 3, + in_buffer_mode: 7, + out_buffer_mode: 8, + ..Compress2TestContext::default() + }; + let state = compress2_test_state(&mut context, 7, 8); + let src = [0u8; 3]; + let mut dst = [0u8; 4]; + + let result = unsafe { + ZSTD_rust_compress2( + &state, + dst.as_mut_ptr().cast(), + dst.len(), + src.as_ptr().cast(), + src.len(), + ) + }; + + assert_eq!(result, ERROR(ZstdErrorCode::MemoryAllocation)); + assert_eq!(context.events, ["reset", "stable", "stream-end", "restore"]); + assert_eq!((context.in_buffer_mode, context.out_buffer_mode), (7, 8)); + } + + #[test] + fn compress2_fallback_maps_remaining_output_to_dst_size_too_small() { + let mut context = Compress2TestContext { + stream_result: 1, + output_pos: 4, + input_pos: 3, + in_buffer_mode: 7, + out_buffer_mode: 8, + ..Compress2TestContext::default() + }; + let state = compress2_test_state(&mut context, 7, 8); + let src = [0u8; 3]; + let mut dst = [0u8; 4]; + + let result = unsafe { + ZSTD_rust_compress2( + &state, + dst.as_mut_ptr().cast(), + dst.len(), + src.as_ptr().cast(), + src.len(), + ) + }; + + assert_eq!(result, ERROR(ZstdErrorCode::DstSizeTooSmall)); + assert_eq!(context.events, ["reset", "stable", "stream-end", "restore"]); + assert_eq!((context.in_buffer_mode, context.out_buffer_mode), (7, 8)); + } + #[derive(Default)] struct CompressStreamInitTestContext { events: Vec<&'static str>,