use aho_corasick::{AhoCorasick, AhoCorasickBuilder, MatchKind};
use crate::classify::errors::{ClassifyError, Result};
use crate::classify::rules::Rule;
#[derive(Clone, Copy)]
struct Boundaries {
left: bool,
right: bool,
}
fn is_word_char(c: char) -> bool {
c.is_alphanumeric() || c == '_'
}
fn boundaries_for(keyword: &str) -> Boundaries {
let mut chars = keyword.chars();
let first = chars.next();
let last = chars.next_back().or(first);
Boundaries {
left: first.is_some_and(is_word_char),
right: last.is_some_and(is_word_char),
}
}
fn match_is_bounded(haystack: &str, start: usize, end: usize, edges: Boundaries) -> bool {
if edges.left {
let before = haystack.get(..start).and_then(|s| s.chars().next_back());
if before.is_some_and(is_word_char) {
return false;
}
}
if edges.right {
let after = haystack.get(end..).and_then(|s| s.chars().next());
if after.is_some_and(is_word_char) {
return false;
}
}
true
}
pub struct ExactMatcher {
automaton: Option<AhoCorasick>,
pattern_rule_idx: Vec<usize>,
pattern_boundaries: Vec<Boundaries>,
rules: Vec<Rule>,
}
impl ExactMatcher {
pub fn new(rules: &[Rule]) -> Result<Self> {
let mut patterns: Vec<String> = Vec::new();
let mut pattern_rule_idx: Vec<usize> = Vec::new();
let mut pattern_boundaries: Vec<Boundaries> = Vec::new();
for (idx, rule) in rules.iter().enumerate() {
for kw in &rule.keywords {
if kw.is_empty() {
continue;
}
pattern_boundaries.push(boundaries_for(kw));
patterns.push(kw.clone());
pattern_rule_idx.push(idx);
}
}
let automaton = if patterns.is_empty() {
None
} else {
let ac = AhoCorasickBuilder::new()
.ascii_case_insensitive(true)
.match_kind(MatchKind::LeftmostLongest)
.build(&patterns)
.map_err(|e| ClassifyError::RuleLoad(format!("aho-corasick build: {e}")))?;
Some(ac)
};
Ok(Self {
automaton,
pattern_rule_idx,
pattern_boundaries,
rules: rules.to_vec(),
})
}
pub fn classify(&self, message: &str) -> Option<&Rule> {
let ac = self.automaton.as_ref()?;
let mut best: Option<&Rule> = None;
for m in ac.find_iter(message) {
let pattern = m.pattern().as_usize();
let edges = self.pattern_boundaries[pattern];
if !match_is_bounded(message, m.start(), m.end(), edges) {
continue;
}
let rule_idx = self.pattern_rule_idx[pattern];
let rule = &self.rules[rule_idx];
best = match best {
Some(prev) if prev.priority >= rule.priority => Some(prev),
_ => Some(rule),
};
}
best
}
}
#[cfg(test)]
mod tests {
use super::*;
fn rule(id: &str, keywords: &[&str]) -> Rule {
Rule {
id: id.into(),
category: id.into(),
subcategory: None,
keywords: keywords.iter().map(|k| (*k).to_string()).collect(),
patterns: vec![],
priority: 50,
confidence: 0.9,
}
}
#[test]
fn boundaries_follow_the_keywords_own_edges() {
let word = boundaries_for("rce");
assert!(word.left && word.right);
let trailing_punct = boundaries_for("cve-");
assert!(trailing_punct.left && !trailing_punct.right);
let leading_punct = boundaries_for("#fix");
assert!(!leading_punct.left && leading_punct.right);
let single = boundaries_for("v");
assert!(single.left && single.right);
}
#[test]
fn a_short_keyword_does_not_match_inside_a_longer_word() {
let m = ExactMatcher::new(&[rule("security", &["rce"])]).expect("build");
for msg in [
"SignalStore schema for source events",
"extract shared resource pool",
"bump the e-commerce sdk",
"enforce the rate limit",
] {
assert!(m.classify(msg).is_none(), "{msg}");
}
assert!(m.classify("mitigate the RCE the pentest found").is_some());
}
#[test]
fn a_keyword_ending_in_punctuation_still_matches_a_following_word() {
let m = ExactMatcher::new(&[rule("security", &["cve-"])]).expect("build");
assert!(m.classify("patch CVE-2024-1234").is_some());
assert!(m.classify("recve-2024-1234").is_none());
}
#[test]
fn multi_word_keywords_are_unaffected() {
let m = ExactMatcher::new(&[rule("bugfix", &["fix bug", "closes #"])]).expect("build");
assert!(m.classify("fix bug in the parser").is_some());
assert!(m.classify("closes #4331").is_some());
assert!(m.classify("prefix bug").is_none());
}
#[test]
fn a_multibyte_haystack_is_scored_without_panicking() {
let m = ExactMatcher::new(&[rule("security", &["rce"])]).expect("build");
assert!(m.classify("refactor — drop the résumé parser").is_none());
assert!(m.classify("— RCE —").is_some());
}
}