Skip to main content

gossan_graph/store/
memory.rs

1//! In-memory graph backend.
2//!
3//! Holds nodes + edges in `Vec`s. Useful for short-lived scans where
4//! the persistence cost of sqlite/graphml/json isn't justified, and as
5//! the simplest implementation against which the
6//! [`GraphBackend`] trait shape can be verified.
7
8use crate::schema::{EdgeType, NodeType};
9use crate::{Edge, Node};
10
11use super::GraphBackend;
12
13/// Errors returned by the in-memory backend.
14#[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/// In-memory store of nodes and edges.
26#[derive(Debug, Default, Clone)]
27pub struct MemoryStore {
28    nodes: Vec<Node>,
29    edges: Vec<Edge>,
30}
31
32impl MemoryStore {
33    /// Construct an empty store.
34    #[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}