use crate::optimization::literal::LiteralSet;
use memchr::{memchr, memmem};
pub struct Prefilter {
strategy: PrefilterStrategy,
}
enum PrefilterStrategy {
SingleByte(u8),
SingleString(memmem::Finder<'static>),
MultiString {
searcher: aho_corasick::AhoCorasick,
patterns: Vec<String>,
},
None,
}
impl std::fmt::Debug for Prefilter {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match &self.strategy {
PrefilterStrategy::SingleByte(b) => f
.debug_struct("Prefilter")
.field("strategy", &"SingleByte")
.field("byte", b)
.finish(),
PrefilterStrategy::SingleString(_) => f
.debug_struct("Prefilter")
.field("strategy", &"SingleString")
.finish(),
PrefilterStrategy::MultiString { patterns, .. } => f
.debug_struct("Prefilter")
.field("strategy", &"MultiString")
.field("patterns", patterns)
.finish(),
PrefilterStrategy::None => f
.debug_struct("Prefilter")
.field("strategy", &"None")
.finish(),
}
}
}
impl Clone for Prefilter {
fn clone(&self) -> Self {
match &self.strategy {
PrefilterStrategy::SingleByte(b) => Prefilter {
strategy: PrefilterStrategy::SingleByte(*b),
},
PrefilterStrategy::SingleString(finder) => {
let pattern = finder.needle();
let new_finder = memmem::Finder::new(pattern).into_owned();
Prefilter {
strategy: PrefilterStrategy::SingleString(new_finder),
}
}
PrefilterStrategy::MultiString { patterns, .. } => {
let searcher = aho_corasick::AhoCorasick::builder()
.match_kind(aho_corasick::MatchKind::LeftmostLongest)
.build(patterns)
.expect("Failed to rebuild aho-corasick in clone");
Prefilter {
strategy: PrefilterStrategy::MultiString {
searcher,
patterns: patterns.clone(),
},
}
}
PrefilterStrategy::None => Prefilter {
strategy: PrefilterStrategy::None,
},
}
}
}
impl Prefilter {
pub fn from_literals(literals: &LiteralSet) -> Self {
if literals.is_empty() {
return Prefilter {
strategy: PrefilterStrategy::None,
};
}
if literals.literals.len() == 1 {
let text = &literals.literals[0].text;
if text.len() == 1 {
let byte = text.as_bytes()[0];
return Prefilter {
strategy: PrefilterStrategy::SingleByte(byte),
};
}
let finder = memmem::Finder::new(text).into_owned();
return Prefilter {
strategy: PrefilterStrategy::SingleString(finder),
};
}
if let Some(prefix) = literals.longest_common_prefix() {
if prefix.len() >= 3 {
if prefix.len() == 1 {
let byte = prefix.as_bytes()[0];
return Prefilter {
strategy: PrefilterStrategy::SingleByte(byte),
};
}
let finder = memmem::Finder::new(prefix).into_owned();
return Prefilter {
strategy: PrefilterStrategy::SingleString(finder),
};
}
}
if literals.literals.len() <= 100 {
let patterns: Vec<String> = literals
.literals
.iter()
.map(|lit| lit.text.clone())
.collect();
if let Ok(searcher) = aho_corasick::AhoCorasick::builder()
.match_kind(aho_corasick::MatchKind::LeftmostLongest)
.build(&patterns)
{
return Prefilter {
strategy: PrefilterStrategy::MultiString { searcher, patterns },
};
}
}
Prefilter {
strategy: PrefilterStrategy::None,
}
}
pub fn is_available(&self) -> bool {
!matches!(self.strategy, PrefilterStrategy::None)
}
pub fn find_candidate(&self, haystack: &[u8], from: usize) -> Option<usize> {
if from >= haystack.len() {
return None;
}
match &self.strategy {
PrefilterStrategy::SingleByte(byte) => {
memchr(*byte, &haystack[from..]).map(|pos| from + pos)
}
PrefilterStrategy::SingleString(finder) => {
finder.find(&haystack[from..]).map(|pos| from + pos)
}
PrefilterStrategy::MultiString { searcher, .. } => {
searcher.find(&haystack[from..]).map(|m| from + m.start())
}
PrefilterStrategy::None => Some(from),
}
}
pub fn candidates<'a>(&'a self, haystack: &'a [u8]) -> CandidateIter<'a> {
CandidateIter {
prefilter: self,
haystack,
pos: 0,
}
}
}
pub struct CandidateIter<'a> {
prefilter: &'a Prefilter,
haystack: &'a [u8],
pos: usize,
}
impl<'a> Iterator for CandidateIter<'a> {
type Item = usize;
fn next(&mut self) -> Option<usize> {
let candidate = self.prefilter.find_candidate(self.haystack, self.pos)?;
self.pos = candidate + 1;
Some(candidate)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::optimization::literal::{Literal, LiteralKind, LiteralSet};
#[test]
fn test_single_byte_prefilter() {
let mut literals = LiteralSet::empty();
literals.literals.push(Literal {
text: "@".to_string(),
is_exact: false,
});
literals.kind = LiteralKind::Inner;
let prefilter = Prefilter::from_literals(&literals);
assert!(prefilter.is_available());
let haystack = b"foo@example.com bar@test.org";
let candidates: Vec<usize> = prefilter.candidates(haystack).collect();
assert_eq!(candidates, vec![3, 19]);
}
#[test]
fn test_single_string_prefilter() {
let mut literals = LiteralSet::empty();
literals.literals.push(Literal {
text: "http".to_string(),
is_exact: false,
});
literals.kind = LiteralKind::Prefix;
let prefilter = Prefilter::from_literals(&literals);
assert!(prefilter.is_available());
let haystack = b"Visit https://example.com or http://test.org";
let candidates: Vec<usize> = prefilter.candidates(haystack).collect();
assert_eq!(candidates, vec![6, 29]);
}
#[test]
fn test_multi_string_prefilter() {
let mut literals = LiteralSet::empty();
literals.literals.push(Literal {
text: "foo".to_string(),
is_exact: true,
});
literals.literals.push(Literal {
text: "bar".to_string(),
is_exact: true,
});
literals.kind = LiteralKind::Prefix;
let prefilter = Prefilter::from_literals(&literals);
assert!(prefilter.is_available());
let haystack = b"foo bar baz foo";
let candidates: Vec<usize> = prefilter.candidates(haystack).collect();
assert_eq!(candidates, vec![0, 4, 12]);
}
}