1use 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#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
11pub enum QueryPlan {
12 VectorSearch {
13 query: Vec<f32>,
14 k: usize,
15 filters: Vec<Filter>,
16 #[serde(default)]
19 condition: Option<Condition>,
20 #[serde(default)]
22 order_by: Option<OrderBy>,
23 },
24 TagScan {
25 tags: Vec<String>,
26 filters: Vec<Filter>,
27 limit: Option<usize>,
28 #[serde(default)]
31 condition: Option<Condition>,
32 #[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
67pub 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 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 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 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 return Ok(QueryPlan::VectorSearch {
114 query: Vec::new(), k: limit,
116 filters: recall.filters.clone(),
117 condition: None,
118 order_by: recall.order_by.clone(),
119 });
120 }
121 return Err(MenteError::Query(
123 "~> operator requires a text value on the right-hand side".into(),
124 ));
125 }
126
127 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 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 let _ = s; }
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 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}