diff --git a/lib/compress/zstdmt_compress.c b/lib/compress/zstdmt_compress.c index cbfbb2e95..5fea6ff95 100644 --- a/lib/compress/zstdmt_compress.c +++ b/lib/compress/zstdmt_compress.c @@ -116,6 +116,18 @@ void ZSTDMT_rust_buffer_pool_release(ZSTDMT_RustBufferPool* pool, ZSTDMT_RustBuffer ZSTDMT_rust_buffer_pool_resize(ZSTDMT_RustBufferPool* pool, ZSTDMT_RustBuffer buffer); +typedef struct { + size_t error; + size_t lastBlockSize; +} ZSTDMT_chunkProcessResult; + +typedef void (*ZSTDMT_chunkProgressFn)(void* opaque, size_t cSize, size_t consumed); + +ZSTDMT_chunkProcessResult ZSTDMT_rust_compressJobChunks( + ZSTD_CCtx* cctx, const void* src, size_t srcSize, + void* dst, size_t dstCapacity, size_t chunkSize, unsigned lastJob, + void* progressContext, ZSTDMT_chunkProgressFn progressCallback); + typedef struct { rawSeq* seq; size_t pos; @@ -660,6 +672,18 @@ typedef struct { unsigned frameChecksumNeeded; /* used only by mtctx */ } ZSTDMT_jobDescription; +static void ZSTDMT_compressionJobProgress(void* opaque, size_t cSize, size_t consumed) +{ + ZSTDMT_jobDescription* const job = (ZSTDMT_jobDescription*)opaque; + ZSTD_PTHREAD_MUTEX_LOCK(&job->job_mutex); + job->cSize += cSize; + job->consumed = consumed; + DEBUGLOG(5, "ZSTDMT_compressionJob: compress new block : cSize==%u bytes (total: %u)", + (U32)cSize, (U32)job->cSize); + ZSTD_pthread_cond_signal(&job->job_cond); /* warns some more data is ready to be flushed */ + ZSTD_pthread_mutex_unlock(&job->job_mutex); +} + #define JOB_ERROR(e) \ do { \ ZSTD_PTHREAD_MUTEX_LOCK(&job->job_mutex); \ @@ -739,40 +763,19 @@ static void ZSTDMT_compressionJob(void* jobDescription) /* compress the entire job by smaller chunks, for better granularity */ { size_t const chunkSize = 4*ZSTD_BLOCKSIZE_MAX; int const nbChunks = (int)((job->src.size + (chunkSize-1)) / chunkSize); - const BYTE* ip = (const BYTE*) job->src.start; - BYTE* const ostart = (BYTE*)dstBuff.start; - BYTE* op = ostart; - BYTE* oend = op + dstBuff.capacity; - int chunkNb; if (sizeof(size_t) > sizeof(int)) assert(job->src.size < ((size_t)INT_MAX) * chunkSize); /* check overflow */ DEBUGLOG(5, "ZSTDMT_compressionJob: compress %u bytes in %i blocks", (U32)job->src.size, nbChunks); assert(job->cSize == 0); - for (chunkNb = 1; chunkNb < nbChunks; chunkNb++) { - size_t const cSize = ZSTD_compressContinue_public(cctx, op, oend-op, ip, chunkSize); - if (ZSTD_isError(cSize)) JOB_ERROR(cSize); - ip += chunkSize; - op += cSize; assert(op < oend); - /* stats */ - ZSTD_PTHREAD_MUTEX_LOCK(&job->job_mutex); - job->cSize += cSize; - job->consumed = chunkSize * chunkNb; - DEBUGLOG(5, "ZSTDMT_compressionJob: compress new block : cSize==%u bytes (total: %u)", - (U32)cSize, (U32)job->cSize); - ZSTD_pthread_cond_signal(&job->job_cond); /* warns some more data is ready to be flushed */ - ZSTD_pthread_mutex_unlock(&job->job_mutex); - } - /* last block */ assert(chunkSize > 0); assert((chunkSize & (chunkSize - 1)) == 0); /* chunkSize must be power of 2 for mask==(chunkSize-1) to work */ - if ((nbChunks > 0) | job->lastJob /*must output a "last block" flag*/ ) { - size_t const lastBlockSize1 = job->src.size & (chunkSize-1); - size_t const lastBlockSize = ((lastBlockSize1==0) & (job->src.size>=chunkSize)) ? chunkSize : lastBlockSize1; - size_t const cSize = (job->lastJob) ? - ZSTD_compressEnd_public(cctx, op, oend-op, ip, lastBlockSize) : - ZSTD_compressContinue_public(cctx, op, oend-op, ip, lastBlockSize); - if (ZSTD_isError(cSize)) JOB_ERROR(cSize); - lastCBlockSize = cSize; - } } + { ZSTDMT_chunkProcessResult const result = ZSTDMT_rust_compressJobChunks( + cctx, job->src.start, job->src.size, + dstBuff.start, dstBuff.capacity, chunkSize, job->lastJob, + job, ZSTDMT_compressionJobProgress); + if (ZSTD_isError(result.error)) JOB_ERROR(result.error); + lastCBlockSize = result.lastBlockSize; + } + } if (!job->firstJob) { /* Double check that we don't have an ext-dict, because then our * repcode invalidation doesn't work. diff --git a/rust/src/zstdmt_compress.rs b/rust/src/zstdmt_compress.rs index 4659f4963..b8db02830 100644 --- a/rust/src/zstdmt_compress.rs +++ b/rust/src/zstdmt_compress.rs @@ -19,7 +19,7 @@ use std::os::raw::{c_int, c_uint, c_void}; use std::ptr; use std::sync::Mutex; -use crate::errors::{ZstdErrorCode, ERROR}; +use crate::errors::{ERR_isError, ZstdErrorCode, ERROR}; use crate::zstd_compress::ZSTD_frameProgression; const ZSTDMT_JOBLOG_MAX: c_uint = if mem::size_of::() == 4 { 29 } else { 30 }; @@ -45,6 +45,146 @@ const PRIME8_BYTES: u64 = 0xCF1B_BCDC_B7A5_6463; const ROLL_HASH_CHAR_OFFSET: u64 = 10; const _: () = assert!(RSYNC_MIN_BLOCK_SIZE >= RSYNC_LENGTH); +type ZstdMtCompressFn = + unsafe extern "C" fn(*mut c_void, *mut c_void, usize, *const c_void, usize) -> usize; +pub type ZSTDMT_chunkProgressFn = unsafe extern "C" fn(*mut c_void, usize, usize); + +#[repr(C)] +#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)] +pub struct ZSTDMT_chunkProcessResult { + pub error: usize, + pub lastBlockSize: usize, +} + +#[cfg(not(test))] +unsafe extern "C" { + fn ZSTD_compressContinue_public( + cctx: *mut c_void, + dst: *mut c_void, + dst_capacity: usize, + src: *const c_void, + src_size: usize, + ) -> usize; + fn ZSTD_compressEnd_public( + cctx: *mut c_void, + dst: *mut c_void, + dst_capacity: usize, + src: *const c_void, + src_size: usize, + ) -> usize; +} + +#[inline] +unsafe fn compress_job_chunks_with( + cctx: *mut c_void, + src: *const c_void, + src_size: usize, + dst: *mut c_void, + dst_capacity: usize, + chunk_size: usize, + last_job: c_uint, + progress_context: *mut c_void, + progress_callback: Option, + compress_continue: ZstdMtCompressFn, + compress_end: ZstdMtCompressFn, +) -> ZSTDMT_chunkProcessResult { + debug_assert!(chunk_size > 0); + debug_assert!(chunk_size.is_power_of_two()); + + let nb_chunks = src_size.div_ceil(chunk_size); + let mut input = src.cast::(); + let mut produced = 0usize; + + for chunk_number in 1..nb_chunks { + let c_size = unsafe { + compress_continue( + cctx, + dst.cast::().wrapping_add(produced).cast(), + dst_capacity.wrapping_sub(produced), + input.cast(), + chunk_size, + ) + }; + if ERR_isError(c_size) { + return ZSTDMT_chunkProcessResult { + error: c_size, + lastBlockSize: 0, + }; + } + input = input.wrapping_add(chunk_size); + produced = produced.wrapping_add(c_size); + debug_assert!(produced < dst_capacity); + + if let Some(callback) = progress_callback { + unsafe { callback(progress_context, c_size, chunk_size * chunk_number) }; + } + } + + if nb_chunks > 0 || last_job != 0 { + let last_block_size1 = src_size & (chunk_size - 1); + let last_block_size = if last_block_size1 == 0 && src_size >= chunk_size { + chunk_size + } else { + last_block_size1 + }; + let c_size = unsafe { + let dst = dst.cast::().wrapping_add(produced).cast(); + let dst_capacity = dst_capacity.wrapping_sub(produced); + let compress = if last_job != 0 { + compress_end + } else { + compress_continue + }; + compress(cctx, dst, dst_capacity, input.cast(), last_block_size) + }; + if ERR_isError(c_size) { + return ZSTDMT_chunkProcessResult { + error: c_size, + lastBlockSize: 0, + }; + } + return ZSTDMT_chunkProcessResult { + error: 0, + lastBlockSize: c_size, + }; + } + + ZSTDMT_chunkProcessResult { + error: 0, + lastBlockSize: 0, + } +} + +#[cfg(not(test))] +#[no_mangle] +pub unsafe extern "C" fn ZSTDMT_rust_compressJobChunks( + cctx: *mut c_void, + src: *const c_void, + src_size: usize, + dst: *mut c_void, + dst_capacity: usize, + chunk_size: usize, + last_job: c_uint, + progress_context: *mut c_void, + progress_callback: Option, +) -> ZSTDMT_chunkProcessResult { + unsafe { + compress_job_chunks_with( + cctx, + src, + src_size, + dst, + dst_capacity, + chunk_size, + last_job, + progress_context, + progress_callback, + ZSTD_compressContinue_public, + ZSTD_compressEnd_public, + ) + } +} + #[inline] fn cycle_log(chain_log: c_uint, strategy: c_int) -> c_uint { chain_log.wrapping_sub((strategy >= ZSTD_BTLAZY2) as c_uint) @@ -1225,6 +1365,201 @@ mod tests { static JOB_TABLE_FAIL_INIT: AtomicBool = AtomicBool::new(false); static JOB_TABLE_TEST_LOCK: Mutex<()> = Mutex::new(()); + struct MockChunkCompressor { + continue_inputs: Vec, + end_inputs: Vec, + continue_results: Vec, + end_results: Vec, + } + + #[derive(Default)] + struct MockChunkProgress { + calls: Vec<(usize, usize)>, + } + + unsafe extern "C" fn mock_compress_continue( + cctx: *mut c_void, + _dst: *mut c_void, + _dst_capacity: usize, + _src: *const c_void, + src_size: usize, + ) -> usize { + let compressor = unsafe { &mut *cctx.cast::() }; + let call_number = compressor.continue_inputs.len(); + compressor.continue_inputs.push(src_size); + compressor + .continue_results + .get(call_number) + .copied() + .unwrap_or(0) + } + + unsafe extern "C" fn mock_compress_end( + cctx: *mut c_void, + _dst: *mut c_void, + _dst_capacity: usize, + _src: *const c_void, + src_size: usize, + ) -> usize { + let compressor = unsafe { &mut *cctx.cast::() }; + let call_number = compressor.end_inputs.len(); + compressor.end_inputs.push(src_size); + compressor + .end_results + .get(call_number) + .copied() + .unwrap_or(0) + } + + unsafe extern "C" fn mock_chunk_progress(context: *mut c_void, c_size: usize, consumed: usize) { + let progress = unsafe { &mut *context.cast::() }; + progress.calls.push((c_size, consumed)); + } + + fn run_mock_chunk_loop( + compressor: &mut MockChunkCompressor, + progress: &mut MockChunkProgress, + src_size: usize, + chunk_size: usize, + last_job: c_uint, + ) -> ZSTDMT_chunkProcessResult { + let src = [0u8; 32]; + let mut dst = [0u8; 64]; + unsafe { + compress_job_chunks_with( + (compressor as *mut MockChunkCompressor).cast(), + src.as_ptr().cast(), + src_size, + dst.as_mut_ptr().cast(), + dst.len(), + chunk_size, + last_job, + (progress as *mut MockChunkProgress).cast(), + Some(mock_chunk_progress), + mock_compress_continue, + mock_compress_end, + ) + } + } + + #[test] + fn chunk_loop_handles_empty_jobs_without_compression() { + let mut compressor = MockChunkCompressor { + continue_inputs: Vec::new(), + end_inputs: Vec::new(), + continue_results: Vec::new(), + end_results: vec![9], + }; + let mut progress = MockChunkProgress::default(); + + let result = run_mock_chunk_loop(&mut compressor, &mut progress, 0, 8, 0); + assert_eq!(result, ZSTDMT_chunkProcessResult::default()); + assert!(compressor.continue_inputs.is_empty()); + assert!(compressor.end_inputs.is_empty()); + assert!(progress.calls.is_empty()); + + let result = run_mock_chunk_loop(&mut compressor, &mut progress, 0, 8, 1); + assert_eq!(result.error, 0); + assert_eq!(result.lastBlockSize, 9); + assert!(compressor.continue_inputs.is_empty()); + assert_eq!(compressor.end_inputs, [0]); + assert!(progress.calls.is_empty()); + } + + #[test] + fn chunk_loop_handles_one_exact_chunk_and_last_job_dispatch() { + let mut compressor = MockChunkCompressor { + continue_inputs: Vec::new(), + end_inputs: Vec::new(), + continue_results: vec![4], + end_results: vec![6], + }; + let mut progress = MockChunkProgress::default(); + + let result = run_mock_chunk_loop(&mut compressor, &mut progress, 8, 8, 0); + assert_eq!(result.lastBlockSize, 4); + assert_eq!(compressor.continue_inputs, [8]); + assert!(compressor.end_inputs.is_empty()); + assert!(progress.calls.is_empty()); + + let result = run_mock_chunk_loop(&mut compressor, &mut progress, 8, 8, 1); + assert_eq!(result.lastBlockSize, 6); + assert_eq!(compressor.end_inputs, [8]); + assert!(progress.calls.is_empty()); + } + + #[test] + fn chunk_loop_handles_multiple_exact_chunks_and_orders_progress() { + let mut compressor = MockChunkCompressor { + continue_inputs: Vec::new(), + end_inputs: Vec::new(), + continue_results: vec![2, 3], + end_results: Vec::new(), + }; + let mut progress = MockChunkProgress::default(); + + let result = run_mock_chunk_loop(&mut compressor, &mut progress, 16, 8, 0); + assert_eq!(result.lastBlockSize, 3); + assert_eq!(compressor.continue_inputs, [8, 8]); + assert!(compressor.end_inputs.is_empty()); + assert_eq!(progress.calls, [(2, 8)]); + } + + #[test] + fn chunk_loop_handles_partial_tail_and_final_dispatch() { + let mut compressor = MockChunkCompressor { + continue_inputs: Vec::new(), + end_inputs: Vec::new(), + continue_results: vec![2], + end_results: vec![5], + }; + let mut progress = MockChunkProgress::default(); + + let result = run_mock_chunk_loop(&mut compressor, &mut progress, 11, 8, 1); + assert_eq!(result.lastBlockSize, 5); + assert_eq!(compressor.continue_inputs, [8]); + assert_eq!(compressor.end_inputs, [3]); + assert_eq!(progress.calls, [(2, 8)]); + } + + #[test] + fn chunk_loop_propagates_intermediate_errors_after_prior_progress() { + let error = ERROR(ZstdErrorCode::DstSizeTooSmall); + let mut compressor = MockChunkCompressor { + continue_inputs: Vec::new(), + end_inputs: Vec::new(), + continue_results: vec![2, error], + end_results: vec![7], + }; + let mut progress = MockChunkProgress::default(); + + let result = run_mock_chunk_loop(&mut compressor, &mut progress, 24, 8, 1); + assert_eq!(result.error, error); + assert_eq!(result.lastBlockSize, 0); + assert_eq!(compressor.continue_inputs, [8, 8]); + assert!(compressor.end_inputs.is_empty()); + assert_eq!(progress.calls, [(2, 8)]); + } + + #[test] + fn chunk_loop_propagates_final_errors_without_progress_callback() { + let error = ERROR(ZstdErrorCode::Generic); + let mut compressor = MockChunkCompressor { + continue_inputs: Vec::new(), + end_inputs: Vec::new(), + continue_results: Vec::new(), + end_results: vec![error], + }; + let mut progress = MockChunkProgress::default(); + + let result = run_mock_chunk_loop(&mut compressor, &mut progress, 3, 8, 1); + assert_eq!(result.error, error); + assert_eq!(result.lastBlockSize, 0); + assert!(compressor.continue_inputs.is_empty()); + assert_eq!(compressor.end_inputs, [3]); + assert!(progress.calls.is_empty()); + } + unsafe extern "C" fn probe_job_table_init( _job_table: *mut c_void, _nb_jobs: c_uint,