diff --git a/lib/compress/zstd_compress_internal.h b/lib/compress/zstd_compress_internal.h index 5ed9b02bd..2142dc12c 100644 --- a/lib/compress/zstd_compress_internal.h +++ b/lib/compress/zstd_compress_internal.h @@ -655,6 +655,23 @@ size_t ZSTD_rust_noCompressBlock(void* dst, size_t dstCapacity, const void* src, size_t srcSize, U32 lastBlock); size_t ZSTD_rust_rleCompressBlock(void* dst, size_t dstCapacity, BYTE src, size_t srcSize, U32 lastBlock); +typedef struct { + const void* nextSrc; + const void* base; + const void* dictBase; + U32 dictLimit; + U32 lowLimit; +} ZSTD_rust_windowUpdateState; +typedef char ZSTD_rust_window_update_state_layout[ + (offsetof(ZSTD_rust_windowUpdateState, nextSrc) == 0 + && offsetof(ZSTD_rust_windowUpdateState, base) == sizeof(void*) + && offsetof(ZSTD_rust_windowUpdateState, dictBase) == 2 * sizeof(void*) + && offsetof(ZSTD_rust_windowUpdateState, dictLimit) == 3 * sizeof(void*) + && offsetof(ZSTD_rust_windowUpdateState, lowLimit) == 3 * sizeof(void*) + sizeof(U32) + && sizeof(ZSTD_rust_windowUpdateState) == 3 * sizeof(void*) + 2 * sizeof(U32)) ? 1 : -1]; +U32 ZSTD_rust_windowUpdate(ZSTD_rust_windowUpdateState* state, + const void* src, size_t srcSize, + int forceNonContiguous); void ZSTD_rust_windowClear(size_t endT, U32* lowLimit, U32* dictLimit); U32 ZSTD_rust_windowCorrectOverflow(U32 curr, U32 cycleLog, U32 maxDist); U32 ZSTD_rust_windowCanOverflowCorrect(U32 curr, U32 cycleLog, U32 maxDist, @@ -1311,37 +1328,21 @@ U32 ZSTD_window_update(ZSTD_window_t* window, const void* src, size_t srcSize, int forceNonContiguous) { - BYTE const* const ip = (BYTE const*)src; - U32 contiguous = 1; + ZSTD_rust_windowUpdateState state; + U32 contiguous; DEBUGLOG(5, "ZSTD_window_update"); - if (srcSize == 0) - return contiguous; - assert(window->base != NULL); - assert(window->dictBase != NULL); - /* Check if blocks follow each other */ - if (src != window->nextSrc || forceNonContiguous) { - /* not contiguous */ - size_t const distanceFromBase = (size_t)(window->nextSrc - window->base); - DEBUGLOG(5, "Non contiguous blocks, new segment starts at %u", window->dictLimit); - window->lowLimit = window->dictLimit; - assert(distanceFromBase == (size_t)(U32)distanceFromBase); /* should never overflow */ - window->dictLimit = (U32)distanceFromBase; - window->dictBase = window->base; - window->base = ip - distanceFromBase; - /* ms->nextToUpdate = window->dictLimit; */ - if (window->dictLimit - window->lowLimit < HASH_READ_SIZE) window->lowLimit = window->dictLimit; /* too small extDict */ - contiguous = 0; - } - window->nextSrc = ip + srcSize; - /* if input and dictionary overlap : reduce dictionary (area presumed modified by input) */ - if ( (ip+srcSize > window->dictBase + window->lowLimit) - & (ip < window->dictBase + window->dictLimit)) { - size_t const highInputIdx = (size_t)((ip + srcSize) - window->dictBase); - U32 const lowLimitMax = (highInputIdx > (size_t)window->dictLimit) ? window->dictLimit : (U32)highInputIdx; - assert(highInputIdx < UINT_MAX); - window->lowLimit = lowLimitMax; - DEBUGLOG(5, "Overlapping extDict and input : new lowLimit = %u", window->lowLimit); - } + state.nextSrc = window->nextSrc; + state.base = window->base; + state.dictBase = window->dictBase; + state.dictLimit = window->dictLimit; + state.lowLimit = window->lowLimit; + contiguous = ZSTD_rust_windowUpdate(&state, src, srcSize, + forceNonContiguous); + window->nextSrc = (BYTE const*)state.nextSrc; + window->base = (BYTE const*)state.base; + window->dictBase = (BYTE const*)state.dictBase; + window->dictLimit = state.dictLimit; + window->lowLimit = state.lowLimit; return contiguous; } diff --git a/rust/src/zstd_compress.rs b/rust/src/zstd_compress.rs index ac25d73a7..4739e0b1e 100644 --- a/rust/src/zstd_compress.rs +++ b/rust/src/zstd_compress.rs @@ -6789,6 +6789,88 @@ pub unsafe extern "C" fn ZSTD_rust_windowClear( } } +const HASH_READ_SIZE: u32 = 8; + +/// Private C window fields projected for the shared window-update policy. +/// Pointer arithmetic is intentionally performed as wrapping address math to +/// mirror the original C implementation's pointer-overflow contract. +#[repr(C)] +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub struct ZSTD_rust_windowUpdateState { + pub nextSrc: *const c_void, + pub base: *const c_void, + pub dictBase: *const c_void, + pub dictLimit: u32, + pub lowLimit: u32, +} + +const _: () = { + assert!(offset_of!(ZSTD_rust_windowUpdateState, nextSrc) == 0); + assert!(offset_of!(ZSTD_rust_windowUpdateState, base) == size_of::()); + assert!(offset_of!(ZSTD_rust_windowUpdateState, dictBase) == 2 * size_of::()); + assert!(offset_of!(ZSTD_rust_windowUpdateState, dictLimit) == 3 * size_of::()); + assert!( + offset_of!(ZSTD_rust_windowUpdateState, lowLimit) + == 3 * size_of::() + size_of::() + ); + assert!( + size_of::() == 3 * size_of::() + 2 * size_of::() + ); +}; + +/// Update the rolling prefix/ext-dictionary window for one input segment. +/// C owns the containing match state; Rust owns this five-field transition. +#[no_mangle] +pub unsafe extern "C" fn ZSTD_rust_windowUpdate( + state: *mut ZSTD_rust_windowUpdateState, + src: *const c_void, + src_size: usize, + force_non_contiguous: c_int, +) -> c_uint { + if state.is_null() { + return 1; + } + let state = unsafe { &mut *state }; + if src_size == 0 { + return 1; + } + + debug_assert!(!state.base.is_null()); + debug_assert!(!state.dictBase.is_null()); + + let ip = src as usize; + let input_end = ip.wrapping_add(src_size); + let mut contiguous = 1; + if ip != state.nextSrc as usize || force_non_contiguous != 0 { + let distance_from_base = (state.nextSrc as usize).wrapping_sub(state.base as usize); + state.lowLimit = state.dictLimit; + debug_assert!(distance_from_base <= u32::MAX as usize); + state.dictLimit = distance_from_base as u32; + state.dictBase = state.base; + state.base = ip.wrapping_sub(distance_from_base) as *const c_void; + if state.dictLimit.wrapping_sub(state.lowLimit) < HASH_READ_SIZE { + state.lowLimit = state.dictLimit; + } + contiguous = 0; + } + + state.nextSrc = input_end as *const c_void; + let dict_base = state.dictBase as usize; + if input_end > dict_base.wrapping_add(state.lowLimit as usize) + && ip < dict_base.wrapping_add(state.dictLimit as usize) + { + let high_input_idx = input_end.wrapping_sub(dict_base); + debug_assert!(high_input_idx < u32::MAX as usize); + let low_limit_max = if high_input_idx > state.dictLimit as usize { + state.dictLimit + } else { + high_input_idx as u32 + }; + state.lowLimit = low_limit_max; + } + contiguous +} + #[inline] fn set_pledged_src_size( stream_stage: c_int, @@ -12509,6 +12591,86 @@ mod tests { assert_eq!(dict_limit, 0x89ab_cdef); } + #[test] + fn window_update_leaves_empty_input_untouched() { + let mut state = ZSTD_rust_windowUpdateState { + nextSrc: ptr::null(), + base: ptr::null(), + dictBase: ptr::null(), + dictLimit: 2, + lowLimit: 2, + }; + let original = state; + + let result = unsafe { ZSTD_rust_windowUpdate(&mut state, ptr::null(), 0, 0) }; + + assert_eq!(result, 1); + assert_eq!(state, original); + } + + #[test] + fn window_update_preserves_contiguous_input_and_clips_overlapping_extdict() { + let storage = [0u8; 64]; + let storage_start = storage.as_ptr() as usize; + let base = unsafe { storage.as_ptr().add(16) }; + let next_src = unsafe { base.add(2) }; + let mut contiguous_state = ZSTD_rust_windowUpdateState { + nextSrc: next_src.cast(), + base: base.cast(), + dictBase: base.cast(), + dictLimit: 2, + lowLimit: 2, + }; + + let contiguous = + unsafe { ZSTD_rust_windowUpdate(&mut contiguous_state, next_src.cast(), 3, 0) }; + + assert_eq!(contiguous, 1); + assert_eq!(contiguous_state.nextSrc as usize, storage_start + 21); + assert_eq!(contiguous_state.base as usize, storage_start + 16); + assert_eq!(contiguous_state.dictBase as usize, storage_start + 16); + assert_eq!(contiguous_state.dictLimit, 2); + assert_eq!(contiguous_state.lowLimit, 2); + + let next_src = unsafe { base.add(6) }; + let mut forced_state = ZSTD_rust_windowUpdateState { + nextSrc: next_src.cast(), + base: base.cast(), + dictBase: base.cast(), + dictLimit: 5, + lowLimit: 2, + }; + + let forced = unsafe { ZSTD_rust_windowUpdate(&mut forced_state, next_src.cast(), 3, 1) }; + + assert_eq!(forced, 0); + assert_eq!(forced_state.base as usize, storage_start + 16); + assert_eq!(forced_state.nextSrc as usize, storage_start + 25); + assert_eq!(forced_state.dictBase as usize, storage_start + 16); + assert_eq!(forced_state.dictLimit, 6); + assert_eq!(forced_state.lowLimit, 6); + + let next_src = unsafe { base.add(20) }; + let src = unsafe { base.add(4) }; + let mut non_contiguous_state = ZSTD_rust_windowUpdateState { + nextSrc: next_src.cast(), + base: base.cast(), + dictBase: base.cast(), + dictLimit: 5, + lowLimit: 2, + }; + + let non_contiguous = + unsafe { ZSTD_rust_windowUpdate(&mut non_contiguous_state, src.cast(), 10, 0) }; + + assert_eq!(non_contiguous, 0); + assert_eq!(non_contiguous_state.base as usize, storage_start); + assert_eq!(non_contiguous_state.nextSrc as usize, storage_start + 30); + assert_eq!(non_contiguous_state.dictBase as usize, storage_start + 16); + assert_eq!(non_contiguous_state.dictLimit, 20); + assert_eq!(non_contiguous_state.lowLimit, 14); + } + #[derive(Default)] struct ResetCCtxTestContext { events: Vec<&'static str>,