diff --git a/lib/compress/zstd_compress.c b/lib/compress/zstd_compress.c index 8c8e6e40d..238b98308 100644 --- a/lib/compress/zstd_compress.c +++ b/lib/compress/zstd_compress.c @@ -103,6 +103,8 @@ size_t ZSTD_rust_estimateWorkspaceSize(size_t cctxSpace, size_t tokenSpace, size_t bufferSpace, size_t externalSeqSpace); +size_t ZSTD_rust_maxEstimateCCtxSize(size_t estimate0, size_t estimate1, + size_t estimate2, size_t estimate3); ZSTD_inBuffer ZSTD_rust_inBufferForEndFlush(int inBufferMode, const void* expectedSrc, size_t expectedSize, @@ -1346,14 +1348,15 @@ size_t ZSTD_estimateCCtxSize_usingCParams(ZSTD_compressionParameters cParams) static size_t ZSTD_estimateCCtxSize_internal(int compressionLevel) { int tier = 0; - size_t largestSize = 0; + size_t estimates[4]; static const unsigned long long srcSizeTiers[4] = {16 KB, 128 KB, 256 KB, ZSTD_CONTENTSIZE_UNKNOWN}; for (; tier < 4; ++tier) { /* Choose the set of cParams for a given level across all srcSizes that give the largest cctxSize */ ZSTD_compressionParameters const cParams = ZSTD_getCParams_internal(compressionLevel, srcSizeTiers[tier], 0, ZSTD_cpm_noAttachDict); - largestSize = MAX(ZSTD_estimateCCtxSize_usingCParams(cParams), largestSize); + estimates[tier] = ZSTD_estimateCCtxSize_usingCParams(cParams); } - return largestSize; + return ZSTD_rust_maxEstimateCCtxSize( + estimates[0], estimates[1], estimates[2], estimates[3]); } size_t ZSTD_estimateCCtxSize(int compressionLevel) diff --git a/rust/src/zstd_compress.rs b/rust/src/zstd_compress.rs index 5255dd8b4..d282938af 100644 --- a/rust/src/zstd_compress.rs +++ b/rust/src/zstd_compress.rs @@ -611,6 +611,27 @@ pub extern "C" fn ZSTD_rust_estimateWorkspaceSize( ) } +#[inline] +fn max_estimate_cctx_size( + estimate0: usize, + estimate1: usize, + estimate2: usize, + estimate3: usize, +) -> usize { + estimate0.max(estimate1).max(estimate2).max(estimate3) +} + +/// Return the largest raw estimate, including any C size_t error values. +#[no_mangle] +pub extern "C" fn ZSTD_rust_maxEstimateCCtxSize( + estimate0: usize, + estimate1: usize, + estimate2: usize, + estimate3: usize, +) -> usize { + max_estimate_cctx_size(estimate0, estimate1, estimate2, estimate3) +} + #[inline] fn reduce_table_internal(table: &mut [u32], reducer_value: u32, preserve_mark: bool) { debug_assert_eq!(table.len() % ZSTD_ROWSIZE, 0); @@ -1860,6 +1881,48 @@ mod tests { ); } + #[test] + fn max_estimate_cctx_size_handles_zero_values() { + assert_eq!(max_estimate_cctx_size(0, 0, 0, 0), 0); + assert_eq!(ZSTD_rust_maxEstimateCCtxSize(0, 0, 0, 0), 0); + } + + #[test] + fn max_estimate_cctx_size_selects_the_largest_ordinary_value() { + assert_eq!(max_estimate_cctx_size(17, 42, 9, 31), 42); + assert_eq!(ZSTD_rust_maxEstimateCCtxSize(17, 42, 9, 31), 42); + } + + #[test] + fn max_estimate_cctx_size_preserves_equal_values() { + assert_eq!(max_estimate_cctx_size(42, 42, 42, 42), 42); + assert_eq!(ZSTD_rust_maxEstimateCCtxSize(42, 42, 42, 42), 42); + } + + #[test] + fn max_estimate_cctx_size_accepts_size_max() { + assert_eq!(max_estimate_cctx_size(1, usize::MAX, 3, 2), usize::MAX); + assert_eq!( + ZSTD_rust_maxEstimateCCtxSize(1, usize::MAX, 3, 2), + usize::MAX + ); + } + + #[test] + fn max_estimate_cctx_size_compares_error_like_raw_values() { + let smaller_error = ERROR(ZstdErrorCode::DstSizeTooSmall); + let larger_error = ERROR(ZstdErrorCode::ParameterUnsupported); + + assert_eq!( + max_estimate_cctx_size(128, smaller_error, 256, larger_error), + larger_error + ); + assert_eq!( + ZSTD_rust_maxEstimateCCtxSize(128, smaller_error, 256, larger_error), + larger_error + ); + } + #[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 =