1use 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
39type TypeHandler = Box<dyn Fn(Value) -> Result<StixObject> + Send + Sync>;
41
42#[derive(Default)]
52pub struct ModelRegistry {
53 handlers: HashMap<String, TypeHandler>,
54}
55
56impl ModelRegistry {
57 pub fn new() -> Self {
59 ModelRegistry::default()
60 }
61
62 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 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 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 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 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 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 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 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}