Skip to main content

stix_model/
registry.rs

1//! `ModelRegistry`: register consumer-supplied handling for STIX object types.
2//!
3//! # Example: validate and add a computed property
4//!
5//! ```
6//! use stix_model::ModelRegistry;
7//!
8//! let mut registry = ModelRegistry::new();
9//! registry.register_handler("x-acme-widget", |mut obj| {
10//!     let score = obj.get("risk_score").and_then(|v| v.as_i64()).unwrap_or(0);
11//!     obj["risk_band"] = serde_json::json!(if score > 80 { "high" } else { "low" });
12//!     Ok(obj)
13//! });
14//!
15//! let bundle = registry
16//!     .parse_bundle(r#"{"type":"bundle","objects":[
17//!         {"type":"x-acme-widget","id":"x-acme-widget--1","risk_score":90}
18//!     ]}"#)
19//!     .unwrap();
20//!
21//! use stix_model::ObjectView;
22//! assert_eq!(
23//!     bundle.objects[0].property("risk_band"),
24//!     Some(stix_model::StixValue::String("high".into()))
25//! );
26//! ```
27
28use std::collections::HashMap;
29use std::sync::Arc;
30
31use serde::de::DeserializeOwned;
32use serde_json::Value;
33
34use crate::bundle::Bundle;
35use crate::error::{ModelError, Result};
36use crate::object::StixObject;
37use crate::view::GenericObject;
38
39/// A handler that turns a raw JSON object of a registered type into a `StixObject`.
40type TypeHandler = Box<dyn Fn(Value) -> Result<StixObject> + Send + Sync>;
41
42/// Maps a STIX `type` to a consumer-supplied handler. Registered handlers take
43/// precedence over the built-in dispatch in [`StixObject::from_json`].
44///
45/// Two registration forms:
46/// - [`register_handler`](ModelRegistry::register_handler): a data-level
47///   `Value -> Result<Value>` validate/normalize hook (the form bindings bridge a
48///   host callable onto). The result is stored as a generic object.
49/// - [`register`](ModelRegistry::register): a typed Rust convenience that stores a
50///   `StixObject::Custom`.
51#[derive(Default)]
52pub struct ModelRegistry {
53    handlers: HashMap<String, TypeHandler>,
54}
55
56impl ModelRegistry {
57    /// An empty registry (parsing behaves like the built-in dispatch).
58    pub fn new() -> Self {
59        ModelRegistry::default()
60    }
61
62    /// Register a data-level validate/normalize hook for `type_name`. The hook may
63    /// reject the object (return `Err`) or return an enriched object (e.g. with a
64    /// computed property). The result is stored as a [`StixObject::Generic`].
65    pub fn register_handler<F>(&mut self, type_name: impl Into<String>, hook: F)
66    where
67        F: Fn(Value) -> Result<Value> + Send + Sync + 'static,
68    {
69        self.handlers.insert(
70            type_name.into(),
71            Box::new(move |value| {
72                let normalized = hook(value)?;
73                Ok(StixObject::Generic(GenericObject::from_json(normalized)?))
74            }),
75        );
76    }
77
78    /// Register a typed Rust struct for `type_name`. Objects of that type
79    /// deserialize into `T` and are stored as [`StixObject::Custom`], retrievable
80    /// with [`StixObject::downcast_ref`]. `T` only needs to implement
81    /// [`ObjectView`](crate::view::ObjectView) (plus `Serialize`/`Deserialize`).
82    pub fn register<T>(&mut self, type_name: impl Into<String>)
83    where
84        T: DeserializeOwned + crate::view::CustomObject + 'static,
85    {
86        self.handlers.insert(
87            type_name.into(),
88            Box::new(|value| {
89                let obj: T = serde_json::from_value(value)?;
90                Ok(StixObject::Custom(Arc::new(obj)))
91            }),
92        );
93    }
94
95    /// Parse a single JSON object, dispatching on its `type`: a registered handler
96    /// wins; otherwise the built-in dispatch applies.
97    pub fn parse_object(&self, value: Value) -> Result<StixObject> {
98        let type_ = value
99            .get("type")
100            .and_then(Value::as_str)
101            .ok_or_else(|| ModelError::InvalidObject("missing 'type' property".to_string()))?
102            .to_string();
103
104        match self.handlers.get(&type_) {
105            Some(handler) => handler(value),
106            None => StixObject::from_json(value),
107        }
108    }
109
110    /// Parse a bundle, routing every object through [`parse_object`](Self::parse_object).
111    pub fn parse_bundle(&self, json: &str) -> Result<Bundle> {
112        let value: Value = serde_json::from_str(json)?;
113        let type_ = value
114            .get("type")
115            .and_then(Value::as_str)
116            .unwrap_or_default()
117            .to_string();
118        if type_ != "bundle" {
119            return Err(ModelError::NotABundle(format!("type was '{type_}'")));
120        }
121        let id = value
122            .get("id")
123            .and_then(Value::as_str)
124            .map(|s| s.to_string());
125        let objects = match value.get("objects") {
126            Some(Value::Array(arr)) => arr
127                .iter()
128                .map(|o| self.parse_object(o.clone()))
129                .collect::<Result<Vec<_>>>()?,
130            _ => Vec::new(),
131        };
132        Ok(Bundle { type_, id, objects })
133    }
134}
135
136#[cfg(test)]
137mod tests {
138    use super::*;
139    use crate::error::ModelError;
140    use crate::object::{StixObject, TypedObject};
141    use crate::view::ObjectView;
142
143    #[test]
144    fn handler_validates_and_rejects() {
145        let mut reg = ModelRegistry::new();
146        reg.register_handler("x-thing", |obj| {
147            if obj.get("risk_score").is_none() {
148                return Err(ModelError::InvalidObject("missing risk_score".into()));
149            }
150            Ok(obj)
151        });
152        let ok = reg.parse_object(serde_json::json!({"type":"x-thing","id":"x--1","risk_score":1}));
153        assert!(ok.is_ok());
154        let bad = reg.parse_object(serde_json::json!({"type":"x-thing","id":"x--1"}));
155        assert!(matches!(bad, Err(ModelError::InvalidObject(_))));
156    }
157
158    #[test]
159    fn handler_adds_computed_property() {
160        let mut reg = ModelRegistry::new();
161        reg.register_handler("x-thing", |mut obj| {
162            let score = obj.get("risk_score").and_then(|v| v.as_i64()).unwrap_or(0);
163            obj["risk_band"] = serde_json::json!(if score > 80 { "high" } else { "low" });
164            Ok(obj)
165        });
166        let parsed = reg
167            .parse_object(serde_json::json!({"type":"x-thing","id":"x--1","risk_score":90}))
168            .unwrap();
169        // The enriched object is stored as data (a Generic object).
170        assert!(matches!(parsed, StixObject::Generic(_)));
171        assert_eq!(
172            parsed.property("risk_band"),
173            Some(crate::value::StixValue::String("high".into()))
174        );
175    }
176
177    #[test]
178    fn unregistered_types_use_builtin_dispatch() {
179        let reg = ModelRegistry::new();
180        // Built-in observed-data dispatch still works.
181        let od = reg
182            .parse_object(serde_json::json!({
183                "type":"observed-data","id":"observed-data--1",
184                "first_observed":"2020-01-01T00:00:00Z","last_observed":"2020-01-01T00:00:00Z",
185                "number_observed":1,"object_refs":[]
186            }))
187            .unwrap();
188        assert!(matches!(
189            od,
190            StixObject::Typed(TypedObject::ObservedData(_))
191        ));
192        // Unknown types fall back to Generic.
193        let g = reg
194            .parse_object(
195                serde_json::json!({"type":"ipv4-addr","id":"ipv4-addr--1","value":"1.2.3.4"}),
196            )
197            .unwrap();
198        assert!(matches!(g, StixObject::Generic(_)));
199    }
200
201    #[test]
202    fn parse_bundle_routes_objects_through_handlers() {
203        let mut reg = ModelRegistry::new();
204        reg.register_handler("x-thing", |mut obj| {
205            obj["seen"] = serde_json::json!(true);
206            Ok(obj)
207        });
208        let bundle = reg
209            .parse_bundle(
210                r#"{"type":"bundle","id":"bundle--1","objects":[
211                    {"type":"x-thing","id":"x--1"},
212                    {"type":"ipv4-addr","id":"ipv4-addr--1","value":"1.2.3.4"}
213                ]}"#,
214            )
215            .unwrap();
216        assert_eq!(bundle.objects.len(), 2);
217        assert_eq!(
218            bundle.objects[0].property("seen"),
219            Some(crate::value::StixValue::Bool(true))
220        );
221    }
222
223    #[test]
224    fn parse_bundle_rejects_non_bundle() {
225        let reg = ModelRegistry::new();
226        let err = reg
227            .parse_bundle(r#"{"type":"ipv4-addr","id":"x--1"}"#)
228            .unwrap_err();
229        assert!(matches!(err, ModelError::NotABundle(_)));
230    }
231
232    use crate::value::StixValue;
233
234    #[derive(Debug, serde::Serialize, serde::Deserialize)]
235    struct Widget {
236        #[serde(rename = "type")]
237        type_: String,
238        id: String,
239        risk: i64,
240    }
241
242    impl ObjectView for Widget {
243        fn id(&self) -> Option<&str> {
244            Some(&self.id)
245        }
246        fn type_(&self) -> Option<&str> {
247            Some(&self.type_)
248        }
249        fn property(&self, name: &str) -> Option<StixValue> {
250            match name {
251                "risk" => Some(StixValue::Integer(self.risk)),
252                _ => None,
253            }
254        }
255    }
256
257    #[test]
258    fn register_typed_yields_custom_and_downcasts() {
259        let mut reg = ModelRegistry::new();
260        reg.register::<Widget>("x-widget");
261        let parsed = reg
262            .parse_object(serde_json::json!({"type":"x-widget","id":"x-widget--1","risk":90}))
263            .unwrap();
264        assert!(matches!(parsed, StixObject::Custom(_)));
265        assert_eq!(parsed.property("risk"), Some(StixValue::Integer(90)));
266        let w = parsed.downcast_ref::<Widget>().expect("downcast");
267        assert_eq!(w.risk, 90);
268    }
269
270    #[test]
271    fn registered_handler_overrides_builtin() {
272        // A consumer may even override a core type's dispatch.
273        let mut reg = ModelRegistry::new();
274        reg.register::<Widget>("observed-data");
275        let parsed = reg
276            .parse_object(serde_json::json!({"type":"observed-data","id":"x--1","risk":7}))
277            .unwrap();
278        assert!(matches!(parsed, StixObject::Custom(_)));
279    }
280}