Skip to main content

sz_orm_graph/
engine.rs

1//! InMemoryGraphEngine — 内存图引擎
2//!
3//! 真实执行 Cypher 子集查询,替代 stub `Ok(vec![])` 实现。
4
5use crate::cypher_parser::{CypherSubsetParser, ParsedQuery, ReturnItem};
6use crate::error::GraphError;
7use crate::query::{GraphNode, GraphRelationship, GraphResult};
8use std::collections::HashMap;
9use std::sync::atomic::{AtomicU64, Ordering};
10
11#[derive(Debug)]
12pub struct InMemoryGraphEngine {
13    nodes: HashMap<String, GraphNode>,
14    relationships: HashMap<String, GraphRelationship>,
15    query_count: AtomicU64,
16    next_id: AtomicU64,
17}
18
19impl InMemoryGraphEngine {
20    pub fn new() -> Self {
21        Self {
22            nodes: HashMap::new(),
23            relationships: HashMap::new(),
24            query_count: AtomicU64::new(0),
25            next_id: AtomicU64::new(1),
26        }
27    }
28
29    fn generate_id(&self) -> String {
30        let id = self.next_id.fetch_add(1, Ordering::Relaxed);
31        format!("auto_{}", id)
32    }
33
34    pub fn add_node(&mut self, node: GraphNode) -> Result<(), GraphError> {
35        if self.nodes.contains_key(&node.id) {
36            return Err(GraphError::QueryError(format!(
37                "duplicate node id: {}",
38                node.id
39            )));
40        }
41        self.nodes.insert(node.id.clone(), node);
42        Ok(())
43    }
44
45    pub fn add_relationship(&mut self, rel: GraphRelationship) -> Result<(), GraphError> {
46        if !self.nodes.contains_key(&rel.start_node_id) {
47            return Err(GraphError::QueryError(format!(
48                "start node not found: {}",
49                rel.start_node_id
50            )));
51        }
52        if !self.nodes.contains_key(&rel.end_node_id) {
53            return Err(GraphError::QueryError(format!(
54                "end node not found: {}",
55                rel.end_node_id
56            )));
57        }
58        if self.relationships.contains_key(&rel.id) {
59            return Err(GraphError::QueryError(format!(
60                "duplicate relationship id: {}",
61                rel.id
62            )));
63        }
64        self.relationships.insert(rel.id.clone(), rel);
65        Ok(())
66    }
67
68    pub fn execute(
69        &self,
70        query: &crate::query::CypherQuery,
71    ) -> Result<Vec<GraphResult>, GraphError> {
72        self.query_count.fetch_add(1, Ordering::Relaxed);
73        let parsed = CypherSubsetParser::parse(&query.cypher, &query.parameters)?;
74
75        match parsed {
76            ParsedQuery::MatchNode {
77                alias: _,
78                label,
79                where_clause,
80                return_items,
81            } => {
82                if return_items
83                    .iter()
84                    .any(|ri| matches!(ri, ReturnItem::Count(_)))
85                {
86                    return self.match_count(label, where_clause, &query.parameters);
87                }
88                self.match_node(label, where_clause, &query.parameters)
89            }
90            ParsedQuery::MatchRelationship {
91                from,
92                rel,
93                to,
94                return_items: _,
95            } => self.match_relationship(from, rel, to),
96            ParsedQuery::CreateNode { .. }
97            | ParsedQuery::MergeNode { .. }
98            | ParsedQuery::Delete { .. }
99            | ParsedQuery::Set { .. } => Err(GraphError::QueryError(
100                "write operations require execute_mut(), use execute_mut() for CREATE/MERGE/DELETE/SET".into(),
101            )),
102        }
103    }
104
105    fn match_node(
106        &self,
107        label: Option<String>,
108        where_clause: Option<crate::cypher_parser::WhereClause>,
109        params: &HashMap<String, serde_json::Value>,
110    ) -> Result<Vec<GraphResult>, GraphError> {
111        let mut results = Vec::new();
112        for node in self.nodes.values() {
113            if let Some(ref label) = label {
114                if !node.labels.iter().any(|l| l == label) {
115                    continue;
116                }
117            }
118            if let Some(ref wc) = where_clause {
119                let param_value = params.get(&wc.param_name).ok_or_else(|| {
120                    GraphError::QueryError(format!("parameter not found: ${}", wc.param_name))
121                })?;
122                let node_prop = &node.properties;
123                if let Some(prop_val) = node_prop.get(&wc.prop) {
124                    if prop_val != param_value {
125                        continue;
126                    }
127                } else {
128                    continue;
129                }
130            }
131            results.push(GraphResult::Node { node: node.clone() });
132        }
133        Ok(results)
134    }
135
136    fn match_count(
137        &self,
138        label: Option<String>,
139        where_clause: Option<crate::cypher_parser::WhereClause>,
140        params: &HashMap<String, serde_json::Value>,
141    ) -> Result<Vec<GraphResult>, GraphError> {
142        let nodes = self.match_node(label, where_clause, params)?;
143        let count = nodes.len();
144        Ok(vec![GraphResult::Scalar {
145            value: serde_json::json!(count),
146        }])
147    }
148
149    fn match_relationship(
150        &self,
151        from: crate::cypher_parser::NodePattern,
152        rel: crate::cypher_parser::RelPattern,
153        to: crate::cypher_parser::NodePattern,
154    ) -> Result<Vec<GraphResult>, GraphError> {
155        let mut results = Vec::new();
156        for relationship in self.relationships.values() {
157            if let Some(ref rel_type) = rel.rel_type {
158                if &relationship.rel_type != rel_type {
159                    continue;
160                }
161            }
162            let start_node = match self.nodes.get(&relationship.start_node_id) {
163                Some(n) => n,
164                None => continue,
165            };
166            let end_node = match self.nodes.get(&relationship.end_node_id) {
167                Some(n) => n,
168                None => continue,
169            };
170            if let Some(ref label) = from.label {
171                if !start_node.labels.iter().any(|l| l == label) {
172                    continue;
173                }
174            }
175            if let Some(ref label) = to.label {
176                if !end_node.labels.iter().any(|l| l == label) {
177                    continue;
178                }
179            }
180            results.push(GraphResult::Node {
181                node: start_node.clone(),
182            });
183            results.push(GraphResult::Relationship {
184                relationship: relationship.clone(),
185            });
186            results.push(GraphResult::Node {
187                node: end_node.clone(),
188            });
189        }
190        Ok(results)
191    }
192
193    pub fn query_count(&self) -> u64 {
194        self.query_count.load(Ordering::Relaxed)
195    }
196
197    pub fn node_count(&self) -> usize {
198        self.nodes.len()
199    }
200
201    pub fn relationship_count(&self) -> usize {
202        self.relationships.len()
203    }
204
205    pub fn execute_mut(
206        &mut self,
207        query: &crate::query::CypherQuery,
208    ) -> Result<Vec<GraphResult>, GraphError> {
209        self.query_count.fetch_add(1, Ordering::Relaxed);
210        let parsed = CypherSubsetParser::parse(&query.cypher, &query.parameters)?;
211
212        match parsed {
213            ParsedQuery::CreateNode {
214                alias: _,
215                label,
216                properties,
217            } => self.execute_create(label, properties, &query.parameters),
218            ParsedQuery::MergeNode {
219                alias: _,
220                label,
221                properties,
222            } => self.execute_merge(label, properties, &query.parameters),
223            ParsedQuery::Delete { alias } => self.execute_delete(alias),
224            ParsedQuery::Set {
225                alias,
226                prop,
227                param_name,
228            } => self.execute_set(alias, prop, param_name, &query.parameters),
229            ParsedQuery::MatchNode { .. } | ParsedQuery::MatchRelationship { .. } => {
230                self.execute(query)
231            }
232        }
233    }
234
235    fn execute_create(
236        &mut self,
237        label: String,
238        properties: Vec<(String, String)>,
239        params: &HashMap<String, serde_json::Value>,
240    ) -> Result<Vec<GraphResult>, GraphError> {
241        let id = self.generate_id();
242        let mut props = serde_json::Map::new();
243        for (prop_name, param_name) in &properties {
244            let value = params.get(param_name).ok_or_else(|| {
245                GraphError::QueryError(format!("parameter not found: ${}", param_name))
246            })?;
247            props.insert(prop_name.clone(), value.clone());
248        }
249        let node = GraphNode {
250            id,
251            labels: vec![label],
252            properties: serde_json::Value::Object(props),
253        };
254        self.add_node(node)?;
255        Ok(Vec::new())
256    }
257
258    fn execute_merge(
259        &mut self,
260        label: String,
261        properties: Vec<(String, String)>,
262        params: &HashMap<String, serde_json::Value>,
263    ) -> Result<Vec<GraphResult>, GraphError> {
264        for node in self.nodes.values() {
265            if !node.labels.iter().any(|l| l == &label) {
266                continue;
267            }
268            let mut all_match = true;
269            for (prop_name, param_name) in &properties {
270                let param_value = params.get(param_name).ok_or_else(|| {
271                    GraphError::QueryError(format!("parameter not found: ${}", param_name))
272                })?;
273                if node.properties.get(prop_name) != Some(param_value) {
274                    all_match = false;
275                    break;
276                }
277            }
278            if all_match {
279                return Ok(Vec::new());
280            }
281        }
282        self.execute_create(label, properties, params)
283    }
284
285    fn execute_delete(&mut self, alias: String) -> Result<Vec<GraphResult>, GraphError> {
286        let node_id = if self.nodes.contains_key(&alias) {
287            alias
288        } else {
289            let mut found = None;
290            for node in self.nodes.values() {
291                if node.properties.get("alias").and_then(|v| v.as_str()) == Some(&alias) {
292                    found = Some(node.id.clone());
293                    break;
294                }
295            }
296            found.ok_or_else(|| GraphError::QueryError(format!("node not found: {}", alias)))?
297        };
298
299        let mut rels_to_remove: Vec<String> = Vec::new();
300        for (rel_id, rel) in &self.relationships {
301            if rel.start_node_id == node_id || rel.end_node_id == node_id {
302                rels_to_remove.push(rel_id.clone());
303            }
304        }
305        for rel_id in rels_to_remove {
306            self.relationships.remove(&rel_id);
307        }
308        self.nodes.remove(&node_id);
309        Ok(Vec::new())
310    }
311
312    fn execute_set(
313        &mut self,
314        alias: String,
315        prop: String,
316        param_name: String,
317        params: &HashMap<String, serde_json::Value>,
318    ) -> Result<Vec<GraphResult>, GraphError> {
319        let value = params.get(&param_name).ok_or_else(|| {
320            GraphError::QueryError(format!("parameter not found: ${}", param_name))
321        })?;
322
323        let node_id = if self.nodes.contains_key(&alias) {
324            alias
325        } else {
326            let mut found = None;
327            for node in self.nodes.values() {
328                if node.properties.get("alias").and_then(|v| v.as_str()) == Some(&alias) {
329                    found = Some(node.id.clone());
330                    break;
331                }
332            }
333            found.ok_or_else(|| GraphError::QueryError(format!("node not found: {}", alias)))?
334        };
335
336        let node = self.nodes.get_mut(&node_id).unwrap();
337        if let Some(obj) = node.properties.as_object_mut() {
338            obj.insert(prop, value.clone());
339        } else {
340            let mut obj = serde_json::Map::new();
341            obj.insert(prop, value.clone());
342            node.properties = serde_json::Value::Object(obj);
343        }
344        Ok(Vec::new())
345    }
346}
347
348impl Default for InMemoryGraphEngine {
349    fn default() -> Self {
350        Self::new()
351    }
352}
353
354#[cfg(test)]
355mod tests {
356    use super::*;
357    use crate::query::{CypherQuery, GraphNode, GraphRelationship};
358    use std::collections::HashMap;
359
360    fn make_node(id: &str, label: &str, props: serde_json::Value) -> GraphNode {
361        GraphNode {
362            id: id.into(),
363            labels: vec![label.into()],
364            properties: props,
365        }
366    }
367
368    #[test]
369    fn test_add_node_duplicate_rejected() {
370        let mut engine = InMemoryGraphEngine::new();
371        let node = make_node("1", "Person", serde_json::json!({}));
372        engine.add_node(node).unwrap();
373        let dup = make_node("1", "Person", serde_json::json!({}));
374        assert!(engine.add_node(dup).is_err());
375    }
376
377    #[test]
378    fn test_add_relationship_endpoint_not_found() {
379        let mut engine = InMemoryGraphEngine::new();
380        let rel = GraphRelationship {
381            id: "r1".into(),
382            rel_type: "KNOWS".into(),
383            start_node_id: "1".into(),
384            end_node_id: "2".into(),
385            properties: serde_json::json!({}),
386        };
387        assert!(engine.add_relationship(rel).is_err());
388    }
389
390    #[test]
391    fn test_add_relationship_success() {
392        let mut engine = InMemoryGraphEngine::new();
393        engine
394            .add_node(make_node("1", "Person", serde_json::json!({})))
395            .unwrap();
396        engine
397            .add_node(make_node("2", "Person", serde_json::json!({})))
398            .unwrap();
399        let rel = GraphRelationship {
400            id: "r1".into(),
401            rel_type: "KNOWS".into(),
402            start_node_id: "1".into(),
403            end_node_id: "2".into(),
404            properties: serde_json::json!({}),
405        };
406        assert!(engine.add_relationship(rel).is_ok());
407        assert_eq!(engine.relationship_count(), 1);
408    }
409
410    #[test]
411    fn test_execute_increments_query_count() {
412        let engine = InMemoryGraphEngine::new();
413        let q = CypherQuery::new("MATCH (n:Person) RETURN n");
414        assert_eq!(engine.query_count(), 0);
415        let _ = engine.execute(&q).unwrap();
416        assert!(engine.query_count() >= 1);
417    }
418
419    #[test]
420    fn test_execute_empty_graph_returns_empty_but_real() {
421        let engine = InMemoryGraphEngine::new();
422        let q = CypherQuery::new("MATCH (n:Person) RETURN n");
423        let result = engine.execute(&q).unwrap();
424        assert!(result.is_empty());
425        assert!(engine.query_count() >= 1);
426    }
427
428    #[test]
429    fn test_execute_returns_real_node() {
430        let mut engine = InMemoryGraphEngine::new();
431        engine
432            .add_node(make_node(
433                "1",
434                "Person",
435                serde_json::json!({"name": "Alice"}),
436            ))
437            .unwrap();
438        let q = CypherQuery::new("MATCH (n:Person) RETURN n");
439        let result = engine.execute(&q).unwrap();
440        assert_eq!(result.len(), 1);
441        assert!(result[0].as_node().is_some());
442    }
443
444    #[test]
445    fn test_execute_with_where_param() {
446        let mut engine = InMemoryGraphEngine::new();
447        engine
448            .add_node(make_node(
449                "1",
450                "Person",
451                serde_json::json!({"name": "Alice"}),
452            ))
453            .unwrap();
454        engine
455            .add_node(make_node("2", "Person", serde_json::json!({"name": "Bob"})))
456            .unwrap();
457        let mut params = HashMap::new();
458        params.insert("name".into(), serde_json::json!("Alice"));
459        let q = CypherQuery::with_params("MATCH (n:Person) WHERE n.name = $name RETURN n", params);
460        let result = engine.execute(&q).unwrap();
461        assert_eq!(result.len(), 1);
462    }
463
464    #[test]
465    fn test_execute_count_aggregation() {
466        let mut engine = InMemoryGraphEngine::new();
467        engine
468            .add_node(make_node("1", "Person", serde_json::json!({})))
469            .unwrap();
470        engine
471            .add_node(make_node("2", "Person", serde_json::json!({})))
472            .unwrap();
473        let q = CypherQuery::new("MATCH (n:Person) RETURN count(n)");
474        let result = engine.execute(&q).unwrap();
475        assert_eq!(result.len(), 1);
476        let scalar = result[0].as_scalar().unwrap();
477        assert_eq!(scalar, &serde_json::json!(2));
478    }
479
480    #[test]
481    fn test_execute_relationship_query() {
482        let mut engine = InMemoryGraphEngine::new();
483        engine
484            .add_node(make_node(
485                "1",
486                "Person",
487                serde_json::json!({"name": "Alice"}),
488            ))
489            .unwrap();
490        engine
491            .add_node(make_node("2", "Person", serde_json::json!({"name": "Bob"})))
492            .unwrap();
493        let rel = GraphRelationship {
494            id: "r1".into(),
495            rel_type: "KNOWS".into(),
496            start_node_id: "1".into(),
497            end_node_id: "2".into(),
498            properties: serde_json::json!({}),
499        };
500        engine.add_relationship(rel).unwrap();
501
502        let q = CypherQuery::new("MATCH (a:Person)-[r:KNOWS]->(b:Person) RETURN a, r, b");
503        let result = engine.execute(&q).unwrap();
504        assert_eq!(result.len(), 3);
505        assert!(result[0].as_node().is_some());
506        assert!(result[1].as_relationship().is_some());
507        assert!(result[2].as_node().is_some());
508    }
509}