diff --git a/lib/compress/zstdmt_compress.c b/lib/compress/zstdmt_compress.c index 6b2eeda7f..b1af71a8d 100644 --- a/lib/compress/zstdmt_compress.c +++ b/lib/compress/zstdmt_compress.c @@ -713,8 +713,6 @@ 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); typedef void (*ZSTDMT_serialStateLockFn)(void* opaque); typedef void (*ZSTDMT_serialStateWaitFn)(void* opaque); @@ -752,11 +750,11 @@ void ZSTDMT_rust_serialStateAdvance(const ZSTDMT_RustSerialAdvanceState* state); void ZSTDMT_rust_serialStateGenSequences( ZSTDMT_RustRawSeqStore* seqStore, const void* src, size_t srcSize, unsigned jobID, int ldmEnabled, int checksumEnabled, void* opaque, + XXH64_state_t* checksumState, ZSTDMT_serialWaitForTurnFn waitForTurn, ZSTDMT_serialLdmWindowUpdateFn updateLdmWindow, ZSTDMT_serialLdmGenerateSequencesFn generateLdmSequences, ZSTDMT_serialLdmPublishWindowFn publishLdmWindow, - ZSTDMT_serialUpdateChecksumFn updateChecksum, ZSTDMT_serialAdvanceFn advance); typedef struct { unsigned skip; @@ -1892,13 +1890,6 @@ static void ZSTDMT_serialState_publishLdmWindow(void* opaque) ZSTD_pthread_mutex_unlock(&serialState->ldmWindowMutex); } -static void ZSTDMT_serialState_updateChecksum( - void* opaque, const void* src, size_t srcSize) -{ - SerialState* const serialState = (SerialState*)opaque; - XXH64_update(&serialState->xxhState, src, srcSize); -} - static void ZSTDMT_serialState_advance(void* opaque) { SerialState* const serialState = (SerialState*)opaque; @@ -2140,11 +2131,11 @@ static void ZSTDMT_compressionJobGenerateSequences(void* opaque) job->serial->params.ldmParams.enableLdm == ZSTD_ps_enable, job->serial->params.fParams.checksumFlag, job->serial, + &job->serial->xxhState, ZSTDMT_serialState_waitForTurn, 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 16abb1ee6..20f2b656d 100644 --- a/rust/src/zstdmt_compress.rs +++ b/rust/src/zstdmt_compress.rs @@ -23,6 +23,7 @@ use std::sync::Mutex; use crate::bits::ZSTD_highbit32; use crate::errors::{ERR_isError, ZstdErrorCode, ERROR}; use crate::mem::MEM_writeLE32; +use crate::xxhash::{XXH64_state_t, XXH64_update}; use crate::zstd_compress::{ZSTD_frameProgression, ZSTD_rust_invalidateRepCodes}; use crate::zstd_compress_frame::ZSTD_rust_writeFrameHeader; use crate::zstd_compress_params::{ @@ -907,7 +908,6 @@ pub type ZSTDMT_serialLdmWindowUpdateFn = 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); pub type ZSTDMT_serialStateWaitFn = unsafe extern "C" fn(*mut c_void); @@ -2905,8 +2905,21 @@ fn serial_state_gen_sequences_with( advance(); } -/// Own MT serial turn/skip and LDM/checksum ordering while C retains the -/// mutexes, LDM window/hash state, checksum state, and raw sequence storage. +#[inline] +unsafe fn serial_state_update_checksum( + checksum_state: *mut XXH64_state_t, + src: *const c_void, + src_size: usize, +) { + if checksum_state.is_null() || src_size == 0 { + return; + } + let _ = unsafe { XXH64_update(checksum_state, src, src_size) }; +} + +/// Own MT serial turn/skip and LDM/checksum-update ordering while C retains +/// the mutexes, LDM window/hash state, checksum state, and raw sequence +/// storage. #[no_mangle] pub unsafe extern "C" fn ZSTDMT_rust_serialStateGenSequences( seq_store: *mut ZstdMtRawSeqStore, @@ -2916,11 +2929,11 @@ pub unsafe extern "C" fn ZSTDMT_rust_serialStateGenSequences( ldm_enabled: c_int, checksum_enabled: c_int, opaque: *mut c_void, + checksum_state: *mut XXH64_state_t, wait_for_turn: Option, update_ldm_window: Option, generate_ldm_sequences: Option, publish_ldm_window: Option, - update_checksum: Option, advance: Option, ) { let ( @@ -2928,20 +2941,18 @@ pub unsafe extern "C" fn ZSTDMT_rust_serialStateGenSequences( 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; }; - if seq_store.is_null() { + if seq_store.is_null() || (checksum_enabled != 0 && checksum_state.is_null()) { return; } @@ -2959,7 +2970,7 @@ pub unsafe extern "C" fn ZSTDMT_rust_serialStateGenSequences( generate_ldm_sequences(opaque, seq_store, src, src_size) }, || publish_ldm_window(opaque), - |src, src_size| update_checksum(opaque, src, src_size), + |src, src_size| serial_state_update_checksum(checksum_state, src, src_size), || advance(opaque), ); } @@ -5955,6 +5966,7 @@ pub unsafe extern "C" fn ZSTDMT_rust_cctx_pool_release(pool: *mut RustCCtxPool, #[cfg(test)] mod tests { use super::*; + use crate::xxhash::{XXH64, XXH64_digest, XXH64_reset}; use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering}; use std::{cell::RefCell, rc::Rc}; @@ -7300,6 +7312,20 @@ mod tests { ); } + #[test] + fn serial_checksum_update_matches_one_shot_xxhash() { + let input = b"serial checksum"; + let mut state = unsafe { MaybeUninit::::zeroed().assume_init() }; + + unsafe { + assert_eq!(XXH64_reset(&mut state, 0), 0); + serial_state_update_checksum(&mut state, input.as_ptr().cast(), input.len()); + } + + let expected = unsafe { XXH64(input.as_ptr().cast(), input.len(), 0) }; + assert_eq!(unsafe { XXH64_digest(&state) }, expected); + } + #[test] fn serial_turn_skips_failed_predecessor_and_keeps_empty_ldm_turn() { let events = Rc::new(RefCell::new(Vec::new()));