1use std::io::Write;
4use std::path::{Path, PathBuf};
5
6use crate::store::GraphBackend;
7use crate::{schema::EdgeType, Edge, Node};
8
9pub struct GraphMlBackend {
11 path: PathBuf,
12 nodes: Vec<Node>,
13 edges: Vec<Edge>,
14}
15
16impl GraphMlBackend {
17 pub fn open<P: AsRef<Path>>(path: P) -> Self {
19 Self {
20 path: path.as_ref().to_path_buf(),
21 nodes: Vec::new(),
22 edges: Vec::new(),
23 }
24 }
25
26 fn flush(&self) -> Result<(), GraphMlError> {
27 let mut f = std::fs::File::create(&self.path)?;
28 write_graphml(&mut f, &self.nodes, &self.edges)?;
29 Ok(())
30 }
31
32 fn load(&mut self) -> Result<(), GraphMlError> {
33 if !self.path.exists() {
34 return Ok(());
35 }
36 let content = std::fs::read_to_string(&self.path)?;
37 let (nodes, edges) = parse_graphml(&content)?;
38 self.nodes = nodes;
39 self.edges = edges;
40 Ok(())
41 }
42}
43
44#[derive(Debug, thiserror::Error)]
46pub enum GraphMlError {
47 #[error("IO error: {0}")]
48 Io(#[from] std::io::Error),
49 #[error("XML parse error: {0}")]
50 Xml(String),
51 #[error("Missing attribute: {0}")]
52 MissingAttr(String),
53}
54
55fn write_graphml(w: &mut impl Write, nodes: &[Node], edges: &[Edge]) -> Result<(), std::io::Error> {
56 writeln!(w, r#"<?xml version="1.0" encoding="UTF-8"?>"#)?;
57 writeln!(
58 w,
59 r#"<graphml xmlns="http://graphml.graphdrawing.org/xmlns">"#
60 )?;
61
62 writeln!(
64 w,
65 r#"<key id="kind" for="node" attr.name="kind" attr.type="string"/>"#
66 )?;
67 writeln!(
68 w,
69 r#"<key id="label" for="node" attr.name="label" attr.type="string"/>"#
70 )?;
71 writeln!(
72 w,
73 r#"<key id="payload" for="node" attr.name="payload" attr.type="string"/>"#
74 )?;
75
76 writeln!(
78 w,
79 r#"<key id="etype" for="edge" attr.name="type" attr.type="string"/>"#
80 )?;
81 writeln!(
82 w,
83 r#"<key id="epayload" for="edge" attr.name="payload" attr.type="string"/>"#
84 )?;
85
86 writeln!(w, r#"<graph id="G" edgedefault="directed">"#)?;
87
88 for n in nodes {
89 write!(w, r#"<node id="{}" >"#, xml_escape(&n.id))?;
90 writeln!(
91 w,
92 r#"<data key="kind">{}</data>"#,
93 xml_escape(&n.kind.to_string())
94 )?;
95 writeln!(w, r#"<data key="label">{}</data>"#, xml_escape(&n.label))?;
96 if let Some(ref p) = n.payload {
97 writeln!(
98 w,
99 r#"<data key="payload">{}</data>"#,
100 xml_escape(&p.to_string())
101 )?;
102 }
103 writeln!(w, "</node>")?;
104 }
105
106 for e in edges {
107 write!(
108 w,
109 r#"<edge source="{}" target="{}">"#,
110 xml_escape(&e.source_id),
111 xml_escape(&e.target_id)
112 )?;
113 writeln!(
114 w,
115 r#"<data key="etype">{}</data>"#,
116 xml_escape(&e.kind.to_string())
117 )?;
118 if let Some(ref p) = e.payload {
119 writeln!(
120 w,
121 r#"<data key="epayload">{}</data>"#,
122 xml_escape(&p.to_string())
123 )?;
124 }
125 writeln!(w, "</edge>")?;
126 }
127
128 writeln!(w, "</graph>")?;
129 writeln!(w, "</graphml>")?;
130 Ok(())
131}
132
133fn xml_escape(s: &str) -> String {
134 s.replace('&', "&")
135 .replace('<', "<")
136 .replace('>', ">")
137 .replace('"', """)
138 .replace('\'', "'")
139}
140
141fn parse_graphml(content: &str) -> Result<(Vec<Node>, Vec<Edge>), GraphMlError> {
142 let mut nodes = Vec::new();
143 let mut edges = Vec::new();
144
145 let node_re = regex::Regex::new(r#"(?s)<node\s+id="([^"]+)"[^>]*>(.*?)</node>"#)
153 .map_err(|e| GraphMlError::Xml(e.to_string()))?;
154
155 let edge_re =
156 regex::Regex::new(r#"(?s)<edge\s+source="([^"]+)"\s+target="([^"]+)"[^>]*>(.*?)</edge>"#)
157 .map_err(|e| GraphMlError::Xml(e.to_string()))?;
158
159 let data_re = regex::Regex::new(r#"(?s)<data\s+key="([^"]+)">(.*?)</data>"#)
160 .map_err(|e| GraphMlError::Xml(e.to_string()))?;
161
162 for cap in node_re.captures_iter(content) {
163 let id = cap[1].to_string();
164 let inner = &cap[2];
165 let mut kind = None;
166 let mut label = None;
167 let mut payload = None;
168 for dcap in data_re.captures_iter(inner) {
169 let key = &dcap[1];
170 let value = xml_unescape(&dcap[2]);
171 match key {
172 "kind" => kind = parse_node_type(&value),
173 "label" => label = Some(value),
174 "payload" => payload = serde_json::from_str(&value).ok(),
175 _ => {}
176 }
177 }
178 let kind = kind.unwrap_or(crate::schema::NodeType::Finding);
179 let label = label.unwrap_or_else(|| id.clone());
180 nodes.push(Node {
181 id,
182 kind,
183 label,
184 payload,
185 first_seen_ms: 0,
186 last_seen_ms: 0,
187 });
188 }
189
190 for cap in edge_re.captures_iter(content) {
191 let source_id = cap[1].to_string();
192 let target_id = cap[2].to_string();
193 let inner = &cap[3];
194 let mut kind = None;
195 let mut payload = None;
196 for dcap in data_re.captures_iter(inner) {
197 let key = &dcap[1];
198 let value = xml_unescape(&dcap[2]);
199 match key {
200 "etype" => kind = parse_edge_type(&value),
201 "epayload" => payload = serde_json::from_str(&value).ok(),
202 _ => {}
203 }
204 }
205 let kind = kind.unwrap_or(EdgeType::HasFinding);
206 edges.push(Edge {
207 source_id,
208 target_id,
209 kind,
210 payload,
211 first_seen_ms: 0,
212 last_seen_ms: 0,
213 });
214 }
215
216 Ok((nodes, edges))
217}
218
219fn xml_unescape(s: &str) -> String {
220 s.replace("<", "<")
221 .replace(">", ">")
222 .replace(""", "\"")
223 .replace("'", "'")
224 .replace("&", "&")
225}
226
227fn parse_node_type(s: &str) -> Option<crate::schema::NodeType> {
228 use crate::schema::NodeType;
229 match s {
230 "domain" => Some(NodeType::Domain),
231 "subdomain" => Some(NodeType::Subdomain),
232 "ip" => Some(NodeType::Ip),
233 "port" => Some(NodeType::Port),
234 "service" => Some(NodeType::Service),
235 "tech" => Some(NodeType::Tech),
236 "endpoint" => Some(NodeType::Endpoint),
237 "secret" => Some(NodeType::Secret),
238 "cloud" => Some(NodeType::Cloud),
239 "finding" => Some(NodeType::Finding),
240 _ => None,
241 }
242}
243
244fn parse_edge_type(s: &str) -> Option<EdgeType> {
245 match s {
246 "RESOLVES_TO" => Some(EdgeType::ResolvesTo),
247 "HOSTS" => Some(EdgeType::Hosts),
248 "RUNS" => Some(EdgeType::Runs),
249 "EXPOSES" => Some(EdgeType::Exposes),
250 "LEAKS" => Some(EdgeType::Leaks),
251 "MISCONFIGURED" => Some(EdgeType::Misconfigured),
252 "HAS_FINDING" => Some(EdgeType::HasFinding),
253 "HAS_SERVICE" => Some(EdgeType::HasService),
254 _ => None,
255 }
256}
257
258impl GraphBackend for GraphMlBackend {
259 type Error = GraphMlError;
260
261 fn init(&mut self) -> Result<(), Self::Error> {
262 self.load()?;
263 Ok(())
264 }
265
266 fn write_nodes(&mut self, nodes: &[Node]) -> Result<(), Self::Error> {
267 self.nodes.extend(nodes.iter().cloned());
268 self.flush()?;
269 Ok(())
270 }
271
272 fn write_edges(&mut self, edges: &[Edge]) -> Result<(), Self::Error> {
273 self.edges.extend(edges.iter().cloned());
274 self.flush()?;
275 Ok(())
276 }
277
278 fn read_nodes(&self) -> Result<Vec<Node>, Self::Error> {
279 Ok(self.nodes.clone())
280 }
281
282 fn read_edges(&self) -> Result<Vec<Edge>, Self::Error> {
283 Ok(self.edges.clone())
284 }
285
286 fn find_nodes_by_type(&self, kind: crate::schema::NodeType) -> Result<Vec<Node>, Self::Error> {
287 Ok(self
288 .nodes
289 .iter()
290 .filter(|n| n.kind == kind)
291 .cloned()
292 .collect())
293 }
294
295 fn neighbors(
296 &self,
297 node_id: &str,
298 edge_type: Option<EdgeType>,
299 ) -> Result<Vec<Edge>, Self::Error> {
300 Ok(self
301 .edges
302 .iter()
303 .filter(|e| {
304 e.source_id == node_id && edge_type.as_ref().map_or(true, |et| e.kind == *et)
305 })
306 .cloned()
307 .collect())
308 }
309
310 fn clear(&mut self) -> Result<(), Self::Error> {
311 self.nodes.clear();
312 self.edges.clear();
313 let _ = std::fs::remove_file(&self.path);
314 Ok(())
315 }
316}
317
318#[cfg(test)]
319mod tests {
320 use super::*;
321 use crate::schema::NodeType;
322 use tempfile::NamedTempFile;
323
324 #[test]
325 fn graphml_roundtrip() {
326 let file = NamedTempFile::new().unwrap();
327 let mut backend = GraphMlBackend::open(file.path());
328 backend.init().unwrap();
329
330 let node = Node::new("n1", NodeType::Domain, "example.com");
331 backend.write_nodes(&[node.clone()]).unwrap();
332
333 let edge = Edge::new("n1", "n2", EdgeType::ResolvesTo);
334 backend.write_edges(&[edge.clone()]).unwrap();
335
336 let mut backend2 = GraphMlBackend::open(file.path());
337 backend2.init().unwrap();
338
339 let nodes = backend2.read_nodes().unwrap();
340 assert_eq!(nodes.len(), 1);
341 assert_eq!(nodes[0].id, "n1");
342 assert_eq!(nodes[0].kind, NodeType::Domain);
343 assert_eq!(nodes[0].label, "example.com");
344
345 let edges = backend2.read_edges().unwrap();
346 assert_eq!(edges.len(), 1);
347 assert_eq!(edges[0].kind, EdgeType::ResolvesTo);
348 }
349
350 #[test]
351 fn xml_escape_unescape_roundtrip() {
352 let original = r#"<script>alert("xss")</script>"#;
353 let escaped = xml_escape(original);
354 let unescaped = xml_unescape(&escaped);
355 assert_eq!(original, unescaped);
356 }
357}