Skip to main content

stix_ffi/
engine.rs

1//! The `Engine` handle: owns a registry, parses patterns/bundles, runs matches.
2
3use stix::model::ModelRegistry;
4
5use crate::error::ErrorCode;
6use crate::error::FfiError;
7use crate::handles::{Bundle, MatchOutcome, Pattern};
8
9/// The stateful facade handle. Holds the custom-model registry used by
10/// `parse_bundle`.
11#[derive(Default)]
12pub struct Engine {
13    registry: ModelRegistry,
14}
15
16impl Engine {
17    /// A new engine with an empty registry.
18    pub fn new() -> Self {
19        Engine::default()
20    }
21
22    /// Register a custom object type. `hook` validates and/or normalizes a raw JSON
23    /// object of `type_name`; returning `Err(message)` rejects it (surfaced as a
24    /// `Validation` error from `parse_bundle`). Returning an enriched object adds
25    /// computed properties, stored as data. The hook runs only at parse time.
26    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    /// Parse a STIX pattern string into a [`Pattern`] handle.
39    pub fn parse_pattern(&self, src: &str) -> Result<Pattern, FfiError> {
40        let inner = stix::parse(src)?;
41        Ok(Pattern::new(inner))
42    }
43
44    /// Parse a STIX bundle (consulting registered custom types) into a [`Bundle`].
45    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    /// Match a pattern against a bundle.
59    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}