Skip to main content

mentedb_query/
parser.rs

1//! Hand-written recursive descent parser for MQL.
2
3use mentedb_core::edge::EdgeType;
4use mentedb_core::error::{MenteError, MenteResult};
5use mentedb_core::memory::MemoryType;
6use uuid::Uuid;
7
8use crate::ast::*;
9use crate::lexer::{Token, TokenKind};
10use mentedb_core::types::MemoryId;
11
12pub struct Parser<'a> {
13    tokens: &'a [Token],
14    pos: usize,
15}
16
17impl<'a> Parser<'a> {
18    pub fn new(tokens: &'a [Token]) -> Self {
19        Self { tokens, pos: 0 }
20    }
21
22    pub fn parse(tokens: &[Token]) -> MenteResult<Statement> {
23        let mut parser = Parser::new(tokens);
24        parser.parse_statement()
25    }
26
27    fn peek(&self) -> &Token {
28        &self.tokens[self.pos.min(self.tokens.len() - 1)]
29    }
30
31    fn advance(&mut self) -> &Token {
32        let tok = &self.tokens[self.pos.min(self.tokens.len() - 1)];
33        if self.pos < self.tokens.len() {
34            self.pos += 1;
35        }
36        tok
37    }
38
39    fn expect(&mut self, kind: TokenKind) -> MenteResult<&Token> {
40        let tok = self.peek();
41        if tok.kind != kind {
42            return Err(MenteError::Query(format!(
43                "expected {:?}, found {:?} ('{}') at position {}",
44                kind, tok.kind, tok.lexeme, tok.position
45            )));
46        }
47        Ok(self.advance())
48    }
49
50    fn at(&self, kind: TokenKind) -> bool {
51        self.peek().kind == kind
52    }
53
54    fn parse_statement(&mut self) -> MenteResult<Statement> {
55        match self.peek().kind {
56            TokenKind::Recall => self.parse_recall(),
57            TokenKind::Relate => self.parse_relate(),
58            TokenKind::Forget => self.parse_forget(),
59            TokenKind::Consolidate => self.parse_consolidate(),
60            TokenKind::Traverse => self.parse_traverse(),
61            _ => Err(MenteError::Query(format!(
62                "expected statement keyword, found {:?} at position {}",
63                self.peek().kind,
64                self.peek().position
65            ))),
66        }
67    }
68
69    fn parse_recall(&mut self) -> MenteResult<Statement> {
70        self.advance(); // RECALL
71
72        // Optional "memories" keyword
73        if self.at(TokenKind::Memories) {
74            self.advance();
75        }
76
77        let mut near = None;
78        let mut filters = Vec::new();
79        let mut limit = None;
80        let mut order_by = None;
81
82        // NEAR [vector]
83        if self.at(TokenKind::Near) {
84            self.advance();
85            near = Some(self.parse_vector()?);
86        }
87
88        // WHERE clause
89        if self.at(TokenKind::Where) {
90            self.advance();
91            filters = self.parse_filters()?;
92        }
93
94        // ORDER BY field
95        if self.at(TokenKind::OrderBy) {
96            self.advance();
97            // consume optional "BY"
98            if self.at(TokenKind::By) {
99                self.advance();
100            }
101            let field = self.parse_field()?;
102            let descending = false; // default ascending
103            order_by = Some(OrderBy { field, descending });
104        }
105
106        // LIMIT n
107        if self.at(TokenKind::Limit) {
108            self.advance();
109            let tok = self.advance();
110            let n: usize = tok
111                .lexeme
112                .parse()
113                .map_err(|_| MenteError::Query(format!("invalid limit value: {}", tok.lexeme)))?;
114            limit = Some(n);
115        }
116
117        // AS OF <timestamp> — point-in-time temporal filter. Added as a ValidAt
118        // filter so it flows through the same pipeline as WHERE clauses.
119        if self.at(TokenKind::As) {
120            self.advance();
121            self.expect(TokenKind::Of)?;
122            let tok = self.advance();
123            let t: i64 = tok.lexeme.parse().map_err(|_| {
124                MenteError::Query(format!("invalid AS OF timestamp: {}", tok.lexeme))
125            })?;
126            filters.push(Filter {
127                field: Field::ValidAt,
128                op: Operator::Eq,
129                value: Value::Integer(t),
130            });
131        }
132
133        Ok(Statement::Recall(RecallStatement {
134            filters,
135            near,
136            limit,
137            order_by,
138        }))
139    }
140
141    fn parse_relate(&mut self) -> MenteResult<Statement> {
142        self.advance(); // RELATE
143
144        let source = self.parse_uuid()?;
145        self.expect(TokenKind::Arrow)?;
146        let target = self.parse_uuid()?;
147        self.expect(TokenKind::As)?;
148        let edge_type = self.parse_edge_type()?;
149
150        let mut weight = None;
151        if self.at(TokenKind::With) {
152            self.advance();
153            // expect "weight = <float>"
154            self.expect(TokenKind::Identifier)?; // "weight"
155            self.expect(TokenKind::Eq)?;
156            let tok = self.advance();
157            let w: f32 = tok
158                .lexeme
159                .parse()
160                .map_err(|_| MenteError::Query(format!("invalid weight value: {}", tok.lexeme)))?;
161            weight = Some(w);
162        }
163
164        Ok(Statement::Relate(RelateStatement {
165            source,
166            target,
167            edge_type,
168            weight,
169        }))
170    }
171
172    fn parse_forget(&mut self) -> MenteResult<Statement> {
173        self.advance(); // FORGET
174        let target = self.parse_uuid()?;
175        Ok(Statement::Forget(ForgetStatement { target }))
176    }
177
178    fn parse_consolidate(&mut self) -> MenteResult<Statement> {
179        self.advance(); // CONSOLIDATE
180        let mut filters = Vec::new();
181        if self.at(TokenKind::Where) {
182            self.advance();
183            filters = self.parse_filters()?;
184        }
185        Ok(Statement::Consolidate(ConsolidateStatement { filters }))
186    }
187
188    fn parse_traverse(&mut self) -> MenteResult<Statement> {
189        self.advance(); // TRAVERSE
190        let start = self.parse_uuid()?;
191
192        self.expect(TokenKind::Depth)?;
193        let tok = self.advance();
194        let depth: usize = tok
195            .lexeme
196            .parse()
197            .map_err(|_| MenteError::Query(format!("invalid depth value: {}", tok.lexeme)))?;
198
199        let mut edge_filter = None;
200        if self.at(TokenKind::Where) {
201            self.advance();
202            // edge_type = <type>
203            self.expect(TokenKind::EdgeType)?;
204            self.expect(TokenKind::Eq)?;
205            let et = self.parse_edge_type()?;
206            edge_filter = Some(vec![et]);
207        }
208
209        Ok(Statement::Traverse(TraverseStatement {
210            start,
211            depth,
212            edge_filter,
213        }))
214    }
215
216    fn parse_filters(&mut self) -> MenteResult<Vec<Filter>> {
217        let mut filters = vec![self.parse_filter()?];
218        while self.at(TokenKind::And) {
219            self.advance();
220            filters.push(self.parse_filter()?);
221        }
222        Ok(filters)
223    }
224
225    fn parse_filter(&mut self) -> MenteResult<Filter> {
226        let field = self.parse_field()?;
227        let op = self.parse_operator()?;
228        let value = if op == Operator::In {
229            self.parse_list_value(&field)?
230        } else {
231            self.parse_value(&field)?
232        };
233        Ok(Filter { field, op, value })
234    }
235
236    fn parse_field(&mut self) -> MenteResult<Field> {
237        let tok = self.advance();
238        match tok.kind {
239            TokenKind::Identifier if tok.lexeme.eq_ignore_ascii_case("content") => {
240                Ok(Field::Content)
241            }
242            TokenKind::Type => Ok(Field::Type),
243            TokenKind::Tag => Ok(Field::Tag),
244            TokenKind::Agent => Ok(Field::Agent),
245            TokenKind::Space => Ok(Field::Space),
246            TokenKind::Salience => Ok(Field::Salience),
247            TokenKind::Confidence => Ok(Field::Confidence),
248            TokenKind::Created => Ok(Field::Created),
249            TokenKind::Accessed => Ok(Field::Accessed),
250            _ => Err(MenteError::Query(format!(
251                "expected field name, found '{}' at position {}",
252                tok.lexeme, tok.position
253            ))),
254        }
255    }
256
257    fn parse_operator(&mut self) -> MenteResult<Operator> {
258        let tok = self.advance();
259        match tok.kind {
260            TokenKind::Eq => Ok(Operator::Eq),
261            TokenKind::Neq => Ok(Operator::Neq),
262            TokenKind::Gt => Ok(Operator::Gt),
263            TokenKind::Lt => Ok(Operator::Lt),
264            TokenKind::Gte => Ok(Operator::Gte),
265            TokenKind::Lte => Ok(Operator::Lte),
266            TokenKind::SimilarTo => Ok(Operator::SimilarTo),
267            TokenKind::In => Ok(Operator::In),
268            TokenKind::Contains => Ok(Operator::Contains),
269            _ => Err(MenteError::Query(format!(
270                "expected operator, found '{}' at position {}",
271                tok.lexeme, tok.position
272            ))),
273        }
274    }
275
276    /// Parse a bracketed, comma-separated list of scalar values for `IN`.
277    fn parse_list_value(&mut self, field: &Field) -> MenteResult<Value> {
278        self.expect(TokenKind::LBracket)?;
279        let mut items = Vec::new();
280        if !self.at(TokenKind::RBracket) {
281            loop {
282                items.push(self.parse_value(field)?);
283                if self.at(TokenKind::Comma) {
284                    self.advance();
285                } else {
286                    break;
287                }
288            }
289        }
290        self.expect(TokenKind::RBracket)?;
291        Ok(Value::List(items))
292    }
293
294    fn parse_value(&mut self, field: &Field) -> MenteResult<Value> {
295        // For Type field, parse as MemoryType
296        if *field == Field::Type {
297            return self.parse_memory_type_value();
298        }
299
300        let tok = self.advance();
301        match tok.kind {
302            TokenKind::StringLit => {
303                // Strip surrounding quotes
304                let inner = tok.lexeme[1..tok.lexeme.len() - 1].to_string();
305                // Check if this looks like a UUID inside quotes
306                if let Ok(uuid) = inner.parse::<MemoryId>() {
307                    return Ok(Value::Uuid(uuid.into()));
308                }
309                Ok(Value::Text(inner))
310            }
311            TokenKind::IntegerLit => {
312                let n: i64 = tok
313                    .lexeme
314                    .parse()
315                    .map_err(|_| MenteError::Query(format!("invalid integer: {}", tok.lexeme)))?;
316                Ok(Value::Integer(n))
317            }
318            TokenKind::FloatLit => {
319                let n: f64 = tok
320                    .lexeme
321                    .parse()
322                    .map_err(|_| MenteError::Query(format!("invalid float: {}", tok.lexeme)))?;
323                Ok(Value::Number(n))
324            }
325            TokenKind::UuidLit => {
326                let uuid: Uuid = tok
327                    .lexeme
328                    .parse()
329                    .map_err(|_| MenteError::Query(format!("invalid UUID: {}", tok.lexeme)))?;
330                Ok(Value::Uuid(uuid))
331            }
332            TokenKind::Identifier => {
333                let lower = tok.lexeme.to_lowercase();
334                match lower.as_str() {
335                    "true" => Ok(Value::Bool(true)),
336                    "false" => Ok(Value::Bool(false)),
337                    _ => Ok(Value::Text(tok.lexeme.clone())),
338                }
339            }
340            TokenKind::LBracket => {
341                // put back and parse as vector
342                self.pos -= 1;
343                let v = self.parse_vector()?;
344                Ok(Value::Vector(v))
345            }
346            _ => Err(MenteError::Query(format!(
347                "expected value, found '{}' at position {}",
348                tok.lexeme, tok.position
349            ))),
350        }
351    }
352
353    fn parse_memory_type_value(&mut self) -> MenteResult<Value> {
354        let tok = self.advance();
355        let name = match tok.kind {
356            TokenKind::Identifier | TokenKind::StringLit => {
357                if tok.kind == TokenKind::StringLit {
358                    tok.lexeme[1..tok.lexeme.len() - 1].to_string()
359                } else {
360                    tok.lexeme.clone()
361                }
362            }
363            _ => {
364                return Err(MenteError::Query(format!(
365                    "expected memory type, found '{}' at position {}",
366                    tok.lexeme, tok.position
367                )));
368            }
369        };
370
371        let mt = match name.to_lowercase().as_str() {
372            "episodic" => MemoryType::Episodic,
373            "semantic" => MemoryType::Semantic,
374            "procedural" => MemoryType::Procedural,
375            "antipattern" | "anti_pattern" => MemoryType::AntiPattern,
376            "reasoning" => MemoryType::Reasoning,
377            "correction" => MemoryType::Correction,
378            _ => {
379                return Err(MenteError::Query(format!("unknown memory type: {}", name)));
380            }
381        };
382        Ok(Value::MemoryType(mt))
383    }
384
385    fn parse_edge_type(&mut self) -> MenteResult<EdgeType> {
386        let tok = self.advance();
387        let name = match tok.kind {
388            TokenKind::Identifier | TokenKind::StringLit => {
389                if tok.kind == TokenKind::StringLit {
390                    tok.lexeme[1..tok.lexeme.len() - 1].to_string()
391                } else {
392                    tok.lexeme.clone()
393                }
394            }
395            _ => {
396                return Err(MenteError::Query(format!(
397                    "expected edge type, found '{}' at position {}",
398                    tok.lexeme, tok.position
399                )));
400            }
401        };
402
403        match name.to_lowercase().as_str() {
404            "caused" => Ok(EdgeType::Caused),
405            "before" => Ok(EdgeType::Before),
406            "related" => Ok(EdgeType::Related),
407            "contradicts" => Ok(EdgeType::Contradicts),
408            "supports" => Ok(EdgeType::Supports),
409            "supersedes" => Ok(EdgeType::Supersedes),
410            "derived" => Ok(EdgeType::Derived),
411            "partof" | "part_of" => Ok(EdgeType::PartOf),
412            _ => Err(MenteError::Query(format!("unknown edge type: {}", name))),
413        }
414    }
415
416    fn parse_uuid(&mut self) -> MenteResult<MemoryId> {
417        let tok = self.advance();
418        match tok.kind {
419            TokenKind::UuidLit => tok
420                .lexeme
421                .parse()
422                .map_err(|_| MenteError::Query(format!("invalid UUID: {}", tok.lexeme))),
423            TokenKind::StringLit => {
424                let inner = &tok.lexeme[1..tok.lexeme.len() - 1];
425                inner.parse().map_err(|_| {
426                    MenteError::Query(format!("invalid UUID in string: {}", tok.lexeme))
427                })
428            }
429            _ => Err(MenteError::Query(format!(
430                "expected UUID, found '{}' at position {}",
431                tok.lexeme, tok.position
432            ))),
433        }
434    }
435
436    fn parse_vector(&mut self) -> MenteResult<Vec<f32>> {
437        self.expect(TokenKind::LBracket)?;
438        let mut values = Vec::new();
439        if !self.at(TokenKind::RBracket) {
440            let tok = self.advance();
441            let v: f32 = tok.lexeme.parse().map_err(|_| {
442                MenteError::Query(format!("invalid float in vector: {}", tok.lexeme))
443            })?;
444            values.push(v);
445            while self.at(TokenKind::Comma) {
446                self.advance();
447                let tok = self.advance();
448                let v: f32 = tok.lexeme.parse().map_err(|_| {
449                    MenteError::Query(format!("invalid float in vector: {}", tok.lexeme))
450                })?;
451                values.push(v);
452            }
453        }
454        self.expect(TokenKind::RBracket)?;
455        Ok(values)
456    }
457}
458
459#[cfg(test)]
460mod tests {
461    use super::*;
462    use crate::lexer::tokenize;
463
464    #[test]
465    fn test_parse_recall_with_type_filter() {
466        let tokens = tokenize("RECALL memories WHERE type = episodic LIMIT 5").unwrap();
467        let stmt = Parser::parse(&tokens).unwrap();
468        match stmt {
469            Statement::Recall(r) => {
470                assert_eq!(r.filters.len(), 1);
471                assert_eq!(r.filters[0].field, Field::Type);
472                assert_eq!(r.filters[0].value, Value::MemoryType(MemoryType::Episodic));
473                assert_eq!(r.limit, Some(5));
474            }
475            _ => panic!("expected Recall"),
476        }
477    }
478
479    #[test]
480    fn test_parse_in_operator() {
481        let tokens =
482            tokenize("RECALL memories WHERE type IN [episodic, semantic] LIMIT 5").unwrap();
483        match Parser::parse(&tokens).unwrap() {
484            Statement::Recall(r) => {
485                assert_eq!(r.filters.len(), 1);
486                assert_eq!(r.filters[0].op, Operator::In);
487                assert_eq!(
488                    r.filters[0].value,
489                    Value::List(vec![
490                        Value::MemoryType(MemoryType::Episodic),
491                        Value::MemoryType(MemoryType::Semantic),
492                    ])
493                );
494            }
495            _ => panic!("expected Recall"),
496        }
497    }
498
499    #[test]
500    fn test_parse_in_operator_text_list() {
501        let tokens = tokenize("RECALL memories WHERE tag IN [\"work\", \"home\"]").unwrap();
502        match Parser::parse(&tokens).unwrap() {
503            Statement::Recall(r) => {
504                assert_eq!(r.filters[0].field, Field::Tag);
505                assert_eq!(r.filters[0].op, Operator::In);
506                assert_eq!(
507                    r.filters[0].value,
508                    Value::List(vec![Value::Text("work".into()), Value::Text("home".into()),])
509                );
510            }
511            _ => panic!("expected Recall"),
512        }
513    }
514
515    #[test]
516    fn test_parse_contains_operator() {
517        let tokens = tokenize("RECALL memories WHERE content CONTAINS \"coffee\"").unwrap();
518        match Parser::parse(&tokens).unwrap() {
519            Statement::Recall(r) => {
520                assert_eq!(r.filters[0].field, Field::Content);
521                assert_eq!(r.filters[0].op, Operator::Contains);
522                assert_eq!(r.filters[0].value, Value::Text("coffee".into()));
523            }
524            _ => panic!("expected Recall"),
525        }
526    }
527
528    #[test]
529    fn test_parse_recall_as_of() {
530        // AS OF <t> lowers to a ValidAt filter carrying the timestamp, on top of
531        // any WHERE filters, and coexists with LIMIT.
532        let tokens =
533            tokenize("RECALL memories WHERE type = semantic LIMIT 5 AS OF 1700000000").unwrap();
534        let stmt = Parser::parse(&tokens).unwrap();
535        match stmt {
536            Statement::Recall(r) => {
537                assert_eq!(r.limit, Some(5));
538                assert_eq!(r.filters.len(), 2);
539                let valid_at = r
540                    .filters
541                    .iter()
542                    .find(|f| f.field == Field::ValidAt)
543                    .expect("expected a ValidAt filter from AS OF");
544                assert_eq!(valid_at.value, Value::Integer(1_700_000_000));
545            }
546            _ => panic!("expected Recall"),
547        }
548    }
549
550    #[test]
551    fn test_parse_recall_as_of_without_where() {
552        let tokens = tokenize("RECALL memories AS OF 42").unwrap();
553        let stmt = Parser::parse(&tokens).unwrap();
554        match stmt {
555            Statement::Recall(r) => {
556                assert_eq!(r.filters.len(), 1);
557                assert_eq!(r.filters[0].field, Field::ValidAt);
558                assert_eq!(r.filters[0].value, Value::Integer(42));
559            }
560            _ => panic!("expected Recall"),
561        }
562    }
563
564    #[test]
565    fn test_parse_recall_similar_to() {
566        let tokens =
567            tokenize(r#"RECALL memories WHERE content ~> "database migration" LIMIT 10"#).unwrap();
568        let stmt = Parser::parse(&tokens).unwrap();
569        match stmt {
570            Statement::Recall(r) => {
571                assert_eq!(r.filters.len(), 1);
572                assert_eq!(r.filters[0].op, Operator::SimilarTo);
573                assert_eq!(r.limit, Some(10));
574            }
575            _ => panic!("expected Recall"),
576        }
577    }
578
579    #[test]
580    fn test_parse_recall_near() {
581        let tokens = tokenize("RECALL memories NEAR [0.1, 0.2, 0.3] LIMIT 10").unwrap();
582        let stmt = Parser::parse(&tokens).unwrap();
583        match stmt {
584            Statement::Recall(r) => {
585                assert_eq!(r.near, Some(vec![0.1, 0.2, 0.3]));
586                assert_eq!(r.limit, Some(10));
587            }
588            _ => panic!("expected Recall"),
589        }
590    }
591
592    #[test]
593    fn test_parse_relate() {
594        let tokens = tokenize(
595            "RELATE 550e8400-e29b-41d4-a716-446655440000 -> 660e8400-e29b-41d4-a716-446655440000 AS caused WITH weight = 0.9"
596        ).unwrap();
597        let stmt = Parser::parse(&tokens).unwrap();
598        match stmt {
599            Statement::Relate(r) => {
600                assert_eq!(r.edge_type, EdgeType::Caused);
601                assert_eq!(r.weight, Some(0.9));
602            }
603            _ => panic!("expected Relate"),
604        }
605    }
606
607    #[test]
608    fn test_parse_forget() {
609        let tokens = tokenize("FORGET 550e8400-e29b-41d4-a716-446655440000").unwrap();
610        let stmt = Parser::parse(&tokens).unwrap();
611        match stmt {
612            Statement::Forget(f) => {
613                assert_eq!(
614                    f.target,
615                    "550e8400-e29b-41d4-a716-446655440000"
616                        .parse::<MemoryId>()
617                        .unwrap()
618                );
619            }
620            _ => panic!("expected Forget"),
621        }
622    }
623
624    #[test]
625    fn test_parse_consolidate() {
626        let tokens =
627            tokenize(r#"CONSOLIDATE WHERE type = episodic AND accessed < "2024-01-01""#).unwrap();
628        let stmt = Parser::parse(&tokens).unwrap();
629        match stmt {
630            Statement::Consolidate(c) => {
631                assert_eq!(c.filters.len(), 2);
632            }
633            _ => panic!("expected Consolidate"),
634        }
635    }
636
637    #[test]
638    fn test_parse_traverse() {
639        let tokens = tokenize(
640            "TRAVERSE 550e8400-e29b-41d4-a716-446655440000 DEPTH 3 WHERE edge_type = caused",
641        )
642        .unwrap();
643        let stmt = Parser::parse(&tokens).unwrap();
644        match stmt {
645            Statement::Traverse(t) => {
646                assert_eq!(t.depth, 3);
647                assert_eq!(t.edge_filter, Some(vec![EdgeType::Caused]));
648            }
649            _ => panic!("expected Traverse"),
650        }
651    }
652}