diff --git a/lib/compress/zstdmt_compress.c b/lib/compress/zstdmt_compress.c index df1999e23..50d3c28aa 100644 --- a/lib/compress/zstdmt_compress.c +++ b/lib/compress/zstdmt_compress.c @@ -374,6 +374,11 @@ typedef struct { } ZSTDMT_RustSerialEnsureFinishedResult; ZSTDMT_RustSerialEnsureFinishedResult ZSTDMT_rust_serialStateEnsureFinished( unsigned nextJobID, unsigned jobID); +typedef void (*ZSTDMT_waitForJobCompleteFn)( + void* opaque, unsigned jobID, unsigned doneJobID); +unsigned ZSTDMT_rust_waitForAllJobsCompleted( + unsigned doneJobID, unsigned nextJobID, unsigned jobIDMask, + void* opaque, ZSTDMT_waitForJobCompleteFn waitForJob); typedef void (*ZSTDMT_waitForLdmLockFn)(void* opaque); typedef int (*ZSTDMT_waitForLdmOverlapFn)( void* opaque, void* bufferStart, size_t bufferCapacity); @@ -1502,19 +1507,26 @@ static void ZSTDMT_releaseAllJobResources(ZSTDMT_CCtx* mtctx) mtctx->allJobsCompleted = 1; } +static void ZSTDMT_waitForJobComplete( + void* opaque, unsigned jobID, unsigned doneJobID) +{ + ZSTDMT_CCtx* const mtctx = (ZSTDMT_CCtx*)opaque; + ZSTDMT_jobDescription* const job = &mtctx->jobs[jobID]; + (void)doneJobID; + ZSTD_PTHREAD_MUTEX_LOCK(&job->job_mutex); + while (job->consumed < job->src.size) { + DEBUGLOG(4, "waiting for jobCompleted signal from job %u", doneJobID); + ZSTD_pthread_cond_wait(&job->job_cond, &job->job_mutex); + } + ZSTD_pthread_mutex_unlock(&job->job_mutex); +} + static void ZSTDMT_waitForAllJobsCompleted(ZSTDMT_CCtx* mtctx) { DEBUGLOG(4, "ZSTDMT_waitForAllJobsCompleted"); - while (mtctx->doneJobID < mtctx->nextJobID) { - unsigned const jobID = mtctx->doneJobID & mtctx->jobIDMask; - ZSTD_PTHREAD_MUTEX_LOCK(&mtctx->jobs[jobID].job_mutex); - while (mtctx->jobs[jobID].consumed < mtctx->jobs[jobID].src.size) { - DEBUGLOG(4, "waiting for jobCompleted signal from job %u", mtctx->doneJobID); /* we want to block when waiting for data to flush */ - ZSTD_pthread_cond_wait(&mtctx->jobs[jobID].job_cond, &mtctx->jobs[jobID].job_mutex); - } - ZSTD_pthread_mutex_unlock(&mtctx->jobs[jobID].job_mutex); - mtctx->doneJobID++; - } + mtctx->doneJobID = ZSTDMT_rust_waitForAllJobsCompleted( + mtctx->doneJobID, mtctx->nextJobID, mtctx->jobIDMask, + mtctx, ZSTDMT_waitForJobComplete); } size_t ZSTDMT_freeCCtx(ZSTDMT_CCtx* mtctx) diff --git a/rust/src/zstdmt_compress.rs b/rust/src/zstdmt_compress.rs index c8d308a6e..39a3a4315 100644 --- a/rust/src/zstdmt_compress.rs +++ b/rust/src/zstdmt_compress.rs @@ -96,6 +96,7 @@ pub struct ZSTDMT_serialStateEnsureFinishedResult { pub skip: c_uint, pub nextJobID: c_uint, } +pub type ZSTDMT_waitForJobCompleteFn = unsafe extern "C" fn(*mut c_void, c_uint, c_uint); type ZSTDMT_waitForLdmLockFn = unsafe extern "C" fn(*mut c_void); type ZSTDMT_waitForLdmOverlapFn = unsafe extern "C" fn(*mut c_void, *mut c_void, usize) -> c_int; @@ -864,6 +865,27 @@ pub extern "C" fn ZSTDMT_rust_serialStateEnsureFinished( } } +/// Wait for each submitted MT job in ring order. C retains the per-job +/// mutex, condition variable, and consumed/source counters behind one +/// callback; Rust owns the ring-slot and completion progression policy. +#[no_mangle] +pub unsafe extern "C" fn ZSTDMT_rust_waitForAllJobsCompleted( + mut done_job_id: c_uint, + next_job_id: c_uint, + job_id_mask: c_uint, + opaque: *mut c_void, + wait_for_job: Option, +) -> c_uint { + let Some(wait_for_job) = wait_for_job else { + return done_job_id; + }; + while done_job_id < next_job_id { + unsafe { wait_for_job(opaque, done_job_id & job_id_mask, done_job_id) }; + done_job_id = done_job_id.wrapping_add(1); + } + done_job_id +} + /// C ABI entry point for the worker-job orchestration. C supplies callbacks /// that keep the private descriptor, pools, mutexes, and codec operations on /// the C side of this narrow projection. @@ -3541,6 +3563,40 @@ mod tests { ); } + struct WaitForAllJobsTestContext { + events: Vec<(c_uint, c_uint)>, + } + + unsafe extern "C" fn wait_for_all_jobs_test_callback( + context: *mut c_void, + job_id: c_uint, + done_job_id: c_uint, + ) { + unsafe { + (*context.cast::()) + .events + .push((job_id, done_job_id)); + } + } + + #[test] + fn wait_for_all_jobs_uses_ring_order_and_advances_to_next_job() { + let mut context = WaitForAllJobsTestContext { events: Vec::new() }; + + let done_job_id = unsafe { + ZSTDMT_rust_waitForAllJobsCompleted( + 3, + 6, + 3, + (&mut context as *mut WaitForAllJobsTestContext).cast(), + Some(wait_for_all_jobs_test_callback), + ) + }; + + assert_eq!(done_job_id, 6); + assert_eq!(context.events, vec![(3, 3), (0, 4), (1, 5)]); + } + #[test] fn compression_job_stops_on_non_first_chunk_error_and_cleans_up() { let state = Rc::new(RefCell::new(MockCompressionJob::default()));