use aho_corasick::{AhoCorasick, MatchKind};
use forbidden_regex::{CompileError, RegexSet};
use std::collections::HashMap;
use crate::rule::frx::{parse_runtime_rules, LoadError, RuntimeRuleInput, RuntimeRuleKind};
const SHORT_LITERAL_THRESHOLD: usize = 8;
#[derive(Clone, Debug, PartialEq, Eq)]
pub(crate) struct LiteralGroup {
pub(crate) bytes: Vec<u8>,
pub(crate) rule_ids: Vec<usize>,
}
#[derive(Debug)]
pub(crate) enum RuntimeMatcherError {
Load(
LoadError,
),
Regex {
index: usize,
reason: CompileError,
},
RegexSet {
reason: CompileError,
},
LiteralBuild,
InvalidArtifact,
}
pub(crate) struct RuntimeRules {
names: Vec<Option<String>>,
literal_groups: Vec<LiteralGroup>,
literal_matcher: Option<AhoCorasick>,
regex_set: Option<RegexSet>,
regex_rule_ids: Vec<usize>,
}
impl RuntimeRules {
pub(crate) fn compile(text: &str) -> Result<Self, RuntimeMatcherError> {
let inputs = parse_runtime_rules(text).map_err(RuntimeMatcherError::Load)?;
return Self::from_inputs(inputs)
}
pub(crate) fn from_artifact(
names: Vec<Option<String>>,
literal_groups: Vec<LiteralGroup>,
regex_set: Option<RegexSet>,
regex_rule_ids: Vec<usize>,
) -> Result<Self, RuntimeMatcherError> {
validate_mappings(&names, &literal_groups, regex_set.as_ref(), ®ex_rule_ids)?;
let literal_matcher = build_literal_matcher(&literal_groups)?;
return Ok(Self {
names,
literal_groups,
literal_matcher,
regex_set,
regex_rule_ids,
})
}
fn from_inputs(inputs: Vec<RuntimeRuleInput>) -> Result<Self, RuntimeMatcherError> {
let names: Vec<Option<String>> = inputs
.iter()
.map(|input| return input.name.clone())
.collect();
let mut literal_groups: Vec<LiteralGroup> = Vec::new();
let mut literal_indices: HashMap<Vec<u8>, usize> = HashMap::new();
let mut regex_patterns: Vec<String> = Vec::new();
let mut regex_rule_ids: Vec<usize> = Vec::new();
for (rule_id, input) in inputs.into_iter().enumerate() {
match input.kind {
RuntimeRuleKind::ExactLiteral(bytes) => {
if let Some(&group_index) = literal_indices.get(&bytes) {
literal_groups[group_index].rule_ids.push(rule_id);
} else {
let group_index = literal_groups.len();
literal_indices.insert(bytes.clone(), group_index);
literal_groups.push(LiteralGroup { bytes, rule_ids: vec![rule_id] });
}
}
RuntimeRuleKind::RestrictedRegex(pattern) => {
if let Err(reason) = RegexSet::new(std::slice::from_ref(&pattern)) {
return Err(RuntimeMatcherError::Regex { index: rule_id, reason });
}
regex_patterns.push(pattern);
regex_rule_ids.push(rule_id);
}
}
}
let regex_set = if regex_patterns.is_empty() {
None
} else {
Some(RegexSet::new(®ex_patterns).map_err(|reason| {
return RuntimeMatcherError::RegexSet { reason }
})?)
};
return Self::from_artifact(names, literal_groups, regex_set, regex_rule_ids)
}
pub(crate) fn names(&self) -> &[Option<String>] {
return &self.names
}
pub(crate) fn literal_groups(&self) -> &[LiteralGroup] {
return &self.literal_groups
}
pub(crate) fn regex_set(&self) -> Option<&RegexSet> {
return self.regex_set.as_ref()
}
pub(crate) fn regex_rule_ids(&self) -> &[usize] {
return &self.regex_rule_ids
}
pub(crate) fn len(&self) -> usize {
return self.names.len()
}
pub(crate) fn line_matches(&self, buf: &[u8], starts: &[usize]) -> Vec<(usize, usize)> {
let mut hits = self.regex_pairs(buf, starts);
for line_index in 0..starts.len() {
let start = starts[line_index];
let end = line_end(buf, starts, line_index);
if end == start {
continue;
}
self.append_literal_matches(&buf[start..end], line_index, &mut hits);
}
hits.sort_unstable();
hits.dedup();
return hits
}
fn regex_pairs(&self, buf: &[u8], starts: &[usize]) -> Vec<(usize, usize)> {
let Some(regex_set) = &self.regex_set else {
return Vec::new();
};
return regex_set
.line_matches(buf, starts)
.into_iter()
.map(|(line, local_rule)| return (line, self.regex_rule_ids[local_rule]))
.collect()
}
fn append_literal_matches(
&self,
line: &[u8],
line_index: usize,
hits: &mut Vec<(usize, usize)>,
) {
let Some(matcher) = &self.literal_matcher else {
return;
};
for found in matcher.find_overlapping_iter(line) {
let group_index = found.pattern().as_usize();
let group = &self.literal_groups[group_index];
if literal_boundaries_match(line, found.start(), found.end(), &group.bytes) {
hits.extend(group.rule_ids.iter().map(|&rule_id| return (line_index, rule_id)));
}
}
}
}
fn build_literal_matcher(
groups: &[LiteralGroup],
) -> Result<Option<AhoCorasick>, RuntimeMatcherError> {
if groups.is_empty() {
return Ok(None);
}
let patterns: Vec<&[u8]> = groups.iter().map(|group| return group.bytes.as_slice()).collect();
let matcher = AhoCorasick::builder()
.match_kind(MatchKind::Standard)
.build(&patterns)
.map_err(|_| return RuntimeMatcherError::LiteralBuild)?;
return Ok(Some(matcher))
}
fn validate_mappings(
names: &[Option<String>],
groups: &[LiteralGroup],
regex_set: Option<&RegexSet>,
regex_rule_ids: &[usize],
) -> Result<(), RuntimeMatcherError> {
if regex_set.map_or(0, RegexSet::len) != regex_rule_ids.len() {
return Err(RuntimeMatcherError::InvalidArtifact);
}
let rule_count = names.len();
let mut seen_rule_ids = vec![false; rule_count];
for &rule_id in regex_rule_ids {
if rule_id >= rule_count || seen_rule_ids[rule_id] {
return Err(RuntimeMatcherError::InvalidArtifact);
}
seen_rule_ids[rule_id] = true;
}
let mut seen_literals: std::collections::HashSet<&[u8]> = std::collections::HashSet::new();
for group in groups {
if group.bytes.is_empty()
|| group.rule_ids.is_empty()
|| !seen_literals.insert(&group.bytes)
{
return Err(RuntimeMatcherError::InvalidArtifact);
}
for &rule_id in &group.rule_ids {
if rule_id >= rule_count || seen_rule_ids[rule_id] {
return Err(RuntimeMatcherError::InvalidArtifact);
}
seen_rule_ids[rule_id] = true;
}
}
if seen_rule_ids.iter().any(|seen| return !seen) {
return Err(RuntimeMatcherError::InvalidArtifact);
}
return Ok(())
}
fn line_end(buf: &[u8], starts: &[usize], line_index: usize) -> usize {
let start = starts[line_index];
let mut end = starts.get(line_index + 1).copied().unwrap_or(buf.len());
if end > start && buf[end - 1] == b'\n' {
end -= 1;
}
if end > start && buf[end - 1] == b'\r' {
end -= 1;
}
return end
}
fn is_word_byte(byte: u8) -> bool {
return byte.is_ascii_alphanumeric() || byte == b'_'
}
fn literal_boundaries_match(
line: &[u8],
start: usize,
end: usize,
literal: &[u8],
) -> bool {
if literal.len() >= SHORT_LITERAL_THRESHOLD {
return true;
}
let left_ok = !is_word_byte(literal[0]) || start == 0 || !is_word_byte(line[start - 1]);
let right_ok = !is_word_byte(literal[literal.len() - 1])
|| end == line.len()
|| !is_word_byte(line[end]);
return left_ok && right_ok
}
impl std::fmt::Display for RuntimeMatcherError {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Load(error) => return write!(formatter, "{error}"),
Self::Regex { index, reason } => return write!(formatter, "rule {index}: {reason}"),
Self::RegexSet { reason } => return write!(formatter, "runtime regex set: {reason}"),
Self::LiteralBuild => return formatter.write_str("runtime literal matcher build failed"),
Self::InvalidArtifact => return formatter.write_str("runtime matcher artifact is invalid"),
}
}
}
impl std::error::Error for RuntimeMatcherError {}
#[cfg(test)]
#[path = "runtime_matcher_tests.rs"]
mod tests;