pub type Trigram = [u8; 3];
pub type TrigramGroup = Vec<Trigram>;
pub type TrigramDnf = Vec<TrigramGroup>;
pub(crate) fn key_of(w: &[u8]) -> u32 {
(u32::from(w[0]) << 16) | (u32::from(w[1]) << 8) | u32::from(w[2])
}
pub(crate) fn tri_of(key: u32) -> Trigram {
[(key >> 16) as u8, (key >> 8) as u8, key as u8]
}
std::thread_local! {
static SEEN: std::cell::RefCell<Box<[u64]>> =
std::cell::RefCell::new(vec![0u64; 1 << 18].into_boxed_slice());
}
pub fn extract(data: &[u8]) -> Vec<Trigram> {
if data.len() < 3 {
return Vec::new();
}
SEEN.with(|seen| {
let mut seen = seen.borrow_mut();
let mut keys: Vec<u32> = Vec::new();
for w in data.windows(3) {
let k = key_of(w);
let word = (k >> 6) as usize;
let bit = 1u64 << (k & 63);
if seen[word] & bit == 0 {
seen[word] |= bit;
keys.push(k);
}
}
for &k in &keys {
seen[(k >> 6) as usize] = 0;
}
keys.sort_unstable();
keys.into_iter().map(tri_of).collect()
})
}
pub fn literal_trigrams(needle: &[u8]) -> Vec<Trigram> {
extract(needle)
}
#[derive(Debug, Default, Clone)]
pub struct TrigramQuery {
pub dnfs: Vec<TrigramDnf>,
}
impl TrigramQuery {
pub fn is_unconstrained(&self) -> bool {
!self.dnfs.iter().any(dnf_filters)
}
pub fn from_literal(needle: &[u8]) -> TrigramQuery {
let tris = literal_trigrams(needle);
if tris.is_empty() {
TrigramQuery::default()
} else {
TrigramQuery {
dnfs: vec![vec![tris]],
}
}
}
pub fn from_literal_ci(needle: &[u8]) -> TrigramQuery {
if needle.len() < 3 {
return TrigramQuery::default();
}
let mut dnfs: Vec<TrigramDnf> = Vec::new();
for w in needle.windows(3) {
if let Some(clause) = ci_window_trigrams([w[0], w[1], w[2]]) {
dnfs.push(clause.into_iter().map(|t| vec![t]).collect());
}
}
if dnfs.iter().any(dnf_filters) {
TrigramQuery { dnfs }
} else {
TrigramQuery::default()
}
}
}
pub fn dnf_filters(dnf: &TrigramDnf) -> bool {
!dnf.is_empty() && dnf.iter().all(|g| !g.is_empty())
}
type FoldForm = ([u8; 3], usize);
fn fold_forms(b: u8, out: &mut [FoldForm; 3]) -> Option<usize> {
if b >= 0x80 {
return None;
}
if !b.is_ascii_alphabetic() {
out[0] = ([b, 0, 0], 1);
return Some(1);
}
out[0] = ([b.to_ascii_lowercase(), 0, 0], 1);
out[1] = ([b.to_ascii_uppercase(), 0, 0], 1);
match b.to_ascii_lowercase() {
b's' => {
out[2] = ([0xC5, 0xBF, 0], 2);
Some(3)
}
b'k' => {
out[2] = ([0xE2, 0x84, 0xAA], 3);
Some(3)
}
_ => Some(2),
}
}
fn ci_window_trigrams(w: Trigram) -> Option<Vec<Trigram>> {
let mut forms = [[([0u8; 3], 0usize); 3]; 3];
let mut counts = [0usize; 3];
for i in 0..3 {
counts[i] = fold_forms(w[i], &mut forms[i])?;
}
let mut out: Vec<Trigram> = Vec::with_capacity(counts[0] * counts[1] * counts[2]);
let mut buf = [0u8; 9];
for a in &forms[0][..counts[0]] {
for b in &forms[1][..counts[1]] {
for c in &forms[2][..counts[2]] {
buf[..a.1].copy_from_slice(&a.0[..a.1]);
buf[a.1..a.1 + b.1].copy_from_slice(&b.0[..b.1]);
buf[a.1 + b.1..a.1 + b.1 + c.1].copy_from_slice(&c.0[..c.1]);
out.push([buf[0], buf[1], buf[2]]);
}
}
}
out.sort_unstable();
out.dedup();
Some(out)
}
pub fn regex_trigrams(pattern: &str, case_insensitive: bool) -> TrigramQuery {
use regex_syntax::hir::literal::{ExtractKind, Extractor};
use regex_syntax::ParserBuilder;
let hir = match ParserBuilder::new()
.case_insensitive(case_insensitive)
.build()
.parse(pattern)
{
Ok(h) => h,
Err(_) => return TrigramQuery::default(),
};
fn dnf_of(seq: ®ex_syntax::hir::literal::Seq) -> Option<TrigramDnf> {
let lits = seq.literals()?;
let mut dnf: TrigramDnf = Vec::with_capacity(lits.len());
for lit in lits {
let tris = literal_trigrams(lit.as_bytes());
if tris.is_empty() {
return None;
}
dnf.push(tris);
}
if dnf.is_empty() {
None
} else {
Some(dnf)
}
}
let prefix = dnf_of(&Extractor::new().extract(&hir));
let suffix = {
let mut ex = Extractor::new();
ex.kind(ExtractKind::Suffix);
dnf_of(&ex.extract(&hir))
};
let mut dnfs: Vec<TrigramDnf> = Vec::new();
if let Some(p) = prefix {
dnfs.push(p);
}
if let Some(s) = suffix {
if dnfs.first() != Some(&s) {
dnfs.push(s);
}
}
if dnfs.is_empty() {
TrigramQuery::default()
} else {
TrigramQuery { dnfs }
}
}
#[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 extract_is_sorted_and_deduped() {
let set = extract(b"abcabcabc");
let mut sorted = set.clone();
sorted.sort_unstable();
sorted.dedup();
assert_eq!(set, sorted, "extract must return sorted, distinct trigrams");
}
#[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());
assert_eq!(q.dnfs.len(), 1);
}
#[test]
fn regex_extracts_suffix_literals() {
let q = regex_trigrams(r"fn \w+_handler", false);
assert!(
!q.is_unconstrained(),
"suffix literal should constrain the query: {q:?}"
);
}
#[test]
fn case_insensitive_literal_is_constrained() {
let q = TrigramQuery::from_literal_ci(b"Foo");
assert!(!q.is_unconstrained());
assert_eq!(q.dnfs.len(), 1);
let tris: Vec<Trigram> = q.dnfs[0].iter().map(|g| g[0]).collect();
assert!(tris.contains(b"foo"));
assert!(tris.contains(b"FOO"));
assert!(tris.contains(b"Foo"));
assert_eq!(tris.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.dnfs.len(), 1);
let tris: Vec<Trigram> = q.dnfs[0].iter().map(|g| g[0]).collect();
assert!(tris.contains(b"caf"));
assert!(tris.contains(b"CAF"));
}
#[test]
fn ci_kelvin_and_long_s_windows_stay_constrained() {
for needle in [&b"class"[..], b"list", b"make", b"kayak"] {
let q = TrigramQuery::from_literal_ci(needle);
assert!(
!q.is_unconstrained(),
"{needle:?} should be constrained: {q:?}"
);
}
let q = TrigramQuery::from_literal_ci(b"las");
let tris: Vec<Trigram> = q.dnfs[0].iter().map(|g| g[0]).collect();
assert!(tris.contains(b"las"));
assert!(tris.contains(b"LAS"));
assert!(
tris.contains(&[b'l', b'a', 0xC5]),
"expected long-s prefix variant, got {tris:?}"
);
}
#[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 {}"), ("class", "cla\u{017F}s X"), ("foobar", "FOOBAR()"), ("make", "MAKEFILE"), ("make", "ma\u{212A}e it"), ("kayak", "KAYAK"), ("kayak", "kaya\u{212A} trip"), ("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!(
filter_keeps(&q, hay.as_bytes()),
"filter wrongly dropped {hay:?} for needle {needle:?}"
);
}
}
#[test]
fn literal_and_regex_filters_keep_their_matches() {
let hay = b"pub fn segment_writer_flush(x: u32) {}";
let lit = TrigramQuery::from_literal(b"segment_writer");
assert!(filter_keeps(&lit, hay));
let re = regex_trigrams(r"fn \w+_flush", false);
assert!(filter_keeps(&re, hay));
}
fn filter_keeps(q: &TrigramQuery, haystack: &[u8]) -> bool {
let doc = extract(haystack);
let has = |t: &Trigram| doc.binary_search(t).is_ok();
q.dnfs
.iter()
.filter(|d| dnf_filters(d))
.all(|dnf| dnf.iter().any(|group| group.iter().all(&has)))
}
}