gossan_graph/store/
memory.rs1use crate::schema::{EdgeType, NodeType};
9use crate::{Edge, Node};
10
11use super::GraphBackend;
12
13#[derive(Debug)]
15pub struct MemoryError(String);
16
17impl std::fmt::Display for MemoryError {
18 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
19 write!(f, "{}", self.0)
20 }
21}
22
23impl std::error::Error for MemoryError {}
24
25#[derive(Debug, Default, Clone)]
27pub struct MemoryStore {
28 nodes: Vec<Node>,
29 edges: Vec<Edge>,
30}
31
32impl MemoryStore {
33 #[must_use]
35 pub fn new() -> Self {
36 Self::default()
37 }
38}
39
40impl GraphBackend for MemoryStore {
41 type Error = MemoryError;
42
43 fn init(&mut self) -> Result<(), Self::Error> {
44 Ok(())
45 }
46
47 fn write_nodes(&mut self, nodes: &[Node]) -> Result<(), Self::Error> {
48 self.nodes.extend_from_slice(nodes);
49 Ok(())
50 }
51
52 fn write_edges(&mut self, edges: &[Edge]) -> Result<(), Self::Error> {
53 self.edges.extend_from_slice(edges);
54 Ok(())
55 }
56
57 fn read_nodes(&self) -> Result<Vec<Node>, Self::Error> {
58 Ok(self.nodes.clone())
59 }
60
61 fn read_edges(&self) -> Result<Vec<Edge>, Self::Error> {
62 Ok(self.edges.clone())
63 }
64
65 fn find_nodes_by_type(&self, kind: NodeType) -> Result<Vec<Node>, Self::Error> {
66 Ok(self
67 .nodes
68 .iter()
69 .filter(|n| n.kind == kind)
70 .cloned()
71 .collect())
72 }
73
74 fn neighbors(
75 &self,
76 node_id: &str,
77 edge_type: Option<EdgeType>,
78 ) -> Result<Vec<Edge>, Self::Error> {
79 Ok(self
80 .edges
81 .iter()
82 .filter(|e| e.source_id == node_id)
83 .filter(|e| edge_type.map_or(true, |t| e.kind == t))
84 .cloned()
85 .collect())
86 }
87
88 fn clear(&mut self) -> Result<(), Self::Error> {
89 self.nodes.clear();
90 self.edges.clear();
91 Ok(())
92 }
93}
94
95#[cfg(test)]
96mod tests {
97 use super::*;
98 use crate::schema::{EdgeType, NodeType};
99
100 fn sample_node(id: &str, kind: NodeType) -> Node {
101 Node::new(id, kind, id)
102 }
103
104 fn sample_edge(src: &str, dst: &str, kind: EdgeType) -> Edge {
105 Edge::new(src, dst, kind)
106 }
107
108 #[test]
109 fn memory_store_roundtrip() {
110 let mut s = MemoryStore::new();
111 s.init().expect("init");
112 let nodes = vec![
113 sample_node("d1", NodeType::Domain),
114 sample_node("h1", NodeType::Ip),
115 ];
116 let edges = vec![sample_edge("d1", "h1", EdgeType::ResolvesTo)];
117 s.write_nodes(&nodes).unwrap();
118 s.write_edges(&edges).unwrap();
119
120 let read_nodes = s.read_nodes().unwrap();
121 let read_edges = s.read_edges().unwrap();
122 assert_eq!(read_nodes.len(), 2);
123 assert_eq!(read_edges.len(), 1);
124 }
125
126 #[test]
127 fn memory_find_by_type_and_neighbors() {
128 let mut s = MemoryStore::new();
129 s.init().unwrap();
130 s.write_nodes(&[
131 sample_node("d1", NodeType::Domain),
132 sample_node("d2", NodeType::Domain),
133 sample_node("h1", NodeType::Ip),
134 ])
135 .unwrap();
136 s.write_edges(&[
137 sample_edge("d1", "h1", EdgeType::ResolvesTo),
138 sample_edge("d2", "h1", EdgeType::ResolvesTo),
139 ])
140 .unwrap();
141
142 let domains = s.find_nodes_by_type(NodeType::Domain).unwrap();
143 assert_eq!(domains.len(), 2);
144 let hosts = s.find_nodes_by_type(NodeType::Ip).unwrap();
145 assert_eq!(hosts.len(), 1);
146
147 let from_d1 = s.neighbors("d1", None).unwrap();
148 assert_eq!(from_d1.len(), 1);
149 assert_eq!(from_d1[0].target_id, "h1");
150 }
151
152 #[test]
153 fn memory_clear_resets_state() {
154 let mut s = MemoryStore::new();
155 s.init().unwrap();
156 s.write_nodes(&[sample_node("d1", NodeType::Domain)])
157 .unwrap();
158 assert_eq!(s.read_nodes().unwrap().len(), 1);
159 s.clear().unwrap();
160 assert!(s.read_nodes().unwrap().is_empty());
161 }
162}