use std::collections::BTreeSet;
pub type Trigram = [u8; 3];
pub fn extract(data: &[u8]) -> BTreeSet<Trigram> {
let mut set = BTreeSet::new();
if data.len() < 3 {
return set;
}
for w in data.windows(3) {
set.insert([w[0], w[1], w[2]]);
}
set
}
pub fn literal_trigrams(needle: &[u8]) -> Vec<Trigram> {
extract(needle).into_iter().collect()
}
#[derive(Debug, Default, Clone)]
pub struct TrigramQuery {
pub or_groups: Vec<Vec<Trigram>>,
pub and_clauses: Vec<Vec<Trigram>>,
}
impl TrigramQuery {
pub fn is_unconstrained(&self) -> bool {
let dnf_off = self.or_groups.is_empty() || self.or_groups.iter().any(|g| g.is_empty());
let cnf_off = self.and_clauses.is_empty() || self.and_clauses.iter().any(|c| c.is_empty());
dnf_off && cnf_off
}
pub fn from_literal(needle: &[u8]) -> TrigramQuery {
let tris = literal_trigrams(needle);
if tris.is_empty() {
TrigramQuery::default()
} else {
TrigramQuery {
or_groups: vec![tris],
and_clauses: Vec::new(),
}
}
}
pub fn from_literal_ci(needle: &[u8]) -> TrigramQuery {
if needle.len() < 3 {
return TrigramQuery::default();
}
let mut and_clauses: Vec<Vec<Trigram>> = Vec::new();
for w in needle.windows(3) {
if w.iter().all(|&b| ci_safe(b)) {
and_clauses.push(case_variants([w[0], w[1], w[2]]));
}
}
if and_clauses.is_empty() {
return TrigramQuery::default();
}
TrigramQuery {
or_groups: Vec::new(),
and_clauses,
}
}
}
fn ci_safe(b: u8) -> bool {
b < 0x80 && !matches!(b, b's' | b'S' | b'k' | b'K')
}
fn case_variants(w: Trigram) -> Vec<Trigram> {
let mut variants: Vec<Trigram> = vec![[0; 3]];
for (i, &b) in w.iter().enumerate() {
if b.is_ascii_alphabetic() {
let lo = b.to_ascii_lowercase();
let up = b.to_ascii_uppercase();
let mut next = Vec::with_capacity(variants.len() * 2);
for v in &variants {
let mut a = *v;
a[i] = lo;
let mut c = *v;
c[i] = up;
next.push(a);
next.push(c);
}
variants = next;
} else {
for v in &mut variants {
v[i] = b;
}
}
}
variants.sort_unstable();
variants.dedup();
variants
}
pub fn regex_trigrams(pattern: &str, case_insensitive: bool) -> TrigramQuery {
use regex_syntax::hir::literal::Extractor;
use regex_syntax::ParserBuilder;
let hir = match ParserBuilder::new()
.case_insensitive(case_insensitive)
.build()
.parse(pattern)
{
Ok(h) => h,
Err(_) => return TrigramQuery::default(),
};
let seq = Extractor::new().extract(&hir);
let mut or_groups: Vec<Vec<Trigram>> = Vec::new();
if let Some(lits) = seq.literals() {
for lit in lits {
let tris = literal_trigrams(lit.as_bytes());
if tris.is_empty() {
return TrigramQuery::default();
}
or_groups.push(tris);
}
}
if or_groups.is_empty() {
TrigramQuery::default()
} else {
TrigramQuery {
or_groups,
and_clauses: Vec::new(),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn extract_basic() {
let set = extract(b"abcd");
assert!(set.contains(b"abc"));
assert!(set.contains(b"bcd"));
assert_eq!(set.len(), 2);
}
#[test]
fn short_input_has_no_trigrams() {
assert!(extract(b"ab").is_empty());
assert!(literal_trigrams(b"ab").is_empty());
}
#[test]
fn literal_query_is_constrained() {
let q = TrigramQuery::from_literal(b"function");
assert!(!q.is_unconstrained());
assert!(TrigramQuery::from_literal(b"fn").is_unconstrained());
}
#[test]
fn regex_extracts_required_literal() {
let q = regex_trigrams("error_handler", false);
assert!(!q.is_unconstrained());
}
#[test]
fn case_insensitive_literal_is_constrained() {
let q = TrigramQuery::from_literal_ci(b"Foo");
assert!(!q.is_unconstrained());
assert_eq!(q.and_clauses.len(), 1);
let clause = &q.and_clauses[0];
assert!(clause.contains(b"foo"));
assert!(clause.contains(b"FOO"));
assert!(clause.contains(b"Foo"));
assert_eq!(clause.len(), 8);
assert!(TrigramQuery::from_literal_ci(b"fo").is_unconstrained());
}
#[test]
fn ci_skips_windows_with_non_ascii_bytes() {
let q = TrigramQuery::from_literal_ci("café".as_bytes());
assert_eq!(q.and_clauses.len(), 1);
let clause = &q.and_clauses[0];
assert!(clause.contains(b"caf"));
assert!(clause.contains(b"CAF"));
for clause in &q.and_clauses {
for tri in clause {
assert!(
tri.iter().all(|&b| b < 0x80),
"clause kept non-ASCII {tri:?}"
);
}
}
}
#[test]
fn ci_skips_kelvin_and_long_s_windows() {
let q = TrigramQuery::from_literal_ci(b"class");
assert_eq!(q.and_clauses.len(), 1);
assert!(q.and_clauses[0].contains(b"cla"));
assert!(TrigramQuery::from_literal_ci(b"list").is_unconstrained());
assert!(TrigramQuery::from_literal_ci(b"make").is_unconstrained());
}
#[test]
fn regex_ci_folds_kelvin_and_long_s() {
let ci = |pat: &str, hay: &str| {
regex::bytes::RegexBuilder::new(®ex::escape(pat))
.case_insensitive(true)
.build()
.unwrap()
.is_match(hay.as_bytes())
};
assert!(ci("k", "\u{212A}"), "/k/i should match KELVIN SIGN");
assert!(
ci("s", "\u{017F}"),
"/s/i should match LATIN SMALL LETTER LONG S"
);
assert!(!ci("a", "\u{00E5}"), "/a/i should not match 'å'");
}
#[test]
fn ci_filter_never_drops_a_match() {
let matching: &[(&str, &str)] = &[
("café", "a CAFÉ here"), ("café", "tiny café shop"), ("class", "MyClass {}"), ("foobar", "FOOBAR()"), ("make", "MAKEFILE"), ("kayak", "KAYAK"), ("string", "STRING s"), ];
for (needle, hay) in matching {
let re = regex::bytes::RegexBuilder::new(®ex::escape(needle))
.case_insensitive(true)
.build()
.unwrap();
assert!(
re.is_match(hay.as_bytes()),
"test setup: {needle:?} must match {hay:?}"
);
let q = TrigramQuery::from_literal_ci(needle.as_bytes());
assert!(
ci_filter_keeps(&q, hay.as_bytes()),
"filter wrongly dropped {hay:?} for needle {needle:?}"
);
}
}
fn ci_filter_keeps(q: &TrigramQuery, haystack: &[u8]) -> bool {
if q.is_unconstrained() {
return true;
}
let doc = extract(haystack);
q.and_clauses
.iter()
.all(|clause| clause.iter().any(|t| doc.contains(t)))
}
}