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 },
21 TagScan {
22 tags: Vec<String>,
23 filters: Vec<Filter>,
24 limit: Option<usize>,
25 #[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
59pub 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 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 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 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 return Ok(QueryPlan::VectorSearch {
104 query: Vec::new(), k: limit,
106 filters: recall.filters.clone(),
107 condition: None,
108 });
109 }
110 return Err(MenteError::Query(
112 "~> operator requires a text value on the right-hand side".into(),
113 ));
114 }
115
116 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 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 let _ = s; }
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 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}