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