diff --git a/src/hunt.rs b/src/hunt.rs index afd68f9..ed7e7d6 100644 --- a/src/hunt.rs +++ b/src/hunt.rs @@ -45,6 +45,23 @@ pub struct Extensions { preconditions: Option>, } +fn push_precondition( + preconditions: &mut FxHashMap>, + rule_id: Uuid, + filter: Expression, +) { + preconditions.entry(rule_id).or_default().push(filter); +} + +fn preconditions_match(preconditions: &[Expression], document: &D) -> bool +where + D: TauDocument, +{ + preconditions + .iter() + .all(|filter| tau_engine::core::solve(filter, document)) +} + #[derive(Clone, Deserialize)] pub struct Group { #[serde(skip, default = "Uuid::new_v4")] @@ -205,8 +222,6 @@ impl HunterBuilder { for precondition in preconditions { for (rid, rule) in &rules { if let Rule::Sigma(sigma) = rule { - // FIXME: How do we handle multiple matches, for now we just take - // the latest, we chould probably just combine them into an AND? if precondition.for_.is_empty() { continue; } @@ -226,7 +241,7 @@ impl HunterBuilder { } } if matched { - preconds.insert(*rid, precondition.filter.clone()); + push_precondition(&mut preconds, *rid, precondition.filter.clone()); } } } @@ -304,8 +319,10 @@ impl HunterBuilder { .. } => { keys.extend(crate::ext::tau::extract_fields(filter)); - for precondition in preconditions.values() { - keys.extend(crate::ext::tau::extract_fields(precondition)); + for filters in preconditions.values() { + for filter in filters { + keys.extend(crate::ext::tau::extract_fields(filter)); + } } } } @@ -551,7 +568,7 @@ pub enum HuntKind { exclusions: HashSet>, filter: Expression, kind: RuleKind, - preconditions: FxHashMap, + preconditions: FxHashMap>, }, Rule { aggregate: Option, @@ -923,9 +940,10 @@ impl Hunter { if exclusions.contains(rid) { return None; } - if let Some(filter) = preconditions.get(rid) - && !tau_engine::core::solve(filter, &mapped) { - return None; + if let Some(filters) = preconditions.get(rid) + && !preconditions_match(filters, &mapped) + { + return None; } if rule.solve(&mapped) { Some((*rid, rule)) @@ -1144,3 +1162,45 @@ impl Hunter { Ok(false) } } + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + + #[test] + fn multiple_matching_preconditions_are_preserved() { + let rule_id = Uuid::new_v4(); + let mut preconditions = FxHashMap::default(); + + push_precondition(&mut preconditions, rule_id, Expression::Boolean(true)); + push_precondition(&mut preconditions, rule_id, Expression::Boolean(false)); + + assert_eq!(preconditions.get(&rule_id).unwrap().len(), 2); + } + + #[test] + fn matching_preconditions_use_and_semantics() { + let document = Value::from(json!({})); + + let passing = vec![Expression::Boolean(true), Expression::Boolean(true)]; + assert!(preconditions_match(&passing, &document)); + + let blocked = vec![Expression::Boolean(true), Expression::Boolean(false)]; + assert!(!preconditions_match(&blocked, &document)); + } + + #[test] + fn single_precondition_keeps_existing_behavior() { + let document = Value::from(json!({})); + + assert!(preconditions_match( + &[Expression::Boolean(true)], + &document + )); + assert!(!preconditions_match( + &[Expression::Boolean(false)], + &document + )); + } +}