diff --git a/lib/compress/zstdmt_compress.c b/lib/compress/zstdmt_compress.c index e13f63822..8d1de4e07 100644 --- a/lib/compress/zstdmt_compress.c +++ b/lib/compress/zstdmt_compress.c @@ -602,9 +602,13 @@ typedef char ZSTDMT_rust_raw_seq_store_layout[ ZSTDMT_RustRawSeqStore ZSTDMT_rust_bufferToSeq(ZSTDMT_RustBuffer buffer); ZSTDMT_RustBuffer ZSTDMT_rust_seqToBuffer(ZSTDMT_RustRawSeqStore seq); typedef int (*ZSTDMT_serialWaitForTurnFn)(void* opaque, unsigned jobID); -typedef void (*ZSTDMT_serialGenerateLdmFn)( +typedef void (*ZSTDMT_serialLdmWindowUpdateFn)( void* opaque, ZSTDMT_RustRawSeqStore* seqStore, const void* src, size_t srcSize); +typedef void (*ZSTDMT_serialLdmGenerateSequencesFn)( + void* opaque, ZSTDMT_RustRawSeqStore* seqStore, + const void* src, size_t srcSize); +typedef void (*ZSTDMT_serialLdmPublishWindowFn)(void* opaque); typedef void (*ZSTDMT_serialUpdateChecksumFn)( void* opaque, const void* src, size_t srcSize); typedef void (*ZSTDMT_serialAdvanceFn)(void* opaque); @@ -645,7 +649,9 @@ void ZSTDMT_rust_serialStateGenSequences( ZSTDMT_RustRawSeqStore* seqStore, const void* src, size_t srcSize, unsigned jobID, int ldmEnabled, int checksumEnabled, void* opaque, ZSTDMT_serialWaitForTurnFn waitForTurn, - ZSTDMT_serialGenerateLdmFn generateLdm, + ZSTDMT_serialLdmWindowUpdateFn updateLdmWindow, + ZSTDMT_serialLdmGenerateSequencesFn generateLdmSequences, + ZSTDMT_serialLdmPublishWindowFn publishLdmWindow, ZSTDMT_serialUpdateChecksumFn updateChecksum, ZSTDMT_serialAdvanceFn advance); typedef struct { @@ -1740,13 +1746,12 @@ static int ZSTDMT_serialState_waitForTurn(void* opaque, unsigned jobID) return ZSTDMT_rust_serialStateWaitForTurn(&state); } -static void ZSTDMT_serialState_generateLdm( +static void ZSTDMT_serialState_updateLdmWindow( void* opaque, ZSTDMT_RustRawSeqStore* seqStore, const void* src, size_t srcSize) { SerialState* const serialState = (SerialState*)opaque; RawSeqStore_t* const cSeqStore = (RawSeqStore_t*)seqStore; - size_t error; DEBUGLOG(6, "ZSTDMT_serialState_genSequences: LDM update"); assert(cSeqStore->seq != NULL && cSeqStore->pos == 0 && @@ -1754,11 +1759,27 @@ static void ZSTDMT_serialState_generateLdm( assert(srcSize <= serialState->params.jobSize); ZSTD_window_update(&serialState->ldmState.window, src, srcSize, /* forceNonContiguous */ 0); +} + +static void ZSTDMT_serialState_generateLdmSequences( + void* opaque, ZSTDMT_RustRawSeqStore* seqStore, + const void* src, size_t srcSize) +{ + SerialState* const serialState = (SerialState*)opaque; + RawSeqStore_t* const cSeqStore = (RawSeqStore_t*)seqStore; + size_t error; + error = ZSTD_ldm_generateSequences( &serialState->ldmState, cSeqStore, &serialState->params.ldmParams, src, srcSize); /* We provide a large enough buffer to never fail. */ assert(!ZSTD_isError(error)); (void)error; +} + +static void ZSTDMT_serialState_publishLdmWindow(void* opaque) +{ + SerialState* const serialState = (SerialState*)opaque; + /* Update ldmWindow to match the ldmState.window and signal the main * thread if it is waiting for a buffer. */ ZSTD_PTHREAD_MUTEX_LOCK(&serialState->ldmWindowMutex); @@ -2002,7 +2023,9 @@ static void ZSTDMT_compressionJobGenerateSequences(void* opaque) job->serial->params.fParams.checksumFlag, job->serial, ZSTDMT_serialState_waitForTurn, - ZSTDMT_serialState_generateLdm, + ZSTDMT_serialState_updateLdmWindow, + ZSTDMT_serialState_generateLdmSequences, + ZSTDMT_serialState_publishLdmWindow, ZSTDMT_serialState_updateChecksum, ZSTDMT_serialState_advance); } diff --git a/rust/src/zstdmt_compress.rs b/rust/src/zstdmt_compress.rs index 29a295f78..9e76263f3 100644 --- a/rust/src/zstdmt_compress.rs +++ b/rust/src/zstdmt_compress.rs @@ -653,8 +653,11 @@ const _: () = { }; pub type ZSTDMT_serialWaitForTurnFn = unsafe extern "C" fn(*mut c_void, c_uint) -> c_int; -pub type ZSTDMT_serialGenerateLdmFn = +pub type ZSTDMT_serialLdmWindowUpdateFn = unsafe extern "C" fn(*mut c_void, *mut ZstdMtRawSeqStore, *const c_void, usize); +pub type ZSTDMT_serialLdmGenerateSequencesFn = + unsafe extern "C" fn(*mut c_void, *mut ZstdMtRawSeqStore, *const c_void, usize); +pub type ZSTDMT_serialLdmPublishWindowFn = unsafe extern "C" fn(*mut c_void); pub type ZSTDMT_serialUpdateChecksumFn = unsafe extern "C" fn(*mut c_void, *const c_void, usize); pub type ZSTDMT_serialAdvanceFn = unsafe extern "C" fn(*mut c_void); pub type ZSTDMT_serialStateLockFn = unsafe extern "C" fn(*mut c_void); @@ -2560,7 +2563,7 @@ fn compression_job_with( /// invokes the advance callback exactly once, including when an earlier job /// has already skipped this turn. #[inline] -fn serial_state_gen_sequences_with( +fn serial_state_gen_sequences_with( seq_store: &mut ZstdMtRawSeqStore, src: *const c_void, src_size: usize, @@ -2568,18 +2571,24 @@ fn serial_state_gen_sequences_with( ldm_enabled: bool, checksum_enabled: bool, mut wait_for_turn: W, - mut generate_ldm: L, + mut update_ldm_window: U, + mut generate_ldm_sequences: G, + mut publish_ldm_window: P, mut update_checksum: C, mut advance: A, ) where W: FnMut(c_uint) -> bool, - L: FnMut(&mut ZstdMtRawSeqStore, *const c_void, usize), + U: FnMut(&mut ZstdMtRawSeqStore, *const c_void, usize), + G: FnMut(&mut ZstdMtRawSeqStore, *const c_void, usize), + P: FnMut(), C: FnMut(*const c_void, usize), A: FnMut(), { if wait_for_turn(job_id) { if ldm_enabled { - generate_ldm(seq_store, src, src_size); + update_ldm_window(seq_store, src, src_size); + generate_ldm_sequences(seq_store, src, src_size); + publish_ldm_window(); } if checksum_enabled && src_size != 0 { update_checksum(src, src_size); @@ -2600,12 +2609,27 @@ pub unsafe extern "C" fn ZSTDMT_rust_serialStateGenSequences( checksum_enabled: c_int, opaque: *mut c_void, wait_for_turn: Option, - generate_ldm: Option, + update_ldm_window: Option, + generate_ldm_sequences: Option, + publish_ldm_window: Option, update_checksum: Option, advance: Option, ) { - let (Some(wait_for_turn), Some(generate_ldm), Some(update_checksum), Some(advance)) = - (wait_for_turn, generate_ldm, update_checksum, advance) + let ( + Some(wait_for_turn), + Some(update_ldm_window), + Some(generate_ldm_sequences), + Some(publish_ldm_window), + Some(update_checksum), + Some(advance), + ) = ( + wait_for_turn, + update_ldm_window, + generate_ldm_sequences, + publish_ldm_window, + update_checksum, + advance, + ) else { return; }; @@ -2622,7 +2646,11 @@ pub unsafe extern "C" fn ZSTDMT_rust_serialStateGenSequences( ldm_enabled != 0, checksum_enabled != 0, |job_id| wait_for_turn(opaque, job_id) != 0, - |seq_store, src, src_size| generate_ldm(opaque, seq_store, src, src_size), + |seq_store, src, src_size| update_ldm_window(opaque, seq_store, src, src_size), + |seq_store, src, src_size| { + generate_ldm_sequences(opaque, seq_store, src, src_size) + }, + || publish_ldm_window(opaque), |src, src_size| update_checksum(opaque, src, src_size), || advance(opaque), ); @@ -6823,7 +6851,9 @@ mod tests { fn serial_turn_runs_ldm_before_checksum_and_advances_once() { let events = Rc::new(RefCell::new(Vec::new())); let wait_events = Rc::clone(&events); - let ldm_events = Rc::clone(&events); + let ldm_window_events = Rc::clone(&events); + let ldm_sequences_events = Rc::clone(&events); + let ldm_publish_events = Rc::clone(&events); let checksum_events = Rc::clone(&events); let advance_events = Rc::clone(&events); let mut seq_store = ZstdMtRawSeqStore::default(); @@ -6842,7 +6872,14 @@ mod tests { }, move |_seq_store, _src, src_size| { assert_eq!(src_size, 8); - ldm_events.borrow_mut().push("ldm"); + ldm_window_events.borrow_mut().push("ldm-window"); + }, + move |_seq_store, _src, src_size| { + assert_eq!(src_size, 8); + ldm_sequences_events.borrow_mut().push("ldm-sequences"); + }, + move || { + ldm_publish_events.borrow_mut().push("ldm-publish"); }, move |_src, src_size| { assert_eq!(src_size, 8); @@ -6851,7 +6888,17 @@ mod tests { move || advance_events.borrow_mut().push("advance"), ); - assert_eq!(&*events.borrow(), &["wait", "ldm", "checksum", "advance"]); + assert_eq!( + &*events.borrow(), + &[ + "wait", + "ldm-window", + "ldm-sequences", + "ldm-publish", + "checksum", + "advance", + ] + ); } #[test] @@ -6872,14 +6919,20 @@ mod tests { wait_events.borrow_mut().push("wait"); false }, - |_seq_store, _src, _src_size| panic!("skipped jobs must not generate LDM sequences"), + |_seq_store, _src, _src_size| panic!("skipped jobs must not update the LDM window"), + |_seq_store, _src, _src_size| { + panic!("skipped jobs must not generate LDM sequences") + }, + || panic!("skipped jobs must not publish the LDM window"), |_src, _src_size| panic!("skipped jobs must not update the checksum"), move || advance_events.borrow_mut().push("advance"), ); assert_eq!(&*events.borrow(), &["wait", "advance"]); let empty_events = Rc::new(RefCell::new(Vec::new())); - let empty_ldm_events = Rc::clone(&empty_events); + let empty_ldm_window_events = Rc::clone(&empty_events); + let empty_ldm_sequences_events = Rc::clone(&empty_events); + let empty_ldm_publish_events = Rc::clone(&empty_events); let empty_advance_events = Rc::clone(&empty_events); serial_state_gen_sequences_with( &mut seq_store, @@ -6891,12 +6944,24 @@ mod tests { |_job_id| true, move |_seq_store, _src, src_size| { assert_eq!(src_size, 0); - empty_ldm_events.borrow_mut().push("ldm"); + empty_ldm_window_events.borrow_mut().push("ldm-window"); + }, + move |_seq_store, _src, src_size| { + assert_eq!(src_size, 0); + empty_ldm_sequences_events + .borrow_mut() + .push("ldm-sequences"); + }, + move || { + empty_ldm_publish_events.borrow_mut().push("ldm-publish"); }, |_src, _src_size| panic!("empty input must not update the checksum"), move || empty_advance_events.borrow_mut().push("advance"), ); - assert_eq!(&*empty_events.borrow(), &["ldm", "advance"]); + assert_eq!( + &*empty_events.borrow(), + &["ldm-window", "ldm-sequences", "ldm-publish", "advance"] + ); } #[test]