1use stix::model::ModelRegistry;
4
5use crate::error::ErrorCode;
6use crate::error::FfiError;
7use crate::handles::{Bundle, MatchOutcome, Pattern};
8
9#[derive(Default)]
12pub struct Engine {
13 registry: ModelRegistry,
14}
15
16impl Engine {
17 pub fn new() -> Self {
19 Engine::default()
20 }
21
22 pub fn register_type(
27 &mut self,
28 type_name: &str,
29 hook: Box<dyn Fn(serde_json::Value) -> Result<serde_json::Value, String> + Send + Sync>,
30 ) {
31 self.registry.register_handler(type_name, move |value| {
32 hook(value).map_err(|message| {
33 stix::model::ModelError::InvalidObject(format!("[Validation] {message}"))
34 })
35 });
36 }
37
38 pub fn parse_pattern(&self, src: &str) -> Result<Pattern, FfiError> {
40 let inner = stix::parse(src)?;
41 Ok(Pattern::new(inner))
42 }
43
44 pub fn parse_bundle(&self, json: &str) -> Result<Bundle, FfiError> {
46 match self.registry.parse_bundle(json) {
47 Ok(inner) => Ok(Bundle::new(inner)),
48 Err(stix::model::ModelError::InvalidObject(m)) if m.starts_with("[Validation] ") => {
49 Err(FfiError::new(
50 ErrorCode::Validation,
51 m.trim_start_matches("[Validation] ").to_string(),
52 ))
53 }
54 Err(e) => Err(e.into()),
55 }
56 }
57
58 pub fn match_bundle(
60 &self,
61 pattern: &Pattern,
62 bundle: &Bundle,
63 ) -> Result<MatchOutcome, FfiError> {
64 let result = stix::matcher::match_bundle(pattern.inner(), bundle.inner())?;
65 Ok(MatchOutcome {
66 matched: result.is_match(),
67 observations: result.observations().iter().map(|&i| i as u64).collect(),
68 })
69 }
70}
71
72#[cfg(test)]
73mod tests {
74 use super::*;
75 use crate::error::ErrorCode;
76 use crate::error::ErrorCode as Code;
77
78 fn custom_bundle_json() -> &'static str {
79 r#"{"type":"bundle","id":"bundle--1","objects":[
80 {"type":"x-acme-widget","id":"x-acme-widget--1","risk_score":90},
81 {"type":"observed-data","id":"observed-data--1",
82 "first_observed":"2020-01-01T00:00:00Z","last_observed":"2020-01-01T00:00:00Z",
83 "number_observed":1,"object_refs":["x-acme-widget--1"]}
84 ]}"#
85 }
86
87 #[test]
88 fn register_type_adds_computed_property_and_matches() {
89 let mut engine = Engine::new();
90 engine.register_type(
91 "x-acme-widget",
92 Box::new(|mut obj| {
93 let score = obj.get("risk_score").and_then(|v| v.as_i64()).unwrap_or(0);
94 obj["risk_band"] = serde_json::json!(if score > 80 { "high" } else { "low" });
95 Ok(obj)
96 }),
97 );
98 let bundle = engine.parse_bundle(custom_bundle_json()).unwrap();
99 let pattern = engine
100 .parse_pattern("[x-acme-widget:risk_band = 'high']")
101 .unwrap();
102 assert!(engine.match_bundle(&pattern, &bundle).unwrap().matched);
103 }
104
105 #[test]
106 fn register_type_rejection_is_validation_error() {
107 let mut engine = Engine::new();
108 engine.register_type(
109 "x-acme-widget",
110 Box::new(|obj| {
111 if obj.get("risk_score").is_none() {
112 return Err("missing risk_score".to_string());
113 }
114 Ok(obj)
115 }),
116 );
117 let err = engine
118 .parse_bundle(r#"{"type":"bundle","objects":[{"type":"x-acme-widget","id":"x--1"}]}"#)
119 .unwrap_err();
120 assert_eq!(err.code, Code::Validation);
121 assert!(err.message.contains("missing risk_score"));
122 }
123
124 fn bundle_json() -> &'static str {
125 r#"{"type":"bundle","id":"bundle--1","objects":[
126 {"type":"ipv4-addr","id":"ipv4-addr--1","value":"198.51.100.5"},
127 {"type":"observed-data","id":"observed-data--1",
128 "first_observed":"2020-01-01T00:00:00Z","last_observed":"2020-01-01T00:00:00Z",
129 "number_observed":1,"object_refs":["ipv4-addr--1"]}
130 ]}"#
131 }
132
133 #[test]
134 fn parse_pattern_ok_and_err() {
135 let engine = Engine::new();
136 assert!(engine
137 .parse_pattern("[ipv4-addr:value = '1.2.3.4']")
138 .is_ok());
139 let err = engine.parse_pattern("[bad").unwrap_err();
140 assert_eq!(err.code, ErrorCode::Parse);
141 }
142
143 #[test]
144 fn parse_bundle_ok_and_non_bundle_err() {
145 let engine = Engine::new();
146 assert!(engine.parse_bundle(bundle_json()).is_ok());
147 let err = engine
148 .parse_bundle(r#"{"type":"ipv4-addr","id":"x--1"}"#)
149 .unwrap_err();
150 assert_eq!(err.code, ErrorCode::Model);
151 }
152
153 #[test]
154 fn match_bundle_match_and_non_match() {
155 let engine = Engine::new();
156 let bundle = engine.parse_bundle(bundle_json()).unwrap();
157
158 let hit = engine
159 .parse_pattern("[ipv4-addr:value = '198.51.100.5']")
160 .unwrap();
161 let outcome = engine.match_bundle(&hit, &bundle).unwrap();
162 assert!(outcome.matched);
163
164 let miss = engine
165 .parse_pattern("[ipv4-addr:value = '203.0.113.9']")
166 .unwrap();
167 assert!(!engine.match_bundle(&miss, &bundle).unwrap().matched);
168 }
169}