diff --git a/lib/compress/zstd_compress.c b/lib/compress/zstd_compress.c index ff6cab2e4..e01afae33 100644 --- a/lib/compress/zstd_compress.c +++ b/lib/compress/zstd_compress.c @@ -41,6 +41,11 @@ size_t ZSTD_rust_writeEpilogue(void* dst, size_t dstCapacity, int* stage, int noDictIDFlag, int checksumFlag, int contentSizeFlag, int format, U32 windowLog, U32 checksum); +int ZSTD_rust_updateFrameProgression(unsigned long long* consumedSrcSize, + unsigned long long* producedCSize, + unsigned long long pledgedSrcSizePlusOne, + size_t srcSize, size_t cSize, + size_t fhSize); size_t ZSTD_rust_resetCCtxForSimpleCompression(void* cctx); size_t ZSTD_rust_prepareCCtxForSimpleCompression(void* cctx, size_t srcSize, @@ -3403,14 +3408,16 @@ static size_t ZSTD_compressContinue_internal (ZSTD_CCtx* cctx, { size_t const cSize = frame ? ZSTD_compress_frameChunk (cctx, dst, dstCapacity, src, srcSize, lastFrameChunk) : ZSTD_compressBlock_internal (cctx, dst, dstCapacity, src, srcSize, 0 /* frame */); + int srcSizeWrong; FORWARD_IF_ERROR(cSize, "%s", frame ? "ZSTD_compress_frameChunk failed" : "ZSTD_compressBlock_internal failed"); - cctx->consumedSrcSize += srcSize; - cctx->producedCSize += (cSize + fhSize); + srcSizeWrong = ZSTD_rust_updateFrameProgression( + &cctx->consumedSrcSize, &cctx->producedCSize, + cctx->pledgedSrcSizePlusOne, srcSize, cSize, fhSize); 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); RETURN_ERROR_IF( - cctx->consumedSrcSize+1 > cctx->pledgedSrcSizePlusOne, + srcSizeWrong, srcSize_wrong, "error : pledgedSrcSize = %u, while realSrcSize >= %u", (unsigned)cctx->pledgedSrcSizePlusOne-1, diff --git a/rust/src/zstd_compress.rs b/rust/src/zstd_compress.rs index 59d7892fa..c9f46f571 100644 --- a/rust/src/zstd_compress.rs +++ b/rust/src/zstd_compress.rs @@ -97,6 +97,46 @@ pub struct ZSTD_outBuffer { * supported pointer widths. */ const TMP_WORKSPACE_SIZE: usize = 16 << 10; +#[inline] +fn update_frame_progression( + consumed_src_size: &mut u64, + produced_c_size: &mut u64, + pledged_src_size_plus_one: u64, + src_size: usize, + c_size: usize, + frame_header_size: usize, +) -> bool { + *consumed_src_size = consumed_src_size.wrapping_add(src_size as u64); + *produced_c_size = produced_c_size.wrapping_add(c_size.wrapping_add(frame_header_size) as u64); + + pledged_src_size_plus_one != 0 && consumed_src_size.wrapping_add(1) > pledged_src_size_plus_one +} + +/// Update C-owned frame counters after successful compression. +/// +/// A zero pledge means that the source size is unknown. C keeps the +/// diagnostic that accompanies a nonzero overrun result. +#[no_mangle] +pub unsafe extern "C" fn ZSTD_rust_updateFrameProgression( + consumed_src_size: *mut u64, + produced_c_size: *mut u64, + pledged_src_size_plus_one: u64, + src_size: usize, + c_size: usize, + frame_header_size: usize, +) -> c_int { + let consumed_src_size = unsafe { &mut *consumed_src_size }; + let produced_c_size = unsafe { &mut *produced_c_size }; + update_frame_progression( + consumed_src_size, + produced_c_size, + pledged_src_size_plus_one, + src_size, + c_size, + frame_header_size, + ) as c_int +} + #[inline] fn bitmix(mut val: u64, len: u64) -> u64 { val ^= val.rotate_right(49) ^ val.rotate_right(24); @@ -712,6 +752,54 @@ mod tests { assert_eq!(ZSTD_rust_dictTooBig(ZSTD_CHUNKSIZE_MAX + 1), 1); } + #[test] + fn frame_progress_unknown_pledge_updates_counters() { + let mut consumed = 7; + let mut produced = 11; + let result = + unsafe { ZSTD_rust_updateFrameProgression(&mut consumed, &mut produced, 0, 5, 13, 2) }; + + assert_eq!(result, 0); + assert_eq!(consumed, 12); + assert_eq!(produced, 26); + } + + #[test] + fn frame_progress_exact_pledge_is_accepted() { + let mut consumed = 7; + let mut produced = 11; + let result = + unsafe { ZSTD_rust_updateFrameProgression(&mut consumed, &mut produced, 13, 5, 13, 2) }; + + assert_eq!(result, 0); + assert_eq!(consumed, 12); + assert_eq!(produced, 26); + } + + #[test] + fn frame_progress_one_byte_overrun_is_reported() { + let mut consumed = 7; + let mut produced = 11; + let result = + unsafe { ZSTD_rust_updateFrameProgression(&mut consumed, &mut produced, 13, 6, 13, 2) }; + + assert_eq!(result, 1); + assert_eq!(consumed, 13); + } + + #[test] + fn frame_progress_counters_update_before_overrun_result() { + let mut consumed = 100; + let mut produced = 200; + let result = unsafe { + ZSTD_rust_updateFrameProgression(&mut consumed, &mut produced, 106, 7, 17, 3) + }; + + assert_eq!(result, 1); + assert_eq!(consumed, 107); + assert_eq!(produced, 220); + } + #[test] fn public_one_shot_abi_is_c_compatible() { let entry: unsafe extern "C" fn(*mut c_void, usize, *const c_void, usize, c_int) -> usize =