use fearless_simd::{Simd, SimdBase, SimdMask, u8x16, u8x32, u8x64};
pub(crate) const SCALAR_PREFIX_BYTES: usize = 16;
#[inline(always)]
pub(crate) fn load_u64_le(data: &[u8], offset: usize) -> u64 {
debug_assert!(offset <= u32::MAX as usize);
let offset = offset & u32::MAX as usize;
match data.get(offset..offset + 8) {
Some(chunk) => u64::from_le_bytes(chunk.try_into().unwrap_or([0; 8])),
None => 0,
}
}
#[inline(always)]
pub(crate) fn current_window(data: &[u8], cur_ix_masked: usize, max_length: usize) -> &[u8] {
data.get(cur_ix_masked..cur_ix_masked + max_length)
.unwrap_or_default()
}
#[inline(always)]
pub(crate) fn match_len_at<S: Simd>(simd: S, data: &[u8], prev_ix: usize, cur: &[u8]) -> usize {
let Some(left) = data.get(prev_ix..prev_ix + cur.len()) else {
return 0;
};
if let (Some(left_word), Some(cur_word)) = (left.first_chunk::<8>(), cur.first_chunk::<8>()) {
let difference = u64::from_le_bytes(*left_word) ^ u64::from_le_bytes(*cur_word);
if difference != 0 {
return difference.trailing_zeros() as usize >> 3;
}
return 8 + match_len_windows(simd, &left[8..], &cur[8..]);
}
match_len_windows(simd, left, cur)
}
#[inline(always)]
pub(crate) fn match_len_windows<S: Simd>(simd: S, left: &[u8], right: &[u8]) -> usize {
let limit = left.len().min(right.len());
let (left, right) = (&left[..limit], &right[..limit]);
let prefix = limit.min(SCALAR_PREFIX_BYTES) & !7;
let mut matched = match_len_words(&left[..prefix], &right[..prefix]);
if matched < prefix {
return matched;
}
let stride = native_vector_stride::<S>();
let whole_vectors = (limit - matched) - (limit - matched) % stride;
let vectored = match_len_native_vectors(simd, &left[matched..], &right[matched..]);
matched += vectored;
if vectored < whole_vectors {
return matched;
}
let whole_words = (limit - matched) & !7;
let tail = match_len_words(
&left[matched..matched + whole_words],
&right[matched..matched + whole_words],
);
matched += tail;
if tail < whole_words || matched == limit {
return matched;
}
if let (Some(left_last), Some(right_last)) = (left.last_chunk::<8>(), right.last_chunk::<8>()) {
let difference = u64::from_le_bytes(*left_last) ^ u64::from_le_bytes(*right_last);
return if difference == 0 {
limit
} else {
limit - 8 + (difference.trailing_zeros() as usize >> 3)
};
}
let mut bytes = 0usize;
while bytes < limit && left[bytes] == right[bytes] {
bytes += 1;
}
bytes
}
#[inline(always)]
fn match_len_words(left: &[u8], right: &[u8]) -> usize {
let (left_words, _) = left.as_chunks::<8>();
let (right_words, _) = right.as_chunks::<8>();
let mut matched = 0usize;
for (left_word, right_word) in left_words.iter().zip(right_words) {
let difference = u64::from_le_bytes(*left_word) ^ u64::from_le_bytes(*right_word);
if difference != 0 {
return matched + (difference.trailing_zeros() as usize >> 3);
}
matched += 8;
}
matched
}
macro_rules! vector_scan {
($name:ident, $vector:ident, $lanes:literal) => {
#[doc = concat!(
"Compares two windows ", stringify!($lanes), " bytes at a time.\n\n",
"Returns the number of leading equal bytes, or the length of the \
whole-vector prefix when every vector matched."
)]
#[inline(always)]
fn $name<S: Simd>(simd: S, left: &[u8], right: &[u8]) -> usize {
let (left_vectors, _) = left.as_chunks::<$lanes>();
let (right_vectors, _) = right.as_chunks::<$lanes>();
let mut matched = 0usize;
for (left_lanes, right_lanes) in left_vectors.iter().zip(right_vectors) {
let equal = $vector::<S>::load_array_ref(simd, left_lanes)
.simd_eq($vector::<S>::load_array_ref(simd, right_lanes));
if equal.any_false() {
return matched + equal.to_bitmask().trailing_ones() as usize;
}
matched += $lanes;
}
matched
}
};
}
vector_scan!(match_len_vectors_16, u8x16, 16);
vector_scan!(match_len_vectors_32, u8x32, 32);
vector_scan!(match_len_vectors_64, u8x64, 64);
#[inline(always)]
const fn native_vector_stride<S: Simd>() -> usize {
match <S::u8s as SimdBase<S>>::N {
16 => 16,
32 => 32,
64 => 64,
_ => 1,
}
}
#[inline(always)]
fn match_len_native_vectors<S: Simd>(simd: S, left: &[u8], right: &[u8]) -> usize {
match <S::u8s as SimdBase<S>>::N {
16 => match_len_vectors_16(simd, left, right),
32 => match_len_vectors_32(simd, left, right),
64 => match_len_vectors_64(simd, left, right),
_ => match_len_bytes(left, right),
}
}
#[inline(always)]
fn match_len_bytes(left: &[u8], right: &[u8]) -> usize {
left.iter()
.zip(right)
.take_while(|(left_byte, right_byte)| left_byte == right_byte)
.count()
}
#[inline(always)]
pub(crate) fn find_match_length<S: Simd>(
simd: S,
data: &[u8],
left: usize,
right: usize,
limit: usize,
) -> usize {
let (Some(left_window), Some(right_window)) = (data.get(left..), data.get(right..)) else {
return 0;
};
if limit > left_window.len() || limit > right_window.len() {
return 0;
}
scan_windows(simd, &left_window[..limit], &right_window[..limit])
}
#[inline(always)]
pub(crate) fn common_prefix_len(left: &[u8], right: &[u8], limit: usize) -> usize {
let limit = limit.min(left.len()).min(right.len());
let (left, right) = (&left[..limit], &right[..limit]);
let whole_words = limit & !7;
let matched = match_len_words(&left[..whole_words], &right[..whole_words]);
if matched < whole_words {
return matched;
}
matched + match_len_bytes(&left[whole_words..], &right[whole_words..])
}
#[inline(always)]
pub(crate) fn common_prefix_len_simd<S: Simd>(
simd: S,
left: &[u8],
right: &[u8],
limit: usize,
) -> usize {
let limit = limit.min(left.len()).min(right.len());
scan_windows(simd, &left[..limit], &right[..limit])
}
#[inline(always)]
fn scan_windows<S: Simd>(simd: S, left_window: &[u8], right_window: &[u8]) -> usize {
let limit = left_window.len().min(right_window.len());
let (left_window, right_window) = (&left_window[..limit], &right_window[..limit]);
let prefix = limit.min(SCALAR_PREFIX_BYTES) & !7;
let mut matched = match_len_words(&left_window[..prefix], &right_window[..prefix]);
if matched < prefix {
return matched;
}
let stride = native_vector_stride::<S>();
let whole_vectors = (limit - matched) - (limit - matched) % stride;
let vectored =
match_len_native_vectors(simd, &left_window[matched..], &right_window[matched..]);
matched += vectored;
if vectored < whole_vectors {
return matched;
}
let whole_words = (limit - matched) & !7;
let tail = match_len_words(
&left_window[matched..matched + whole_words],
&right_window[matched..matched + whole_words],
);
matched += tail;
if tail < whole_words {
return matched;
}
matched + match_len_bytes(&left_window[matched..], &right_window[matched..])
}
#[cfg(test)]
mod tests {
use super::*;
use fearless_simd::{Level, dispatch};
fn baseline(data: &[u8], left: usize, right: usize, limit: usize) -> usize {
(0..limit)
.take_while(|&i| data[left + i] == data[right + i])
.count()
}
fn measure(data: &[u8], left: usize, right: usize, limit: usize) -> usize {
let level = Level::new();
dispatch!(level, simd => find_match_length(simd, data, left, right, limit))
}
fn measure_fallback(data: &[u8], left: usize, right: usize, limit: usize) -> usize {
let level = Level::fallback();
dispatch!(level, simd => find_match_length(simd, data, left, right, limit))
}
#[test]
fn loads_read_little_endian_words_at_every_offset() {
let data: Vec<u8> = (1..=32u8).collect();
for offset in 0..=(data.len() - 8) {
let mut expected = [0u8; 8];
expected.copy_from_slice(&data[offset..offset + 8]);
assert_eq!(load_u64_le(&data, offset), u64::from_le_bytes(expected));
}
}
#[test]
fn loads_work_at_the_exact_end_of_the_slice() {
let data = [1u8, 2, 3, 4, 5, 6, 7, 8];
assert_eq!(load_u64_le(&data, 0), 0x0807_0605_0403_0201);
assert_eq!(load_u64_le(&data, 1), 0);
}
#[test]
fn loads_past_the_end_read_as_zero() {
let data = [1u8, 2, 3];
assert_eq!(load_u64_le(&data, 0), 0);
assert_eq!(load_u64_le(&data, 3), 0);
assert_eq!(load_u64_le(&data, 99), 0);
}
#[test]
fn reports_every_mismatch_position() {
let length = 300usize;
for mismatch in 0..length {
let mut data = vec![0u8; 2 * length];
data[length + mismatch] = 1;
let limit = length;
assert_eq!(measure(&data, 0, length, limit), mismatch);
assert_eq!(measure_fallback(&data, 0, length, limit), mismatch);
assert_eq!(baseline(&data, 0, length, limit), mismatch);
}
}
#[test]
fn respects_every_limit() {
let data = vec![7u8; 512];
for limit in 0..=256 {
assert_eq!(measure(&data, 0, 256, limit), limit);
assert_eq!(measure_fallback(&data, 0, 256, limit), limit);
}
}
#[test]
fn matches_the_baseline_on_pseudo_random_data() {
let mut data = vec![0u8; 4096];
let mut state = 0x1234_5678u32;
for byte in data.iter_mut() {
state = state.wrapping_mul(1_103_515_245).wrapping_add(12_345);
*byte = (state >> 16) as u8 & 0x03;
}
for left in (0..1024).step_by(7) {
for right in (1024..2048).step_by(13) {
let limit = 512;
let expected = baseline(&data, left, right, limit);
assert_eq!(measure(&data, left, right, limit), expected);
assert_eq!(measure_fallback(&data, left, right, limit), expected);
}
}
}
#[test]
fn a_window_that_does_not_fit_the_input_reports_no_match() {
let data = vec![0u8; 32];
for (left, right, limit) in [(0usize, 16usize, 17usize), (24, 0, 9), (33, 0, 1)] {
assert_eq!(measure(&data, left, right, limit), 0);
assert_eq!(measure_fallback(&data, left, right, limit), 0);
}
}
#[test]
fn the_reported_stride_is_the_one_the_vector_scan_takes() {
fn check<S: Simd>(simd: S) {
let stride = native_vector_stride::<S>();
assert!(matches!(stride, 1 | 16 | 32 | 64), "stride {stride}");
let data = vec![0xCDu8; 4 * stride];
let left = vec![0xCDu8; 4 * stride];
assert_eq!(
match_len_native_vectors(simd, &left, &data),
4 * stride,
"a fully matching window has to report every stride"
);
let window = 2 * stride - 1;
assert_eq!(
match_len_native_vectors(simd, &left[..window], &data[..window]),
window - window % stride
);
}
dispatch!(Level::new(), simd => check(simd));
dispatch!(Level::fallback(), simd => check(simd));
}
#[test]
fn common_prefixes_stop_at_the_first_difference_or_the_shortest_bound() {
let left: Vec<u8> = (0..100u8).collect();
let mut right = left.clone();
assert_eq!(common_prefix_len(&left, &right, 100), 100);
assert_eq!(common_prefix_len(&left, &right, 37), 37);
assert_eq!(common_prefix_len(&left[..20], &right, 100), 20);
assert_eq!(common_prefix_len(&[], &right, 100), 0);
for mismatch in [0usize, 3, 7, 8, 9, 15, 16, 31, 32, 63, 64, 99] {
right[mismatch] ^= 1;
assert_eq!(common_prefix_len(&left, &right, 100), mismatch);
for level in [Level::new(), Level::fallback()] {
let wide =
dispatch!(level, simd => common_prefix_len_simd(simd, &left, &right, 100));
assert_eq!(wide, mismatch, "{mismatch}");
let bounded =
dispatch!(level, simd => common_prefix_len_simd(simd, &left, &right[..50], 80));
assert_eq!(bounded, mismatch.min(50), "{mismatch}");
}
right[mismatch] ^= 1;
}
}
#[test]
fn overlapping_ranges_extend_like_the_reference() {
let data = vec![0xABu8; 64];
assert_eq!(measure(&data, 0, 1, 63), 63);
assert_eq!(measure_fallback(&data, 0, 1, 63), 63);
}
}