use super::Matcher;
use super::multi::CompiledPatterns;
use crate::Match;
use crate::r#const::SIMD_CHUNK_BYTES;
use crate::sort::radix_sort_matches;
use alloc::vec::Vec;
const GATHER_BATCH: usize = 64;
#[derive(Debug, Default)]
pub(super) struct ResolvedScratch {
bytes: Vec<u8>,
spans: Vec<(u32, u32, u32)>,
hits: Vec<Match>,
}
impl Matcher {
pub fn match_list_resolved<T, F, const N: usize>(
&mut self,
items: &[T],
resolve: &F,
) -> Vec<Match>
where
F: Fn(&T, &mut [*const u8; N]) -> Option<(usize, u16)>,
{
let mut matches = Vec::new();
self.match_list_resolved_into(items, 0, resolve, &mut matches);
if self.config.sort.is_reversed() {
matches.reverse();
}
if !self.patterns.is_empty() && self.config.sort.is_by_score() {
radix_sort_matches(&mut matches);
}
matches
}
pub fn match_list_resolved_into<T, F, const N: usize>(
&mut self,
items: &[T],
item_index_offset: u32,
resolve: &F,
matches: &mut Vec<Match>,
) where
F: Fn(&T, &mut [*const u8; N]) -> Option<(usize, u16)>,
{
let mut scratch = ResolvedScratch::default();
self.match_list_resolved_into_with(
items,
item_index_offset,
resolve,
matches,
&mut scratch,
);
}
fn resolved_min_haystack_len(&self) -> usize {
match (&self.patterns, self.raw_patterns.as_slice()) {
(CompiledPatterns::Single(compiled), [pattern]) if !compiled.negated => compiled
.max_typos
.map(|max| pattern.needle.chars().count().saturating_sub(max as usize))
.unwrap_or(0),
_ => 0,
}
}
pub(super) fn match_list_resolved_into_with<T, F, const N: usize>(
&mut self,
items: &[T],
item_index_offset: u32,
resolve: &F,
matches: &mut Vec<Match>,
scratch: &mut ResolvedScratch,
) where
F: Fn(&T, &mut [*const u8; N]) -> Option<(usize, u16)>,
{
Self::guard_against_haystack_overflow(items.len(), item_index_offset);
let mut chunk_ptrs = [core::ptr::null::<u8>(); N];
if self.patterns.is_empty() {
matches.extend(
items
.iter()
.enumerate()
.filter(|(_, item)| resolve(item, &mut chunk_ptrs).is_some())
.map(|(i, _)| Match::from_index(i + item_index_offset as usize)),
);
return;
}
let min_haystack_len = self.resolved_min_haystack_len();
let ResolvedScratch { bytes, spans, hits } = scratch;
for (batch_idx, batch) in items.chunks(GATHER_BATCH).enumerate() {
bytes.clear();
spans.clear();
for (i, item) in batch.iter().enumerate() {
let Some((chunk_count, byte_len)) = resolve(item, &mut chunk_ptrs) else {
continue;
};
let len = byte_len as usize;
if len < min_haystack_len {
continue;
}
debug_assert!(
chunk_count == len.div_ceil(SIMD_CHUNK_BYTES),
"chunk_count {chunk_count} does not cover byte_len {len}"
);
debug_assert!(
chunk_count <= N,
"chunk_count {chunk_count} exceeds capacity {N}"
);
let start = bytes.len();
bytes.reserve(chunk_count * SIMD_CHUNK_BYTES);
unsafe {
let dst = bytes
.as_mut_ptr()
.add(start)
.cast::<[u8; SIMD_CHUNK_BYTES]>();
for (chunk, &ptr) in chunk_ptrs[..chunk_count].iter().enumerate() {
dst.add(chunk)
.write_unaligned(ptr.cast::<[u8; SIMD_CHUNK_BYTES]>().read_unaligned());
}
bytes.set_len(start + len);
}
let index = item_index_offset + (batch_idx * GATHER_BATCH + i) as u32;
spans.push((index, start as u32, len as u32));
}
if spans.is_empty() {
continue;
}
let haystacks: Vec<&str> = spans
.iter()
.map(|&(_, start, len)| {
let slice = &bytes[start as usize..(start + len) as usize];
unsafe { core::str::from_utf8_unchecked(slice) }
})
.collect();
hits.clear();
self.match_list_into(&haystacks, 0, hits);
matches.extend(hits.drain(..).map(|mut hit| {
hit.index = spans[hit.index as usize].0;
hit
}));
}
}
}
#[cfg(test)]
pub(crate) mod tests {
use super::*;
use crate::{Config, Matching, Pattern, SortStrategy};
use alloc::{format, string::String, vec};
#[derive(Clone)]
pub(crate) struct ChunkItem {
ptrs: Vec<*const u8>,
chunk_count: usize,
byte_len: u16,
}
unsafe impl Sync for ChunkItem {}
pub(crate) fn string_to_chunks(s: &str) -> ChunkItem {
let bytes = s.as_bytes();
let n_chunks = if bytes.is_empty() {
0
} else {
bytes.len().div_ceil(SIMD_CHUNK_BYTES)
};
let mut arena = vec![[0u8; SIMD_CHUNK_BYTES]; n_chunks];
for (i, chunk) in arena.iter_mut().enumerate() {
let start = i * SIMD_CHUNK_BYTES;
let take = SIMD_CHUNK_BYTES.min(bytes.len() - start);
chunk[..take].copy_from_slice(&bytes[start..start + take]);
}
let ptrs: Vec<*const u8> = arena.iter().map(|c| c.as_ptr()).collect();
core::mem::forget(arena);
ChunkItem {
ptrs,
chunk_count: n_chunks,
byte_len: bytes.len() as u16,
}
}
pub(crate) fn resolve_chunks<const N: usize>(
item: &ChunkItem,
ptrs_buf: &mut [*const u8; N],
) -> Option<(usize, u16)> {
ptrs_buf[..item.ptrs.len()].copy_from_slice(&item.ptrs);
Some((item.chunk_count, item.byte_len))
}
#[test]
fn resolved_matches_contiguous_parity() {
use proptest::prelude::*;
use proptest::test_runner::{Config as PropConfig, TestRunner};
let mut runner = TestRunner::new(PropConfig {
cases: if cfg!(miri) { 16 } else { 2000 },
..PropConfig::default()
});
let strategy = (
"[a-z]{2,12}", proptest::collection::vec("[a-z/_\\.]{5,80}", 1..30), (0u16..=8u16), proptest::bool::ANY, );
runner
.run(&strategy, |(needle, haystacks, max_typos, by_score)| {
let sort = if by_score {
SortStrategy::ScoreThenIndexAsc
} else {
SortStrategy::IndexAsc
};
let config = Config::default().max_typos(Some(max_typos)).sort(sort);
let contiguous = Matcher::new(needle.as_str(), &config).match_list(&haystacks);
let chunk_data: Vec<ChunkItem> =
haystacks.iter().map(|s| string_to_chunks(s)).collect();
let resolved = Matcher::new(needle.as_str(), &config)
.match_list_resolved(&chunk_data, &resolve_chunks::<8>);
prop_assert_eq!(
&contiguous,
&resolved,
"needle={:?} max_typos={} sort={:?}",
needle,
max_typos,
sort,
);
Ok(())
})
.unwrap();
}
#[test]
fn resolved_matches_contiguous_across_gather_batches() {
let haystacks: Vec<String> = (0..3 * GATHER_BATCH + 7)
.map(|i| {
if i % 97 == 0 {
format!("src/abc_{i}.rs")
} else {
format!("nomatch-{i}")
}
})
.collect();
let chunk_data: Vec<ChunkItem> = haystacks.iter().map(|s| string_to_chunks(s)).collect();
for sort in [SortStrategy::ScoreThenIndexAsc, SortStrategy::IndexDesc] {
let config = Config::default().sort(sort);
let contiguous = Matcher::new("abc", &config).match_list(&haystacks);
let resolved =
Matcher::new("abc", &config).match_list_resolved(&chunk_data, &resolve_chunks::<8>);
assert_eq!(contiguous, resolved, "sort={sort:?}");
assert!(!resolved.is_empty());
}
}
#[test]
fn resolved_supports_literal_and_multi_pattern() {
let haystacks = ["foo/bar", "bar/foo", "foo", "foobar", "qux"];
let chunk_data: Vec<ChunkItem> = haystacks.iter().map(|s| string_to_chunks(s)).collect();
let config = Config::default().sort(SortStrategy::IndexAsc);
let patterns = Pattern::parse_query("foo !^bar");
let contiguous = Matcher::from_patterns(&patterns, &config).match_list(&haystacks);
let resolved = Matcher::from_patterns(&patterns, &config)
.match_list_resolved(&chunk_data, &resolve_chunks::<2>);
assert_eq!(contiguous, resolved);
assert_eq!(
resolved.iter().map(|m| m.index).collect::<Vec<_>>(),
vec![0, 2, 3]
);
for matching in [
Matching::Exact,
Matching::Prefix,
Matching::Suffix,
Matching::Substring,
] {
let config = config.matching(matching);
let contiguous = Matcher::new("foo", &config).match_list(&haystacks);
let resolved =
Matcher::new("foo", &config).match_list_resolved(&chunk_data, &resolve_chunks::<2>);
assert_eq!(contiguous, resolved, "matching={matching:?}");
}
}
#[test]
fn resolved_skips_none_items_and_empty_needle_reports_present_items() {
let present = string_to_chunks("hello_world");
let items = [Some(present.clone()), None, Some(present)];
let resolve =
|item: &Option<ChunkItem>, ptrs_buf: &mut [*const u8; 4]| -> Option<(usize, u16)> {
item.as_ref()
.and_then(|item| resolve_chunks(item, ptrs_buf))
};
let config = Config::default().sort(SortStrategy::IndexAsc);
let matches = Matcher::new("hw", &config).match_list_resolved(&items, &resolve);
assert_eq!(
matches.iter().map(|m| m.index).collect::<Vec<_>>(),
vec![0, 2]
);
let matches = Matcher::new("", &config).match_list_resolved(&items, &resolve);
assert_eq!(
matches.iter().map(|m| m.index).collect::<Vec<_>>(),
vec![0, 2]
);
let mut matches = Vec::new();
Matcher::new("hw", &config).match_list_resolved_into(&items, 10, &resolve, &mut matches);
assert_eq!(
matches.iter().map(|m| m.index).collect::<Vec<_>>(),
vec![10, 12]
);
}
}