Skip to main content

mentedb_query/
planner.rs

1//! Query planner: converts a parsed `Statement` into a `QueryPlan`.
2
3use crate::ast::*;
4use mentedb_core::edge::EdgeType;
5use mentedb_core::error::{MenteError, MenteResult};
6use mentedb_core::types::{MemoryId, Timestamp};
7use serde::{Deserialize, Serialize};
8
9/// A physical query plan describing how to execute a query.
10#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
11pub enum QueryPlan {
12    VectorSearch {
13        query: Vec<f32>,
14        k: usize,
15        filters: Vec<Filter>,
16        /// Full boolean WHERE tree, when the clause used OR/NOT/grouping. When set,
17        /// the executor post-filters candidates by this instead of `filters`.
18        #[serde(default)]
19        condition: Option<Condition>,
20    },
21    TagScan {
22        tags: Vec<String>,
23        filters: Vec<Filter>,
24        limit: Option<usize>,
25        /// Full boolean WHERE tree, when the clause used OR/NOT/grouping. When set,
26        /// the executor post-filters candidates by this instead of `filters`.
27        #[serde(default)]
28        condition: Option<Condition>,
29    },
30    TemporalScan {
31        start: Timestamp,
32        end: Timestamp,
33        filters: Vec<Filter>,
34    },
35    GraphTraversal {
36        start: MemoryId,
37        depth: usize,
38        edge_types: Vec<EdgeType>,
39    },
40    PointLookup {
41        id: MemoryId,
42    },
43    EdgeInsert {
44        source: MemoryId,
45        target: MemoryId,
46        edge_type: EdgeType,
47        weight: f32,
48    },
49    Delete {
50        id: MemoryId,
51    },
52    Consolidate {
53        filters: Vec<Filter>,
54    },
55}
56
57const DEFAULT_LIMIT: usize = 20;
58
59/// Produce an execution plan from a parsed statement.
60pub fn plan(statement: &Statement) -> MenteResult<QueryPlan> {
61    match statement {
62        Statement::Recall(recall) => plan_recall(recall),
63        Statement::Relate(relate) => plan_relate(relate),
64        Statement::Forget(forget) => Ok(QueryPlan::Delete { id: forget.target }),
65        Statement::Consolidate(cons) => Ok(QueryPlan::Consolidate {
66            filters: cons.filters.clone(),
67        }),
68        Statement::Traverse(trav) => plan_traverse(trav),
69    }
70}
71
72fn plan_recall(recall: &RecallStatement) -> MenteResult<QueryPlan> {
73    let limit = recall.limit.unwrap_or(DEFAULT_LIMIT);
74
75    // If there's a NEAR clause or a SimilarTo filter, use vector search
76    if let Some(ref vec) = recall.near {
77        return Ok(QueryPlan::VectorSearch {
78            query: vec.clone(),
79            k: limit,
80            filters: recall.filters.clone(),
81            condition: recall.condition.clone(),
82        });
83    }
84
85    // OR/NOT/grouped WHERE: none of the leaf-based index optimizations below are
86    // safe (they assume every filter must hold), so full-scan and let the executor
87    // evaluate the boolean tree. Correct, if less selective, for these rarer queries.
88    if recall.condition.is_some() {
89        return Ok(QueryPlan::TagScan {
90            tags: Vec::new(),
91            filters: Vec::new(),
92            limit: Some(limit),
93            condition: recall.condition.clone(),
94        });
95    }
96
97    // Check for SimilarTo operator in filters — implies vector search via text embedding
98    if let Some(sim_filter) = recall.filters.iter().find(|f| f.op == Operator::SimilarTo) {
99        if let Value::Text(ref _text) = sim_filter.value {
100            // The text is embedded at execution time (no embedder here), so keep
101            // ALL filters, including the SimilarTo one: the executor reads its
102            // text to build the query vector and applies the rest (type, tag).
103            return Ok(QueryPlan::VectorSearch {
104                query: Vec::new(), // placeholder — executor embeds the SimilarTo text
105                k: limit,
106                filters: recall.filters.clone(),
107                condition: None,
108            });
109        }
110        // SimilarTo with non-text value doesn't make sense
111        return Err(MenteError::Query(
112            "~> operator requires a text value on the right-hand side".into(),
113        ));
114    }
115
116    // If only tag filters, use TagScan
117    let tag_filters: Vec<&Filter> = recall
118        .filters
119        .iter()
120        .filter(|f| f.field == Field::Tag)
121        .collect();
122    if !tag_filters.is_empty() && recall.filters.iter().all(|f| f.field == Field::Tag) {
123        let tags: Vec<String> = tag_filters
124            .iter()
125            .filter_map(|f| match &f.value {
126                Value::Text(t) => Some(t.clone()),
127                _ => None,
128            })
129            .collect();
130        return Ok(QueryPlan::TagScan {
131            tags,
132            filters: Vec::new(),
133            limit: Some(limit),
134            condition: None,
135        });
136    }
137
138    // If time-range filters exist (created or accessed with range ops), use TemporalScan
139    let time_filters: Vec<&Filter> = recall
140        .filters
141        .iter()
142        .filter(|f| {
143            (f.field == Field::Created || f.field == Field::Accessed)
144                && matches!(
145                    f.op,
146                    Operator::Gt | Operator::Lt | Operator::Gte | Operator::Lte
147                )
148        })
149        .collect();
150
151    if !time_filters.is_empty() {
152        let remaining: Vec<Filter> = recall
153            .filters
154            .iter()
155            .filter(|f| {
156                !((f.field == Field::Created || f.field == Field::Accessed)
157                    && matches!(
158                        f.op,
159                        Operator::Gt | Operator::Lt | Operator::Gte | Operator::Lte
160                    ))
161            })
162            .cloned()
163            .collect();
164
165        let mut start: Timestamp = 0;
166        let mut end: Timestamp = u64::MAX;
167        for f in &time_filters {
168            if let Value::Text(ref s) = f.value {
169                // Simple heuristic: treat text timestamps as orderable strings for now.
170                // A real implementation would parse dates. We use 0/MAX as placeholders.
171                let _ = s; // acknowledged
172            }
173            if let Value::Integer(ts) = f.value {
174                let ts = ts as u64;
175                match f.op {
176                    Operator::Gt | Operator::Gte => start = ts,
177                    Operator::Lt | Operator::Lte => end = ts,
178                    _ => {}
179                }
180            }
181        }
182
183        return Ok(QueryPlan::TemporalScan {
184            start,
185            end,
186            filters: remaining,
187        });
188    }
189
190    // Fallback: tag scan with no tags (full scan with filters)
191    Ok(QueryPlan::TagScan {
192        tags: Vec::new(),
193        filters: recall.filters.clone(),
194        limit: Some(limit),
195        condition: None,
196    })
197}
198
199fn plan_relate(relate: &RelateStatement) -> MenteResult<QueryPlan> {
200    Ok(QueryPlan::EdgeInsert {
201        source: relate.source,
202        target: relate.target,
203        edge_type: relate.edge_type,
204        weight: relate.weight.unwrap_or(1.0),
205    })
206}
207
208fn plan_traverse(trav: &TraverseStatement) -> MenteResult<QueryPlan> {
209    Ok(QueryPlan::GraphTraversal {
210        start: trav.start,
211        depth: trav.depth,
212        edge_types: trav.edge_filter.clone().unwrap_or_default(),
213    })
214}
215
216#[cfg(test)]
217mod tests {
218    use super::*;
219    use crate::lexer::tokenize;
220    use crate::parser::Parser;
221
222    fn plan_mql(input: &str) -> QueryPlan {
223        let tokens = tokenize(input).unwrap();
224        let stmt = Parser::parse(&tokens).unwrap();
225        plan(&stmt).unwrap()
226    }
227
228    #[test]
229    fn test_near_produces_vector_search() {
230        let qp = plan_mql("RECALL memories NEAR [0.1, 0.2, 0.3] LIMIT 5");
231        match qp {
232            QueryPlan::VectorSearch { query, k, .. } => {
233                assert_eq!(query, vec![0.1, 0.2, 0.3]);
234                assert_eq!(k, 5);
235            }
236            _ => panic!("expected VectorSearch, got {:?}", qp),
237        }
238    }
239
240    #[test]
241    fn test_similar_to_produces_vector_search() {
242        let qp = plan_mql(r#"RECALL memories WHERE content ~> "database migration" LIMIT 10"#);
243        match qp {
244            QueryPlan::VectorSearch { k, .. } => {
245                assert_eq!(k, 10);
246            }
247            _ => panic!("expected VectorSearch, got {:?}", qp),
248        }
249    }
250
251    #[test]
252    fn test_tag_filter_produces_tag_scan() {
253        let qp = plan_mql(r#"RECALL memories WHERE tag = "backend" LIMIT 5"#);
254        match qp {
255            QueryPlan::TagScan { tags, limit, .. } => {
256                assert_eq!(tags, vec!["backend".to_string()]);
257                assert_eq!(limit, Some(5));
258            }
259            _ => panic!("expected TagScan, got {:?}", qp),
260        }
261    }
262
263    #[test]
264    fn test_forget_produces_delete() {
265        let qp = plan_mql("FORGET 550e8400-e29b-41d4-a716-446655440000");
266        match qp {
267            QueryPlan::Delete { id } => {
268                assert_eq!(
269                    id,
270                    "550e8400-e29b-41d4-a716-446655440000"
271                        .parse::<MemoryId>()
272                        .unwrap()
273                );
274            }
275            _ => panic!("expected Delete, got {:?}", qp),
276        }
277    }
278
279    #[test]
280    fn test_traverse_produces_graph_traversal() {
281        let qp = plan_mql(
282            "TRAVERSE 550e8400-e29b-41d4-a716-446655440000 DEPTH 3 WHERE edge_type = caused",
283        );
284        match qp {
285            QueryPlan::GraphTraversal {
286                depth, edge_types, ..
287            } => {
288                assert_eq!(depth, 3);
289                assert_eq!(edge_types, vec![EdgeType::Caused]);
290            }
291            _ => panic!("expected GraphTraversal, got {:?}", qp),
292        }
293    }
294
295    #[test]
296    fn test_relate_produces_edge_insert() {
297        let qp = plan_mql(
298            "RELATE 550e8400-e29b-41d4-a716-446655440000 -> 660e8400-e29b-41d4-a716-446655440000 AS caused WITH weight = 0.8",
299        );
300        match qp {
301            QueryPlan::EdgeInsert {
302                edge_type, weight, ..
303            } => {
304                assert_eq!(edge_type, EdgeType::Caused);
305                assert!((weight - 0.8).abs() < f32::EPSILON);
306            }
307            _ => panic!("expected EdgeInsert, got {:?}", qp),
308        }
309    }
310}