use rayon::iter::{IntoParallelRefIterator, ParallelIterator};
use regex_automata::dfa::regex::Regex;
use std::collections::{HashMap, HashSet};
use crate::RegexTrieError;
const SPECIALS: &str = ".?*+()[]{}";
type ScorerFuncType = Box<dyn Fn(&str, bool) -> usize>;
#[derive(Debug, Default)]
struct TrieNode {
children: HashMap<char, TrieNode>,
pattern_indices: Vec<usize>,
contains_non_regex_prefix: bool,
is_escaped: bool,
}
pub struct RegexTrie {
root: TrieNode,
compiled_patterns: Vec<(String, Regex, usize)>,
scorer: ScorerFuncType,
}
impl Default for RegexTrie {
fn default() -> Self {
Self::new_with_custom_scorer(Box::new(|pattern: &str, is_regex| {
if is_regex {
pattern.len()
} else {
0
}
}))
}
}
impl std::fmt::Debug for RegexTrie {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("RegexTrie")
.field("root", &self.root)
.field("compiled_patterns", &self.compiled_patterns)
.finish()
}
}
impl RegexTrie {
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn new_with_custom_scorer(scorer: ScorerFuncType) -> Self {
Self {
root: TrieNode::default(),
compiled_patterns: Vec::default(),
scorer,
}
}
pub fn from(patterns: &[String]) -> Result<Self, RegexTrieError> {
let mut trie = Self::new();
trie.insert_many(patterns)?;
Ok(trie)
}
pub fn from_with_scorer(
patterns: &[String],
scorer: ScorerFuncType,
) -> Result<Self, RegexTrieError> {
let mut trie = Self::new_with_custom_scorer(scorer);
trie.insert_many(patterns)?;
Ok(trie)
}
pub fn insert(&mut self, pattern: &str) -> Result<(), RegexTrieError> {
self.insert_many_lazy(&[pattern.to_string()])
}
pub fn insert_many(&mut self, patterns: &[String]) -> Result<(), RegexTrieError> {
self.insert_many_lazy(patterns)
}
fn insert_many_lazy(&mut self, patterns: &[String]) -> Result<(), RegexTrieError> {
let placeholder =
Regex::new("").map_err(|err| RegexTrieError::RegexCompilationFailed(Box::new(err)))?;
let mut to_compile = Vec::default();
for pattern in patterns {
let mut current_node = &mut self.root;
let mut previous_char = None;
let mut is_regex = false;
let mut chars = pattern.chars().peekable();
while let Some(ch) = chars.next() {
if ch == '\\' && matches!(chars.peek(), Some(next) if SPECIALS.contains(*next)) {
previous_char = Some(ch);
continue;
}
let mut is_escaped = false;
if SPECIALS.contains(ch) {
if previous_char == Some('\\') {
is_escaped = true;
} else {
is_regex = true;
break;
}
}
current_node = current_node.children.entry(ch).or_default();
current_node.is_escaped = is_escaped;
previous_char = Some(ch);
}
if is_regex {
let pattern_index = self.compiled_patterns.len();
let score = (self.scorer)(pattern, true);
self.compiled_patterns
.push((pattern.clone(), placeholder.clone(), score));
to_compile.push((pattern_index, pattern.clone()));
current_node.pattern_indices.push(pattern_index);
} else {
current_node.contains_non_regex_prefix = true;
}
}
let compiled = to_compile
.par_iter()
.map(|(idx, pattern)| {
let dfa = Regex::new(pattern)
.map_err(|err| RegexTrieError::RegexCompilationFailed(Box::new(err)))?;
Ok::<_, RegexTrieError>((*idx, dfa))
})
.collect::<Result<Vec<_>, _>>()?;
for (idx, dfa) in compiled {
self.compiled_patterns[idx].1 = dfa;
}
Ok(())
}
#[must_use]
pub fn find_matches(&self, input: &str) -> Vec<String> {
let mut candidate_indices = HashSet::new();
let mut current_node = &self.root;
for &index in ¤t_node.pattern_indices {
candidate_indices.insert(index);
}
let mut input_match_entirely = true;
let mut escaped_pattern = String::with_capacity(input.len());
for ch in input.chars() {
if let Some(node) = current_node.children.get(&ch) {
if node.is_escaped {
escaped_pattern.push('\\');
}
escaped_pattern.push(ch);
current_node = node;
for &index in ¤t_node.pattern_indices {
candidate_indices.insert(index);
}
} else {
input_match_entirely = false;
break;
}
}
let mut matching_patterns = Vec::new();
if input_match_entirely && current_node.contains_non_regex_prefix {
matching_patterns.push(escaped_pattern);
}
let input_bytes = input.as_bytes();
for index in candidate_indices {
let (pattern_str, dfa, _) = &self.compiled_patterns[index];
if let Some(m) = dfa.find(input_bytes) {
if m.start() == 0 && m.end() == input_bytes.len() {
matching_patterns.push(pattern_str.clone());
}
}
}
matching_patterns
}
#[must_use]
pub fn find_best_match(&self, input: &str) -> Option<String> {
let mut candidate_indices = HashSet::new();
let mut current_node = &self.root;
for &index in ¤t_node.pattern_indices {
candidate_indices.insert(index);
}
let mut input_match_entirely = true;
let mut escaped_pattern = String::with_capacity(input.len());
for ch in input.chars() {
if let Some(node) = current_node.children.get(&ch) {
if node.is_escaped {
escaped_pattern.push('\\');
}
escaped_pattern.push(ch);
current_node = node;
for &index in ¤t_node.pattern_indices {
candidate_indices.insert(index);
}
} else {
input_match_entirely = false;
break;
}
}
let mut best_match = None;
if input_match_entirely && current_node.contains_non_regex_prefix {
let score = (self.scorer)(&escaped_pattern, false);
best_match = Some((escaped_pattern, score));
}
let input_bytes = input.as_bytes();
for index in candidate_indices {
let (pattern_str, dfa, score) = &self.compiled_patterns[index];
if let Some(m) = dfa.find(input_bytes) {
if m.start() == 0 && m.end() == input_bytes.len() {
match &best_match {
Some((_, best_score)) => {
if score < best_score {
best_match = Some((pattern_str.clone(), *score));
}
}
None => best_match = Some((pattern_str.clone(), *score)),
}
}
}
}
best_match.map(|(pattern, _)| pattern)
}
}