diff --git a/lib/compress/zstd_compress.c b/lib/compress/zstd_compress.c index 5f437ba56..259a2dbbe 100644 --- a/lib/compress/zstd_compress.c +++ b/lib/compress/zstd_compress.c @@ -232,6 +232,50 @@ typedef char ZSTD_rust_compress_continue_state_layout[ == (sizeof(void*) == 8 ? 96 : 52)) ? 1 : -1]; +/* Rust owns the end-of-frame orchestration. The callbacks retain the + * private CCtx-dependent continue, epilogue, and trace operations in C. */ +typedef size_t (*ZSTD_rust_compressEndContinue_f)(void* context, + void* dst, + size_t dstCapacity, + const void* src, + size_t srcSize, + U32 frame, + U32 lastFrameChunk); +typedef size_t (*ZSTD_rust_compressEndEpilogue_f)(void* context, + void* dst, + size_t dstCapacity); +typedef void (*ZSTD_rust_compressEndTrace_f)(void* context, + size_t extraCSize); +typedef struct { + void* callbackContext; + ZSTD_rust_compressEndContinue_f compressContinue; + ZSTD_rust_compressEndEpilogue_f writeEpilogue; + ZSTD_rust_compressEndTrace_f trace; + unsigned long long* consumedSrcSize; + U64 pledgedSrcSizePlusOne; + int contentSizeFlag; +} ZSTD_rust_compressEndState; +size_t ZSTD_rust_compressEnd(const ZSTD_rust_compressEndState* state, + void* dst, size_t dstCapacity, + const void* src, size_t srcSize); +typedef char ZSTD_rust_compress_end_state_layout[ + (offsetof(ZSTD_rust_compressEndState, callbackContext) == 0 + && offsetof(ZSTD_rust_compressEndState, compressContinue) + == sizeof(void*) + && offsetof(ZSTD_rust_compressEndState, writeEpilogue) + == 2 * sizeof(void*) + && offsetof(ZSTD_rust_compressEndState, trace) + == 3 * sizeof(void*) + && offsetof(ZSTD_rust_compressEndState, consumedSrcSize) + == 4 * sizeof(void*) + && offsetof(ZSTD_rust_compressEndState, pledgedSrcSizePlusOne) + == 5 * sizeof(void*) + && offsetof(ZSTD_rust_compressEndState, contentSizeFlag) + == 5 * sizeof(void*) + sizeof(U64) + && sizeof(ZSTD_rust_compressEndState) + == (sizeof(void*) == 8 ? 56 : 32)) + ? 1 : -1]; + /* Rust owns the single-threaded buffered/stable stream state machine. The * projection contains only stream bookkeeping and callback slots; operations * which still need the private CCtx layout remain C callbacks. */ @@ -4045,31 +4089,42 @@ void ZSTD_CCtx_trace(ZSTD_CCtx* cctx, size_t extraCSize) #endif } +static size_t ZSTD_rust_compressEnd_continue( + void* context, void* dst, size_t dstCapacity, + const void* src, size_t srcSize, U32 frame, U32 lastFrameChunk) +{ + ZSTD_CCtx* const cctx = (ZSTD_CCtx*)context; + return ZSTD_compressContinue_dispatch( + cctx, dst, dstCapacity, src, srcSize, + frame, lastFrameChunk, cctx->blockSizeMax, + 0 /* block size already selected */); +} + +static size_t ZSTD_rust_compressEnd_writeEpilogue( + void* context, void* dst, size_t dstCapacity) +{ + return ZSTD_writeEpilogue((ZSTD_CCtx*)context, dst, dstCapacity); +} + +static void ZSTD_rust_compressEnd_trace(void* context, size_t extraCSize) +{ + ZSTD_CCtx_trace((ZSTD_CCtx*)context, extraCSize); +} + size_t ZSTD_compressEnd_public(ZSTD_CCtx* cctx, void* dst, size_t dstCapacity, const void* src, size_t srcSize) { - size_t endResult; - size_t const cSize = ZSTD_compressContinue_dispatch( - cctx, dst, dstCapacity, src, srcSize, - 1 /* frame mode */, 1 /* last chunk */, - cctx->blockSizeMax, 0 /* block size already selected */); - FORWARD_IF_ERROR(cSize, "ZSTD_compressContinue failed"); - endResult = ZSTD_writeEpilogue(cctx, (char*)dst + cSize, dstCapacity-cSize); - FORWARD_IF_ERROR(endResult, "ZSTD_writeEpilogue failed"); - assert(!(cctx->appliedParams.fParams.contentSizeFlag && cctx->pledgedSrcSizePlusOne == 0)); - if (cctx->pledgedSrcSizePlusOne != 0) { /* control src size */ - ZSTD_STATIC_ASSERT(ZSTD_CONTENTSIZE_UNKNOWN == (unsigned long long)-1); - DEBUGLOG(4, "end of frame : controlling src size"); - RETURN_ERROR_IF( - cctx->pledgedSrcSizePlusOne != cctx->consumedSrcSize+1, - srcSize_wrong, - "error : pledgedSrcSize = %u, while realSrcSize = %u", - (unsigned)cctx->pledgedSrcSizePlusOne-1, - (unsigned)cctx->consumedSrcSize); - } - ZSTD_CCtx_trace(cctx, endResult); - return cSize + endResult; + ZSTD_rust_compressEndState state; + state.callbackContext = cctx; + state.compressContinue = ZSTD_rust_compressEnd_continue; + state.writeEpilogue = ZSTD_rust_compressEnd_writeEpilogue; + state.trace = ZSTD_rust_compressEnd_trace; + state.consumedSrcSize = &cctx->consumedSrcSize; + state.pledgedSrcSizePlusOne = cctx->pledgedSrcSizePlusOne; + state.contentSizeFlag = cctx->appliedParams.fParams.contentSizeFlag; + return ZSTD_rust_compressEnd( + &state, dst, dstCapacity, src, srcSize); } /* NOTE: Must just wrap ZSTD_compressEnd_public() */ diff --git a/rust/src/zstd_compress.rs b/rust/src/zstd_compress.rs index 64efc9745..f36deff69 100644 --- a/rust/src/zstd_compress.rs +++ b/rust/src/zstd_compress.rs @@ -586,6 +586,114 @@ pub unsafe extern "C" fn ZSTD_rust_compressContinue( } } +type CompressEndContinueFn = unsafe extern "C" fn( + *mut c_void, + *mut c_void, + usize, + *const c_void, + usize, + c_uint, + c_uint, +) -> usize; +type CompressEndEpilogueFn = unsafe extern "C" fn(*mut c_void, *mut c_void, usize) -> usize; +type CompressEndTraceFn = unsafe extern "C" fn(*mut c_void, usize); + +/// Explicit projection for the public end-of-frame orchestration. +/// +/// Rust owns callback ordering, output offset/capacity accounting, and +/// pledged-size validation. The opaque callback context retains the private +/// CCtx-dependent continue, epilogue, checksum, and trace operations in C. +#[repr(C)] +pub struct ZSTD_rust_compressEndState { + callback_context: *mut c_void, + compress_continue: CompressEndContinueFn, + write_epilogue: CompressEndEpilogueFn, + trace: CompressEndTraceFn, + consumed_src_size: *const u64, + pledged_src_size_plus_one: u64, + content_size_flag: c_int, +} + +const _: () = { + assert!(offset_of!(ZSTD_rust_compressEndState, callback_context) == 0); + assert!(offset_of!(ZSTD_rust_compressEndState, compress_continue) == size_of::()); + assert!(offset_of!(ZSTD_rust_compressEndState, write_epilogue) == 2 * size_of::()); + assert!(offset_of!(ZSTD_rust_compressEndState, trace) == 3 * size_of::()); + assert!(offset_of!(ZSTD_rust_compressEndState, consumed_src_size) == 4 * size_of::()); + assert!( + offset_of!(ZSTD_rust_compressEndState, pledged_src_size_plus_one) == 5 * size_of::() + ); + assert!( + offset_of!(ZSTD_rust_compressEndState, content_size_flag) + == 5 * size_of::() + size_of::() + ); + assert!( + size_of::() == if size_of::() == 8 { 56 } else { 32 } + ); +}; + +unsafe fn compress_end_body_with( + state: &ZSTD_rust_compressEndState, + dst: *mut c_void, + dst_capacity: usize, + src: *const c_void, + src_size: usize, +) -> usize { + if state.consumed_src_size.is_null() { + return ERROR(ZstdErrorCode::Generic); + } + + let c_size = unsafe { + (state.compress_continue)( + state.callback_context, + dst, + dst_capacity, + src, + src_size, + 1, + 1, + ) + }; + if ERR_isError(c_size) { + return c_size; + } + + let end_result = unsafe { + (state.write_epilogue)( + state.callback_context, + dst.cast::().add(c_size).cast(), + dst_capacity.wrapping_sub(c_size), + ) + }; + if ERR_isError(end_result) { + return end_result; + } + + debug_assert!(!(state.content_size_flag != 0 && state.pledged_src_size_plus_one == 0)); + if state.pledged_src_size_plus_one != 0 + && state.pledged_src_size_plus_one != unsafe { (*state.consumed_src_size).wrapping_add(1) } + { + return ERROR(ZstdErrorCode::SrcSizeWrong); + } + + unsafe { (state.trace)(state.callback_context, end_result) }; + c_size.wrapping_add(end_result) +} + +#[no_mangle] +pub unsafe extern "C" fn ZSTD_rust_compressEnd( + state: *const ZSTD_rust_compressEndState, + dst: *mut c_void, + dst_capacity: usize, + src: *const c_void, + src_size: usize, +) -> usize { + if state.is_null() { + return ERROR(ZstdErrorCode::Generic); + } + unsafe { compress_end_body_with(&*state, dst, dst_capacity, src, src_size) } +} + 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( @@ -5141,6 +5249,164 @@ mod tests { assert_eq!((context.in_buffer_mode, context.out_buffer_mode), (7, 8)); } + #[derive(Default)] + struct CompressEndTestContext { + events: Vec<&'static str>, + continue_result: usize, + epilogue_result: usize, + epilogue_offset: usize, + epilogue_capacity: usize, + trace_extra: usize, + continue_frame: c_uint, + continue_last_frame_chunk: c_uint, + dst_base: usize, + } + + unsafe fn compress_end_test_context( + context: *mut c_void, + ) -> &'static mut CompressEndTestContext { + unsafe { &mut *context.cast::() } + } + + unsafe extern "C" fn compress_end_test_continue( + context: *mut c_void, + _dst: *mut c_void, + _dst_capacity: usize, + _src: *const c_void, + _src_size: usize, + frame: c_uint, + last_frame_chunk: c_uint, + ) -> usize { + let context = unsafe { compress_end_test_context(context) }; + context.events.push("continue"); + context.continue_frame = frame; + context.continue_last_frame_chunk = last_frame_chunk; + context.continue_result + } + + unsafe extern "C" fn compress_end_test_epilogue( + context: *mut c_void, + dst: *mut c_void, + dst_capacity: usize, + ) -> usize { + let context = unsafe { compress_end_test_context(context) }; + context.events.push("epilogue"); + context.epilogue_offset = (dst as usize).wrapping_sub(context.dst_base); + context.epilogue_capacity = dst_capacity; + context.epilogue_result + } + + unsafe extern "C" fn compress_end_test_trace(context: *mut c_void, extra_c_size: usize) { + let context = unsafe { compress_end_test_context(context) }; + context.events.push("trace"); + context.trace_extra = extra_c_size; + } + + fn compress_end_test_state( + context: &mut CompressEndTestContext, + consumed_src_size: &u64, + pledged_src_size_plus_one: u64, + content_size_flag: c_int, + ) -> ZSTD_rust_compressEndState { + ZSTD_rust_compressEndState { + callback_context: (context as *mut CompressEndTestContext).cast(), + compress_continue: compress_end_test_continue, + write_epilogue: compress_end_test_epilogue, + trace: compress_end_test_trace, + consumed_src_size: consumed_src_size as *const u64, + pledged_src_size_plus_one, + content_size_flag, + } + } + + #[test] + fn compress_end_preserves_callback_order_and_output_accounting() { + let mut dst = [0u8; 16]; + let mut context = CompressEndTestContext { + continue_result: 3, + epilogue_result: 5, + dst_base: dst.as_mut_ptr() as usize, + ..CompressEndTestContext::default() + }; + let consumed_src_size = 7; + let state = compress_end_test_state(&mut context, &consumed_src_size, 8, 1); + + let result = unsafe { + ZSTD_rust_compressEnd(&state, dst.as_mut_ptr().cast(), dst.len(), ptr::null(), 0) + }; + + assert_eq!(result, 8); + assert_eq!(context.events, ["continue", "epilogue", "trace"]); + assert_eq!( + (context.continue_frame, context.continue_last_frame_chunk), + (1, 1) + ); + assert_eq!( + (context.epilogue_offset, context.epilogue_capacity), + (3, 13) + ); + assert_eq!(context.trace_extra, 5); + } + + #[test] + fn compress_end_stops_before_epilogue_when_continue_fails() { + let mut dst = [0u8; 16]; + let mut context = CompressEndTestContext { + continue_result: ERROR(ZstdErrorCode::MemoryAllocation), + dst_base: dst.as_mut_ptr() as usize, + ..CompressEndTestContext::default() + }; + let consumed_src_size = 7; + let state = compress_end_test_state(&mut context, &consumed_src_size, 8, 0); + + let result = unsafe { + ZSTD_rust_compressEnd(&state, dst.as_mut_ptr().cast(), dst.len(), ptr::null(), 0) + }; + + assert_eq!(result, ERROR(ZstdErrorCode::MemoryAllocation)); + assert_eq!(context.events, ["continue"]); + } + + #[test] + fn compress_end_stops_before_validation_and_trace_when_epilogue_fails() { + let mut dst = [0u8; 16]; + let mut context = CompressEndTestContext { + continue_result: 3, + epilogue_result: ERROR(ZstdErrorCode::MemoryAllocation), + dst_base: dst.as_mut_ptr() as usize, + ..CompressEndTestContext::default() + }; + let consumed_src_size = 7; + let state = compress_end_test_state(&mut context, &consumed_src_size, 99, 1); + + let result = unsafe { + ZSTD_rust_compressEnd(&state, dst.as_mut_ptr().cast(), dst.len(), ptr::null(), 0) + }; + + assert_eq!(result, ERROR(ZstdErrorCode::MemoryAllocation)); + assert_eq!(context.events, ["continue", "epilogue"]); + } + + #[test] + fn compress_end_rejects_pledged_size_before_trace() { + let mut dst = [0u8; 16]; + let mut context = CompressEndTestContext { + continue_result: 3, + epilogue_result: 5, + dst_base: dst.as_mut_ptr() as usize, + ..CompressEndTestContext::default() + }; + let consumed_src_size = 7; + let state = compress_end_test_state(&mut context, &consumed_src_size, 99, 1); + + let result = unsafe { + ZSTD_rust_compressEnd(&state, dst.as_mut_ptr().cast(), dst.len(), ptr::null(), 0) + }; + + assert_eq!(result, ERROR(ZstdErrorCode::SrcSizeWrong)); + assert_eq!(context.events, ["continue", "epilogue"]); + } + #[derive(Default)] struct CompressStreamInitTestContext { events: Vec<&'static str>,