use fearless_simd::Simd;
use crate::compressor::core::shared::constants::{HASH_MUL32, WINDOW_GAP};
use crate::compressor::core::shared::dictionary::MAX_STATIC_DICTIONARY_MATCH_LEN;
use crate::compressor::core::shared::dictionary::all_matches::{self, INVALID_MATCH};
use crate::compressor::core::shared::match_len::find_match_length;
const BUCKET_BITS: u32 = 17;
const BUCKET_SIZE: usize = 1 << BUCKET_BITS;
pub(crate) const MAX_TREE_COMP_LENGTH: usize = 128;
const MAX_TREE_SEARCH_DEPTH: usize = 64;
pub(crate) const HASH_TYPE_LENGTH: usize = 4;
pub(crate) const STORE_LOOKAHEAD: usize = MAX_TREE_COMP_LENGTH;
const MAX_NUM_MATCHES: usize = 64 + MAX_TREE_SEARCH_DEPTH;
const _: () = assert!(MAX_NUM_MATCHES == 128);
const DENSE_TAIL: usize = 63;
const SPARSE_THRESHOLD: usize = 512;
const SPARSE_STRIDE: usize = 8;
#[derive(Copy, Clone, Debug, Default, Eq, PartialEq)]
pub(crate) struct BackwardMatch {
pub(crate) distance: u32,
length_and_code: u32,
}
impl BackwardMatch {
#[inline(always)]
pub(crate) const fn new(distance: usize, length: usize) -> Self {
Self {
distance: distance as u32,
length_and_code: (length as u32) << 5,
}
}
#[inline(always)]
pub(crate) const fn dictionary(distance: usize, length: usize, length_code: usize) -> Self {
let code = if length == length_code {
0
} else {
length_code
};
Self {
distance: distance as u32,
length_and_code: ((length as u32) << 5) | (code as u32),
}
}
#[inline(always)]
pub(crate) const fn length(&self) -> usize {
(self.length_and_code >> 5) as usize
}
#[inline(always)]
pub(crate) const fn length_code(&self) -> usize {
let code = self.length_and_code & 31;
if code != 0 {
code as usize
} else {
self.length()
}
}
}
pub(crate) struct BinaryTreeMatcher {
window_mask: usize,
invalid_pos: u32,
buckets: Vec<u32>,
forest: Vec<u32>,
}
impl BinaryTreeMatcher {
pub(crate) fn new(lgwin: usize) -> Self {
let window_mask = (1usize << lgwin) - 1;
let num_nodes = 1usize << lgwin;
Self {
window_mask,
invalid_pos: (0u32).wrapping_sub(window_mask as u32),
buckets: vec![0u32; BUCKET_SIZE],
forest: vec![0u32; 2 * num_nodes],
}
}
#[inline(always)]
fn hash(data: &[u8], offset: usize) -> usize {
let word = match data.get(offset..).and_then(<[u8]>::first_chunk::<4>) {
Some(chunk) => u32::from_le_bytes(*chunk),
None => 0,
};
(word.wrapping_mul(HASH_MUL32) >> (32 - BUCKET_BITS)) as usize
}
pub(crate) fn retained_bytes(&self) -> usize {
(self.buckets.capacity() + self.forest.capacity()) * size_of::<u32>()
}
pub(crate) fn prepare(&mut self) {
self.buckets.fill(self.invalid_pos);
}
#[inline(always)]
const fn left(&self, pos: usize) -> usize {
2 * (pos & self.window_mask)
}
#[inline(always)]
const fn right(&self, pos: usize) -> usize {
2 * (pos & self.window_mask) + 1
}
#[inline(always)]
fn link(&self, index: usize) -> u32 {
self.forest.get(index).copied().unwrap_or(self.invalid_pos)
}
#[inline(always)]
fn set_link(&mut self, index: usize, value: u32) {
if let Some(slot) = self.forest.get_mut(index) {
*slot = value;
}
}
#[expect(
clippy::too_many_arguments,
reason = "mirrors StoreAndFindMatches, whose parameters are all needed"
)]
#[inline(always)]
fn store_and_find_matches<S: Simd>(
&mut self,
simd: S,
data: &[u8],
cur_ix: usize,
ring_buffer_mask: usize,
max_length: usize,
max_backward: usize,
best_len: &mut usize,
matches: &mut Vec<BackwardMatch>,
collect: bool,
) {
let cur_ix_masked = cur_ix & ring_buffer_mask;
let max_comp_len = max_length.min(MAX_TREE_COMP_LENGTH);
let should_reroot_tree = max_length >= MAX_TREE_COMP_LENGTH;
let key = Self::hash(data, cur_ix_masked);
let mut prev_ix = self.buckets.get(key).copied().unwrap_or(self.invalid_pos) as usize;
let mut node_left = self.left(cur_ix);
let mut node_right = self.right(cur_ix);
let mut best_len_left = 0usize;
let mut best_len_right = 0usize;
if should_reroot_tree && let Some(slot) = self.buckets.get_mut(key) {
*slot = cur_ix as u32;
}
let mut depth_remaining = MAX_TREE_SEARCH_DEPTH;
loop {
let backward = cur_ix.wrapping_sub(prev_ix);
let prev_ix_masked = prev_ix & ring_buffer_mask;
if backward == 0 || backward > max_backward || depth_remaining == 0 {
if should_reroot_tree {
let invalid = self.invalid_pos;
self.set_link(node_left, invalid);
self.set_link(node_right, invalid);
}
return;
}
let cur_len = best_len_left.min(best_len_right);
let len = cur_len
+ find_match_length(
simd,
data,
prev_ix_masked + cur_len,
cur_ix_masked + cur_len,
max_length - cur_len,
);
if collect && len > *best_len {
*best_len = len;
matches.push(BackwardMatch::new(backward, len));
}
if len >= max_comp_len {
if should_reroot_tree {
let left = self.link(self.left(prev_ix));
let right = self.link(self.right(prev_ix));
self.set_link(node_left, left);
self.set_link(node_right, right);
}
return;
}
let cur_byte = data.get(cur_ix_masked + len).copied().unwrap_or(0);
let prev_byte = data.get(prev_ix_masked + len).copied().unwrap_or(0);
if cur_byte > prev_byte {
best_len_left = len;
if should_reroot_tree {
self.set_link(node_left, prev_ix as u32);
}
node_left = self.right(prev_ix);
prev_ix = self.link(node_left) as usize;
} else {
best_len_right = len;
if should_reroot_tree {
self.set_link(node_right, prev_ix as u32);
}
node_right = self.left(prev_ix);
prev_ix = self.link(node_right) as usize;
}
depth_remaining -= 1;
}
}
#[expect(
clippy::too_many_arguments,
reason = "mirrors FindAllMatches, whose parameters are all needed"
)]
pub(crate) fn find_all_matches<S: Simd>(
&mut self,
simd: S,
data: &[u8],
ring_buffer_mask: usize,
cur_ix: usize,
max_length: usize,
max_backward: usize,
dictionary_distance: usize,
max_distance: usize,
short_scan: usize,
matches: &mut Vec<BackwardMatch>,
) -> usize {
simd.vectorize(
#[inline(always)]
|| {
let start = matches.len();
let cur_ix_masked = cur_ix & ring_buffer_mask;
let mut best_len = 1usize;
let stop = cur_ix.saturating_sub(short_scan);
let mut index = cur_ix.wrapping_sub(1);
while index > stop && best_len <= 2 {
let backward = cur_ix.wrapping_sub(index);
if backward > max_backward {
break;
}
let prev_ix = index & ring_buffer_mask;
index = index.wrapping_sub(1);
if data.get(cur_ix_masked) != data.get(prev_ix)
|| data.get(cur_ix_masked + 1) != data.get(prev_ix + 1)
{
continue;
}
let len = find_match_length(simd, data, prev_ix, cur_ix_masked, max_length);
if len > best_len {
best_len = len;
matches.push(BackwardMatch::new(backward, len));
}
}
if best_len < max_length {
self.store_and_find_matches(
simd,
data,
cur_ix,
ring_buffer_mask,
max_length,
max_backward,
&mut best_len,
matches,
true,
);
}
if dictionary_distance > max_distance {
return matches.len() - start;
}
let min_len = best_len.saturating_add(1).max(4);
let mut found = [INVALID_MATCH; MAX_STATIC_DICTIONARY_MATCH_LEN + 1];
if all_matches::find_all(
data.get(cur_ix_masked..).unwrap_or_default(),
min_len,
max_length,
&mut found,
) {
let max_len = max_length.min(MAX_STATIC_DICTIONARY_MATCH_LEN);
for length in min_len..=max_len {
let Some(&packed) = found.get(length) else {
continue;
};
if packed >= INVALID_MATCH {
continue;
}
let distance = dictionary_distance + (packed >> 5) as usize + 1;
if distance <= max_distance {
matches.push(BackwardMatch::dictionary(
distance,
length,
(packed & 31) as usize,
));
}
}
}
matches.len() - start
},
)
}
pub(crate) fn store<S: Simd>(&mut self, simd: S, data: &[u8], mask: usize, ix: usize) {
simd.vectorize(
#[inline(always)]
|| {
let max_backward = self.window_mask - WINDOW_GAP + 1;
let mut best_len = 0usize;
let mut discard = Vec::new();
self.store_and_find_matches(
simd,
data,
ix,
mask,
MAX_TREE_COMP_LENGTH,
max_backward,
&mut best_len,
&mut discard,
false,
);
},
);
}
pub(crate) fn store_range<S: Simd>(
&mut self,
simd: S,
data: &[u8],
mask: usize,
start: usize,
end: usize,
) {
let mut dense_from = start;
if start + DENSE_TAIL <= end {
dense_from = end - DENSE_TAIL;
}
if start + SPARSE_THRESHOLD <= dense_from {
let mut sparse = start;
while sparse < dense_from {
self.store(simd, data, mask, sparse);
sparse += SPARSE_STRIDE;
}
}
for ix in dense_from..end {
self.store(simd, data, mask, ix);
}
}
pub(crate) fn stitch_to_previous_block<S: Simd>(
&mut self,
simd: S,
num_bytes: usize,
position: usize,
data: &[u8],
mask: usize,
) {
if num_bytes < HASH_TYPE_LENGTH - 1 || position < MAX_TREE_COMP_LENGTH {
return;
}
let start = position - MAX_TREE_COMP_LENGTH + 1;
let end = position.min(start + num_bytes);
for ix in start..end {
let max_backward = self.window_mask - (WINDOW_GAP - 1).max(position - ix);
let mut best_len = 0usize;
let mut discard = Vec::new();
self.store_and_find_matches(
simd,
data,
ix,
mask,
MAX_TREE_COMP_LENGTH,
max_backward,
&mut best_len,
&mut discard,
false,
);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use fearless_simd::{Level, dispatch};
fn walk(data: &[u8], end: usize, short_scan: usize) -> Vec<BackwardMatch> {
walk_with(Level::new(), data, end, short_scan)
}
fn walk_with(level: Level, data: &[u8], end: usize, short_scan: usize) -> Vec<BackwardMatch> {
let mut matcher = BinaryTreeMatcher::new(22);
matcher.prepare();
let mut out = Vec::new();
for position in 0..=end {
out.clear();
dispatch!(level, simd => matcher.find_all_matches(
simd,
data,
usize::MAX,
position,
data.len() - position,
position,
usize::MAX >> 1,
0,
short_scan,
&mut out,
));
}
out
}
fn true_length(data: &[u8], cur_ix: usize, distance: usize) -> usize {
let prev = cur_ix - distance;
data[cur_ix..]
.iter()
.zip(&data[prev..])
.take_while(|(a, b)| a == b)
.count()
}
fn brute_force_longest(data: &[u8], cur_ix: usize, stored_end: usize) -> usize {
let mut best = 0usize;
for prev in 0..cur_ix.min(stored_end) {
let length = true_length(data, cur_ix, cur_ix - prev);
best = best.max(length);
}
best
}
fn prefix_occurrences(data: &[u8], cur_ix: usize, stored_end: usize) -> usize {
let Some(prefix) = data.get(cur_ix..cur_ix + 4) else {
return usize::MAX;
};
(0..cur_ix.min(stored_end))
.filter(|&prev| data.get(prev..prev + 4) == Some(prefix))
.count()
}
fn assert_well_formed(data: &[u8], cur_ix: usize, matches: &[BackwardMatch]) {
for (index, m) in matches.iter().enumerate() {
let distance = m.distance as usize;
assert!(distance >= 1, "match {index} has distance zero");
assert!(
distance <= cur_ix,
"match {index} reaches back past the start of the input"
);
assert!(
true_length(data, cur_ix, distance) >= m.length(),
"match {index} claims {} bytes it does not have",
m.length()
);
}
for pair in matches.windows(2) {
assert!(
pair[1].length() > pair[0].length(),
"lengths were not strictly increasing: {matches:?}"
);
}
}
fn repeated() -> Vec<u8> {
let body: Vec<u8> = (0..300u32).map(|i| (i * 7 % 251) as u8 + 1).collect();
let mut data = body.clone();
data.extend_from_slice(&body);
data.extend_from_slice(&[0u8; 256]);
data
}
#[test]
fn a_match_packs_its_length_and_distance() {
let plain = BackwardMatch::new(1234, 56);
assert_eq!(plain.distance, 1234);
assert_eq!(plain.length(), 56);
assert_eq!(plain.length_code(), 56);
}
#[test]
fn a_dictionary_match_keeps_a_separate_length_code() {
let transformed = BackwardMatch::dictionary(9000, 12, 15);
assert_eq!(transformed.length(), 12);
assert_eq!(transformed.length_code(), 15);
let plain = BackwardMatch::dictionary(9000, 12, 12);
assert_eq!(plain.length(), 12);
assert_eq!(plain.length_code(), 12);
}
#[test]
fn an_empty_tree_finds_only_what_the_short_scan_sees() {
let data = repeated();
let mut matcher = BinaryTreeMatcher::new(22);
matcher.prepare();
let mut out = Vec::new();
let level = Level::new();
dispatch!(level, simd => matcher.find_all_matches(
simd, &data, usize::MAX, 300, data.len() - 300, 300,
usize::MAX >> 1, 0, 64, &mut out,
));
assert!(
out.iter().all(|m| (m.distance as usize) <= 64),
"the tree returned a match it never stored: {out:?}"
);
}
#[test]
fn a_repeat_is_found_at_its_full_length() {
let data = repeated();
let found = walk(&data, 300, 64);
assert_well_formed(&data, 300, &found);
let best = found.last().expect("the repeat was missed");
assert_eq!((best.distance, best.length()), (300, 300));
}
#[test]
fn the_longest_match_agrees_with_brute_force() {
let mut rng = 0x1234_5678_9ABC_DEF0u64;
let payload = 4000usize;
let mut data: Vec<u8> = Vec::new();
while data.len() < payload {
rng ^= rng << 13;
rng ^= rng >> 7;
rng ^= rng << 17;
data.push(b'a' + ((rng >> 24) % 8) as u8);
}
data.extend_from_slice(&[0u8; 256]);
let stored_end = payload - MAX_TREE_COMP_LENGTH;
let mut matcher = BinaryTreeMatcher::new(22);
matcher.prepare();
let mut out = Vec::new();
let level = Level::new();
let mut checked = 0usize;
for position in 0..payload {
out.clear();
dispatch!(level, simd => matcher.find_all_matches(
simd, &data, usize::MAX, position, payload - position, position,
usize::MAX >> 1, 0, 64, &mut out,
));
assert_well_formed(&data, position, &out);
let expected =
brute_force_longest(&data[..payload], position, stored_end).min(payload - position);
if !(4..=MAX_TREE_COMP_LENGTH).contains(&expected)
|| prefix_occurrences(&data, position, stored_end) > MAX_TREE_SEARCH_DEPTH
{
continue;
}
let found = out.last().map_or(0, BackwardMatch::length);
assert!(
found >= expected,
"position {position}: found {found}, brute force {expected}"
);
checked += 1;
}
assert!(checked > 500, "only {checked} positions were checked");
}
#[test]
fn the_short_scan_finds_a_two_byte_repeat_the_tree_cannot() {
let mut data = b"Xabcdab".to_vec();
data.extend_from_slice(&[0u8; 256]);
let mut matcher = BinaryTreeMatcher::new(22);
matcher.prepare();
let mut out = Vec::new();
let level = Level::new();
dispatch!(level, simd => matcher.find_all_matches(
simd, &data, usize::MAX, 5, data.len() - 5, 5,
usize::MAX >> 1, 0, 64, &mut out,
));
assert_eq!(out.len(), 1, "{out:?}");
assert_eq!((out[0].distance, out[0].length()), (4, 2));
}
#[test]
fn the_short_scan_length_bounds_how_far_back_it_looks() {
let mut data = b"Xab".to_vec();
data.extend_from_slice(&[b'z'; 27]);
data.extend_from_slice(b"ab");
data.extend_from_slice(&[0u8; 256]);
let cur = 30usize;
for (scan, expected) in [(16usize, 0usize), (64, 1)] {
let mut matcher = BinaryTreeMatcher::new(22);
matcher.prepare();
let mut out = Vec::new();
let level = Level::new();
dispatch!(level, simd => matcher.find_all_matches(
simd, &data, usize::MAX, cur, data.len() - cur, cur,
usize::MAX >> 1, 0, scan, &mut out,
));
assert_eq!(out.len(), expected, "scan {scan} found {out:?}");
}
}
#[test]
fn a_range_store_keeps_the_last_positions_dense() {
let body: Vec<u8> = (0..400u32).map(|i| (i * 31 % 251) as u8 + 1).collect();
let mut data = body.clone();
data.extend_from_slice(&body);
data.extend_from_slice(&[0u8; 256]);
let mut matcher = BinaryTreeMatcher::new(22);
matcher.prepare();
let level = Level::new();
dispatch!(level, simd => matcher.store_range(simd, &data, usize::MAX, 0, 400));
let inside = 400 - 10;
let mut out = Vec::new();
dispatch!(level, simd => matcher.find_all_matches(
simd, &data, usize::MAX, 400 + inside, data.len() - 400 - inside,
400 + inside, usize::MAX >> 1, 0, 0, &mut out,
));
assert!(
out.iter().any(|m| m.distance == 400),
"the dense tail was not stored: {out:?}"
);
let skipped = 10usize;
let mut out = Vec::new();
dispatch!(level, simd => matcher.find_all_matches(
simd, &data, usize::MAX, 400 + skipped, data.len() - 400 - skipped,
400 + skipped, usize::MAX >> 1, 0, 0, &mut out,
));
for found in &out {
assert_ne!(
found.distance, 400,
"a skipped position was stored anyway: {out:?}"
);
}
}
#[test]
fn a_long_range_store_strides_over_its_prefix() {
let body: Vec<u8> = (0..2000u32).map(|i| (i * 17 % 251) as u8 + 1).collect();
let mut data = body.clone();
data.extend_from_slice(&body);
data.extend_from_slice(&[0u8; 256]);
let mut matcher = BinaryTreeMatcher::new(22);
matcher.prepare();
let level = Level::new();
dispatch!(level, simd => matcher.store_range(simd, &data, usize::MAX, 0, 2000));
for (offset, stored) in [(800usize, true), (801, false)] {
let mut out = Vec::new();
dispatch!(level, simd => matcher.find_all_matches(
simd, &data, usize::MAX, 2000 + offset, data.len() - 2000 - offset,
2000 + offset, usize::MAX >> 1, 0, 0, &mut out,
));
assert_eq!(
out.iter().any(|m| m.distance == 2000),
stored,
"offset {offset} was {}stored",
if stored { "not " } else { "" }
);
}
}
#[test]
fn every_backend_finds_the_same_matches() {
let data = repeated();
let mut results = Vec::new();
for level in [Level::new(), Level::baseline(), Level::fallback()] {
results.push(walk_with(level, &data, 300, 64));
}
assert!(results.windows(2).all(|pair| pair[0] == pair[1]));
}
#[test]
fn stitching_stores_the_positions_before_the_boundary() {
let data = repeated();
let level = Level::new();
let mut early = BinaryTreeMatcher::new(22);
early.prepare();
dispatch!(level, simd => early.stitch_to_previous_block(
simd, 64, MAX_TREE_COMP_LENGTH - 1, &data, usize::MAX));
let mut out = Vec::new();
dispatch!(level, simd => early.find_all_matches(
simd, &data, usize::MAX, 300, data.len() - 300, 300,
usize::MAX >> 1, 0, 0, &mut out,
));
assert!(out.is_empty(), "stitching ran too early: {out:?}");
let mut late = BinaryTreeMatcher::new(22);
late.prepare();
dispatch!(level, simd => late.stitch_to_previous_block(
simd, 300, 300, &data, usize::MAX));
let mut out = Vec::new();
dispatch!(level, simd => late.find_all_matches(
simd, &data, usize::MAX, 550, data.len() - 550, 550,
usize::MAX >> 1, 0, 0, &mut out,
));
assert!(
out.iter().any(|m| m.distance == 300),
"stitching stored nothing: {out:?}"
);
}
#[test]
fn the_tree_search_depth_is_bounded() {
let mut data = vec![b'q'; 4096];
data.extend_from_slice(&[0u8; 256]);
let found = walk(&data, 2048, 64);
assert!(
found.len() <= MAX_NUM_MATCHES,
"collected {} matches",
found.len()
);
}
}