From 031592da1e99b0ca2c8d5c60e7953768addcfa40 Mon Sep 17 00:00:00 2001 From: ddidderr Date: Sat, 11 Jul 2026 08:22:41 +0200 Subject: [PATCH] feat(rust): port optimal binary-tree updates Move the binary-tree table maintenance shared by optimal compression strategies into Rust. A small C projection preserves the private match-state boundary; the dynamic-programming parser and match selection remain C-owned for now. The Rust implementation covers prefix and external-dictionary tree updates and preserves the optional C predictor path. The component map now distinguishes this migrated tree layer from the remaining optimal parser. Test Plan: - cargo test --all-targets - cargo test --target i686-unknown-linux-gnu --all-targets - cargo clippy && cargo clippy --benches && cargo clippy --tests - cargo +nightly fmt - make -B -C tests -j2 fuzzer zstreamtest invalidDictionaries - ./tests/fuzzer -s5346 -i1 --no-big-tests - ./tests/zstreamtest -i3000 -s334462 - ./tests/invalidDictionaries - make -B -C tests fuzzer32 MOREFLAGS=-DZSTD_C_PREDICT - ./tests/fuzzer32 -s5346 -i1 --no-big-tests - byte-compare C and Rust-control CLI output at levels 13, 16, 19, and 22 for random, patterned, zero, and dictionary inputs with and without ZSTD_C_PREDICT Refs: rust/README.md --- lib/compress/zstd_opt.c | 188 +++++------------- rust/README.md | 2 + rust/src/lib.rs | 2 + rust/src/zstd_opt_tree.rs | 395 ++++++++++++++++++++++++++++++++++++++ 4 files changed, 449 insertions(+), 138 deletions(-) create mode 100644 rust/src/zstd_opt_tree.rs diff --git a/lib/compress/zstd_opt.c b/lib/compress/zstd_opt.c index 3d7171b75..9b636445a 100644 --- a/lib/compress/zstd_opt.c +++ b/lib/compress/zstd_opt.c @@ -16,6 +16,53 @@ || !defined(ZSTD_EXCLUDE_BTOPT_BLOCK_COMPRESSOR) \ || !defined(ZSTD_EXCLUDE_BTULTRA_BLOCK_COMPRESSOR) +/* The Rust tree updater receives this leaf view instead of the private, + * configuration-dependent ZSTD_MatchState_t layout. */ +typedef struct { + U32* hashTable; + U32* chainTable; + const BYTE* base; + const BYTE* dictBase; + U32 dictLimit; + U32 lowLimit; + U32 loadedDictEnd; + U32* nextToUpdate; + U32 hashLog; + U32 chainLog; + U32 searchLog; + U32 windowLog; + int useCPredict; +} ZSTD_RustOptTreeState; + +static ZSTD_RustOptTreeState ZSTD_rustOptTreeState( + ZSTD_MatchState_t* const ms) +{ + ZSTD_RustOptTreeState state; + state.hashTable = ms->hashTable; + state.chainTable = ms->chainTable; + state.base = ms->window.base; + state.dictBase = ms->window.dictBase; + state.dictLimit = ms->window.dictLimit; + state.lowLimit = ms->window.lowLimit; + state.loadedDictEnd = ms->loadedDictEnd; + state.nextToUpdate = &ms->nextToUpdate; + state.hashLog = ms->cParams.hashLog; + state.chainLog = ms->cParams.chainLog; + state.searchLog = ms->cParams.searchLog; + state.windowLog = ms->cParams.windowLog; +#ifdef ZSTD_C_PREDICT + state.useCPredict = 1; +#else + state.useCPredict = 0; +#endif + return state; +} + +void ZSTD_rust_opt_updateTreeInternal( + ZSTD_RustOptTreeState* state, + const void* ip, const void* iend, + U32 mls, int extDict); + #define ZSTD_LITFREQ_ADD 2 /* scaling factor for litFreq, so that frequencies adapt faster to new stats */ #define ZSTD_MAX_PRICE (1<<30) @@ -433,130 +480,6 @@ U32 ZSTD_insertAndFindFirstIndexHash3 (const ZSTD_MatchState_t* ms, /*-************************************* * Binary Tree search ***************************************/ -/** ZSTD_insertBt1() : add one or multiple positions to tree. - * @param ip assumed <= iend-8 . - * @param target The target of ZSTD_updateTree_internal() - we are filling to this position - * @return : nb of positions added */ -static -ZSTD_ALLOW_POINTER_OVERFLOW_ATTR -U32 ZSTD_insertBt1( - const ZSTD_MatchState_t* ms, - const BYTE* const ip, const BYTE* const iend, - U32 const target, - U32 const mls, const int extDict) -{ - const ZSTD_compressionParameters* const cParams = &ms->cParams; - U32* const hashTable = ms->hashTable; - U32 const hashLog = cParams->hashLog; - size_t const h = ZSTD_hashPtr(ip, hashLog, mls); - U32* const bt = ms->chainTable; - U32 const btLog = cParams->chainLog - 1; - U32 const btMask = (1 << btLog) - 1; - U32 matchIndex = hashTable[h]; - size_t commonLengthSmaller=0, commonLengthLarger=0; - const BYTE* const base = ms->window.base; - const BYTE* const dictBase = ms->window.dictBase; - const U32 dictLimit = ms->window.dictLimit; - const BYTE* const dictEnd = dictBase + dictLimit; - const BYTE* const prefixStart = base + dictLimit; - const BYTE* match; - const U32 curr = (U32)(ip-base); - const U32 btLow = btMask >= curr ? 0 : curr - btMask; - U32* smallerPtr = bt + 2*(curr&btMask); - U32* largerPtr = smallerPtr + 1; - U32 dummy32; /* to be nullified at the end */ - /* windowLow is based on target because - * we only need positions that will be in the window at the end of the tree update. - */ - U32 const windowLow = ZSTD_getLowestMatchIndex(ms, target, cParams->windowLog); - U32 matchEndIdx = curr+8+1; - size_t bestLength = 8; - U32 nbCompares = 1U << cParams->searchLog; -#ifdef ZSTD_C_PREDICT - U32 predictedSmall = *(bt + 2*((curr-1)&btMask) + 0); - U32 predictedLarge = *(bt + 2*((curr-1)&btMask) + 1); - predictedSmall += (predictedSmall>0); - predictedLarge += (predictedLarge>0); -#endif /* ZSTD_C_PREDICT */ - - DEBUGLOG(8, "ZSTD_insertBt1 (%u)", curr); - - assert(curr <= target); - assert(ip <= iend-8); /* required for h calculation */ - hashTable[h] = curr; /* Update Hash Table */ - - assert(windowLow > 0); - for (; nbCompares && (matchIndex >= windowLow); --nbCompares) { - U32* const nextPtr = bt + 2*(matchIndex & btMask); - size_t matchLength = MIN(commonLengthSmaller, commonLengthLarger); /* guaranteed minimum nb of common bytes */ - assert(matchIndex < curr); - -#ifdef ZSTD_C_PREDICT /* note : can create issues when hlog small <= 11 */ - const U32* predictPtr = bt + 2*((matchIndex-1) & btMask); /* written this way, as bt is a roll buffer */ - if (matchIndex == predictedSmall) { - /* no need to check length, result known */ - *smallerPtr = matchIndex; - if (matchIndex <= btLow) { smallerPtr=&dummy32; break; } /* beyond tree size, stop the search */ - smallerPtr = nextPtr+1; /* new "smaller" => larger of match */ - matchIndex = nextPtr[1]; /* new matchIndex larger than previous (closer to current) */ - predictedSmall = predictPtr[1] + (predictPtr[1]>0); - continue; - } - if (matchIndex == predictedLarge) { - *largerPtr = matchIndex; - if (matchIndex <= btLow) { largerPtr=&dummy32; break; } /* beyond tree size, stop the search */ - largerPtr = nextPtr; - matchIndex = nextPtr[0]; - predictedLarge = predictPtr[0] + (predictPtr[0]>0); - continue; - } -#endif - - if (!extDict || (matchIndex+matchLength >= dictLimit)) { - assert(matchIndex+matchLength >= dictLimit); /* might be wrong if actually extDict */ - match = base + matchIndex; - matchLength += ZSTD_count(ip+matchLength, match+matchLength, iend); - } else { - match = dictBase + matchIndex; - matchLength += ZSTD_count_2segments(ip+matchLength, match+matchLength, iend, dictEnd, prefixStart); - if (matchIndex+matchLength >= dictLimit) - match = base + matchIndex; /* to prepare for next usage of match[matchLength] */ - } - - if (matchLength > bestLength) { - bestLength = matchLength; - if (matchLength > matchEndIdx - matchIndex) - matchEndIdx = matchIndex + (U32)matchLength; - } - - if (ip+matchLength == iend) { /* equal : no way to know if inf or sup */ - break; /* drop , to guarantee consistency ; miss a bit of compression, but other solutions can corrupt tree */ - } - - if (match[matchLength] < ip[matchLength]) { /* necessarily within buffer */ - /* match is smaller than current */ - *smallerPtr = matchIndex; /* update smaller idx */ - commonLengthSmaller = matchLength; /* all smaller will now have at least this guaranteed common length */ - if (matchIndex <= btLow) { smallerPtr=&dummy32; break; } /* beyond tree size, stop searching */ - smallerPtr = nextPtr+1; /* new "candidate" => larger than match, which was smaller than target */ - matchIndex = nextPtr[1]; /* new matchIndex, larger than previous and closer to current */ - } else { - /* match is larger than current */ - *largerPtr = matchIndex; - commonLengthLarger = matchLength; - if (matchIndex <= btLow) { largerPtr=&dummy32; break; } /* beyond tree size, stop searching */ - largerPtr = nextPtr; - matchIndex = nextPtr[0]; - } } - - *smallerPtr = *largerPtr = 0; - { U32 positions = 0; - if (bestLength > 384) positions = MIN(192, (U32)(bestLength - 384)); /* speed optimization */ - assert(matchEndIdx > curr + 8); - return MAX(positions, matchEndIdx - (curr + 8)); - } -} - FORCE_INLINE_TEMPLATE ZSTD_ALLOW_POINTER_OVERFLOW_ATTR void ZSTD_updateTree_internal( @@ -564,20 +487,9 @@ void ZSTD_updateTree_internal( const BYTE* const ip, const BYTE* const iend, const U32 mls, const ZSTD_dictMode_e dictMode) { - const BYTE* const base = ms->window.base; - U32 const target = (U32)(ip - base); - U32 idx = ms->nextToUpdate; - DEBUGLOG(7, "ZSTD_updateTree_internal, from %u to %u (dictMode:%u)", - idx, target, dictMode); - - while(idx < target) { - U32 const forward = ZSTD_insertBt1(ms, base+idx, iend, target, mls, dictMode == ZSTD_extDict); - assert(idx < (U32)(idx + forward)); - idx += forward; - } - assert((size_t)(ip - base) <= (size_t)(U32)(-1)); - assert((size_t)(iend - base) <= (size_t)(U32)(-1)); - ms->nextToUpdate = target; + ZSTD_RustOptTreeState state = ZSTD_rustOptTreeState(ms); + ZSTD_rust_opt_updateTreeInternal(&state, ip, iend, mls, + dictMode == ZSTD_extDict); } void ZSTD_updateTree(ZSTD_MatchState_t* ms, const BYTE* ip, const BYTE* iend) { diff --git a/rust/README.md b/rust/README.md index 5e6d19c44..c128afa98 100644 --- a/rust/README.md +++ b/rust/README.md @@ -35,6 +35,8 @@ zstd ABI: fast block match finders, including attached and external dictionary paths. - `zstd_lazy` implements greedy, lazy, lazy2, and binary-tree matching, including row-based and dictionary search variants. + - `zstd_opt_tree` maintains the binary-tree index used by optimal matching; + the dynamic-programming optimal parser itself remains in C for now. - Runtime support - `threading` provides platform pthread wrappers required by zstd headers. - `pool` implements the bounded worker pool used by multithreaded compression. diff --git a/rust/src/lib.rs b/rust/src/lib.rs index 08f1079ab..efbafe9b8 100644 --- a/rust/src/lib.rs +++ b/rust/src/lib.rs @@ -40,4 +40,6 @@ pub mod zstd_fast; #[cfg(feature = "compression")] pub mod zstd_lazy; #[cfg(feature = "compression")] +pub mod zstd_opt_tree; +#[cfg(feature = "compression")] pub mod zstd_presplit; diff --git a/rust/src/zstd_opt_tree.rs b/rust/src/zstd_opt_tree.rs new file mode 100644 index 000000000..4a6354067 --- /dev/null +++ b/rust/src/zstd_opt_tree.rs @@ -0,0 +1,395 @@ +#![allow(non_camel_case_types)] +#![allow(non_snake_case)] +#![allow(clippy::missing_safety_doc)] + +//! Binary-tree maintenance for optimal block compression. +//! +//! `ZSTD_MatchState_t` remains a C-owned, configuration-sensitive structure. +//! The companion C translation unit projects only the matching-table and window +//! leaves used to populate its binary tree. The parser itself remains separate +//! while this module moves the shared tree-maintenance path into Rust. + +use crate::mem::{ + MEM_64bits, MEM_isLittleEndian, MEM_read16, MEM_read32, MEM_readLE32, MEM_readLE64, MEM_readST, +}; +use std::cmp::{max, min}; +use std::ffi::c_void; +use std::mem::size_of; +use std::os::raw::c_int; + +const HASH_READ_SIZE: usize = 8; + +/// Leaf view built by `zstd_opt.c` from its C-owned match state. +#[repr(C)] +pub struct ZSTD_RustOptTreeState { + hash_table: *mut u32, + chain_table: *mut u32, + base: *const u8, + dict_base: *const u8, + dict_limit: u32, + low_limit: u32, + loaded_dict_end: u32, + next_to_update: *mut u32, + hash_log: u32, + chain_log: u32, + search_log: u32, + window_log: u32, + use_c_predict: c_int, +} + +#[inline] +fn ptr_lt(left: *const u8, right: *const u8) -> bool { + (left as usize) < (right as usize) +} + +#[inline] +unsafe fn hash_shift32(value: u32, hbits: u32) -> usize { + if hbits == 0 { + 0 + } else { + (value >> (32 - hbits)) as usize + } +} + +#[inline] +unsafe fn hash_shift64(value: u64, hbits: u32) -> usize { + if hbits == 0 { + 0 + } else { + (value >> (64 - hbits)) as usize + } +} + +#[inline] +unsafe fn hash_ptr(ptr: *const u8, hbits: u32, mls: u32) -> usize { + match mls { + 5 => { + let value = unsafe { MEM_readLE64(ptr.cast::()) }; + unsafe { hash_shift64(value.wrapping_shl(24).wrapping_mul(889_523_592_379), hbits) } + } + 6 => { + let value = unsafe { MEM_readLE64(ptr.cast::()) }; + unsafe { + hash_shift64( + value.wrapping_shl(16).wrapping_mul(227_718_039_650_203), + hbits, + ) + } + } + 7 => { + let value = unsafe { MEM_readLE64(ptr.cast::()) }; + unsafe { + hash_shift64( + value.wrapping_shl(8).wrapping_mul(58_295_818_150_454_627), + hbits, + ) + } + } + 8 => { + let value = unsafe { MEM_readLE64(ptr.cast::()) }; + unsafe { hash_shift64(value.wrapping_mul(0xCF1B_BCDC_B7A5_6463), hbits) } + } + _ => { + let value = unsafe { MEM_readLE32(ptr.cast::()) }; + unsafe { hash_shift32(value.wrapping_mul(2_654_435_761), hbits) } + } + } +} + +#[inline] +fn common_bytes(word: usize) -> usize { + let zeros = if MEM_isLittleEndian() { + word.trailing_zeros() + } else { + word.leading_zeros() + }; + (zeros / 8) as usize +} + +/// Equivalent to C's `ZSTD_count()` including its word-at-a-time tail. +unsafe fn count(mut input: *const u8, mut matched: *const u8, input_limit: *const u8) -> usize { + let input_start = input; + let word_size = size_of::(); + while unsafe { input_limit.offset_from(input) as usize } >= word_size { + let difference = + unsafe { MEM_readST(matched.cast::()) ^ MEM_readST(input.cast::()) }; + if difference != 0 { + return unsafe { input.offset_from(input_start) as usize } + common_bytes(difference); + } + input = input.wrapping_add(word_size); + matched = matched.wrapping_add(word_size); + } + if MEM_64bits() + && unsafe { input_limit.offset_from(input) as usize } >= 4 + && unsafe { MEM_read32(matched.cast::()) == MEM_read32(input.cast::()) } + { + input = input.wrapping_add(4); + matched = matched.wrapping_add(4); + } + if unsafe { input_limit.offset_from(input) as usize } >= 2 + && unsafe { MEM_read16(matched.cast::()) == MEM_read16(input.cast::()) } + { + input = input.wrapping_add(2); + matched = matched.wrapping_add(2); + } + if ptr_lt(input, input_limit) && unsafe { *matched == *input } { + input = input.wrapping_add(1); + } + unsafe { input.offset_from(input_start) as usize } +} + +unsafe fn count_2segments( + input: *const u8, + matched: *const u8, + input_end: *const u8, + match_end: *const u8, + input_start: *const u8, +) -> usize { + let match_remaining = unsafe { match_end.offset_from(matched) as usize }; + let input_remaining = unsafe { input_end.offset_from(input) as usize }; + let first_end = input.wrapping_add(min(match_remaining, input_remaining)); + let first_count = unsafe { count(input, matched, first_end) }; + if matched.wrapping_add(first_count) != match_end { + return first_count; + } + first_count + unsafe { count(input.wrapping_add(first_count), input_start, input_end) } +} + +#[inline] +fn lowest_match_index(state: &ZSTD_RustOptTreeState, current: u32) -> u32 { + let max_distance = 1u32.wrapping_shl(state.window_log); + let within_window = if current.wrapping_sub(state.low_limit) > max_distance { + current.wrapping_sub(max_distance) + } else { + state.low_limit + }; + if state.loaded_dict_end != 0 { + state.low_limit + } else { + within_window + } +} + +unsafe fn insert_bt1( + state: &ZSTD_RustOptTreeState, + ip: *const u8, + iend: *const u8, + target: u32, + mls: u32, + ext_dict: bool, +) -> u32 { + let hash = unsafe { hash_ptr(ip, state.hash_log, mls) }; + let bt_log = state.chain_log - 1; + let bt_mask = (1u32 << bt_log) - 1; + let mut match_index = unsafe { *state.hash_table.add(hash) }; + let mut common_smaller = 0usize; + let mut common_larger = 0usize; + let dict_end = state.dict_base.wrapping_add(state.dict_limit as usize); + let prefix_start = state.base.wrapping_add(state.dict_limit as usize); + let current = unsafe { ip.offset_from(state.base) as u32 }; + let bt_low = current.saturating_sub(bt_mask); + let mut smaller = unsafe { state.chain_table.add(2 * (current & bt_mask) as usize) }; + let mut larger = unsafe { smaller.add(1) }; + let mut dummy = 0u32; + let window_low = lowest_match_index(state, target); + let mut match_end_index = current.wrapping_add(HASH_READ_SIZE as u32 + 1); + let mut best_length = HASH_READ_SIZE; + let mut comparisons = 1u32 << state.search_log; + + unsafe { *state.hash_table.add(hash) = current }; + + let mut predicted_small = 0u32; + let mut predicted_large = 0u32; + if state.use_c_predict != 0 { + let prior = unsafe { + state + .chain_table + .add(2 * (current.wrapping_sub(1) & bt_mask) as usize) + }; + predicted_small = unsafe { *prior }.wrapping_add(u32::from(unsafe { *prior } > 0)); + predicted_large = + unsafe { *prior.add(1) }.wrapping_add(u32::from(unsafe { *prior.add(1) } > 0)); + } + + while comparisons != 0 && match_index >= window_low { + comparisons -= 1; + let next = unsafe { state.chain_table.add(2 * (match_index & bt_mask) as usize) }; + let mut match_length = min(common_smaller, common_larger); + + if state.use_c_predict != 0 && match_index == predicted_small { + let prediction = unsafe { + state + .chain_table + .add(2 * (match_index.wrapping_sub(1) & bt_mask) as usize) + }; + unsafe { *smaller = match_index }; + if match_index <= bt_low { + smaller = &mut dummy; + break; + } + smaller = unsafe { next.add(1) }; + match_index = unsafe { *next.add(1) }; + predicted_small = unsafe { *prediction.add(1) } + .wrapping_add(u32::from(unsafe { *prediction.add(1) } > 0)); + continue; + } + if state.use_c_predict != 0 && match_index == predicted_large { + let prediction = unsafe { + state + .chain_table + .add(2 * (match_index.wrapping_sub(1) & bt_mask) as usize) + }; + unsafe { *larger = match_index }; + if match_index <= bt_low { + larger = &mut dummy; + break; + } + larger = next; + match_index = unsafe { *next }; + predicted_large = + unsafe { *prediction }.wrapping_add(u32::from(unsafe { *prediction } > 0)); + continue; + } + + let mut matched; + if !ext_dict || match_index.wrapping_add(match_length as u32) >= state.dict_limit { + matched = state.base.wrapping_add(match_index as usize); + match_length += unsafe { + count( + ip.wrapping_add(match_length), + matched.wrapping_add(match_length), + iend, + ) + }; + } else { + matched = state.dict_base.wrapping_add(match_index as usize); + match_length += unsafe { + count_2segments( + ip.wrapping_add(match_length), + matched.wrapping_add(match_length), + iend, + dict_end, + prefix_start, + ) + }; + if match_index.wrapping_add(match_length as u32) >= state.dict_limit { + matched = state.base.wrapping_add(match_index as usize); + } + } + + if match_length > best_length { + best_length = match_length; + if match_length > match_end_index.wrapping_sub(match_index) as usize { + match_end_index = match_index.wrapping_add(match_length as u32); + } + } + + if ip.wrapping_add(match_length) == iend { + break; + } + + if unsafe { *matched.add(match_length) < *ip.add(match_length) } { + unsafe { *smaller = match_index }; + common_smaller = match_length; + if match_index <= bt_low { + smaller = &mut dummy; + break; + } + smaller = unsafe { next.add(1) }; + match_index = unsafe { *next.add(1) }; + } else { + unsafe { *larger = match_index }; + common_larger = match_length; + if match_index <= bt_low { + larger = &mut dummy; + break; + } + larger = next; + match_index = unsafe { *next }; + } + } + + unsafe { + *smaller = 0; + *larger = 0; + } + let positions = if best_length > 384 { + min(192, (best_length - 384) as u32) + } else { + 0 + }; + max( + positions, + match_end_index.wrapping_sub(current.wrapping_add(HASH_READ_SIZE as u32)), + ) +} + +#[no_mangle] +pub unsafe extern "C" fn ZSTD_rust_opt_updateTreeInternal( + state: *mut ZSTD_RustOptTreeState, + ip: *const c_void, + iend: *const c_void, + mls: u32, + ext_dict: c_int, +) { + let state = unsafe { &*state }; + let input = ip.cast::(); + let input_end = iend.cast::(); + let target = unsafe { input.offset_from(state.base) as u32 }; + let mut index = unsafe { *state.next_to_update }; + while index < target { + let forward = unsafe { + insert_bt1( + state, + state.base.wrapping_add(index as usize), + input_end, + target, + mls, + ext_dict != 0, + ) + }; + debug_assert!(index < index.wrapping_add(forward)); + index = index.wrapping_add(forward); + } + unsafe { *state.next_to_update = target }; +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn tree_update_records_prefix_positions() { + let input = b"abcabcabcabcabcabcabcabc"; + let mut hashes = [0u32; 1 << 8]; + let mut chain = [0u32; 1 << 9]; + let mut next = 2u32; + let mut state = ZSTD_RustOptTreeState { + hash_table: hashes.as_mut_ptr(), + chain_table: chain.as_mut_ptr(), + base: input.as_ptr(), + dict_base: input.as_ptr(), + dict_limit: 0, + low_limit: 1, + loaded_dict_end: 0, + next_to_update: &mut next, + hash_log: 8, + chain_log: 9, + search_log: 4, + window_log: 10, + use_c_predict: 0, + }; + unsafe { + ZSTD_rust_opt_updateTreeInternal( + &mut state, + input.as_ptr().add(input.len() - HASH_READ_SIZE).cast(), + input.as_ptr().add(input.len()).cast(), + 4, + 0, + ); + } + assert_eq!(next, (input.len() - HASH_READ_SIZE) as u32); + assert!(hashes.iter().any(|&index| index != 0)); + assert!(chain.iter().all(|&index| index < input.len() as u32)); + } +}