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
12/// If `cond` is a pure AND of leaf filters (a single leaf, or ANDs nesting only
13/// leaves and further ANDs), return those leaves flattened. Otherwise return
14/// `None`: the tree has an OR or NOT and must be evaluated as a tree. This keeps
15/// the common `a AND b AND c` clause on the flat-filter fast path.
16fn flatten_pure_and(cond: &Condition) -> Option<Vec<Filter>> {
17    fn collect(c: &Condition, out: &mut Vec<Filter>) -> bool {
18        match c {
19            Condition::Leaf(f) => {
20                out.push(f.clone());
21                true
22            }
23            Condition::And(children) => children.iter().all(|ch| collect(ch, out)),
24            _ => false,
25        }
26    }
27    let mut out = Vec::new();
28    if collect(cond, &mut out) {
29        Some(out)
30    } else {
31        None
32    }
33}
34
35pub struct Parser<'a> {
36    tokens: &'a [Token],
37    pos: usize,
38}
39
40impl<'a> Parser<'a> {
41    pub fn new(tokens: &'a [Token]) -> Self {
42        Self { tokens, pos: 0 }
43    }
44
45    pub fn parse(tokens: &[Token]) -> MenteResult<Statement> {
46        let mut parser = Parser::new(tokens);
47        parser.parse_statement()
48    }
49
50    fn peek(&self) -> &Token {
51        &self.tokens[self.pos.min(self.tokens.len() - 1)]
52    }
53
54    fn advance(&mut self) -> &Token {
55        let tok = &self.tokens[self.pos.min(self.tokens.len() - 1)];
56        if self.pos < self.tokens.len() {
57            self.pos += 1;
58        }
59        tok
60    }
61
62    fn expect(&mut self, kind: TokenKind) -> MenteResult<&Token> {
63        let tok = self.peek();
64        if tok.kind != kind {
65            return Err(MenteError::Query(format!(
66                "expected {:?}, found {:?} ('{}') at position {}",
67                kind, tok.kind, tok.lexeme, tok.position
68            )));
69        }
70        Ok(self.advance())
71    }
72
73    fn at(&self, kind: TokenKind) -> bool {
74        self.peek().kind == kind
75    }
76
77    fn parse_statement(&mut self) -> MenteResult<Statement> {
78        match self.peek().kind {
79            TokenKind::Recall => self.parse_recall(),
80            TokenKind::Relate => self.parse_relate(),
81            TokenKind::Forget => self.parse_forget(),
82            TokenKind::Consolidate => self.parse_consolidate(),
83            TokenKind::Traverse => self.parse_traverse(),
84            _ => Err(MenteError::Query(format!(
85                "expected statement keyword, found {:?} at position {}",
86                self.peek().kind,
87                self.peek().position
88            ))),
89        }
90    }
91
92    fn parse_recall(&mut self) -> MenteResult<Statement> {
93        self.advance(); // RECALL
94
95        // Optional "memories" keyword
96        if self.at(TokenKind::Memories) {
97            self.advance();
98        }
99
100        let mut near = None;
101        let mut filters = Vec::new();
102        let mut condition: Option<Condition> = None;
103        let mut limit = None;
104        let mut order_by = None;
105
106        // NEAR [vector]
107        if self.at(TokenKind::Near) {
108            self.advance();
109            near = Some(self.parse_vector()?);
110        }
111
112        // WHERE clause. Parse the full boolean expression, then keep the common
113        // pure-AND case as a flat filter list (so the planner's index selection is
114        // unchanged); only fall back to the tree when OR, NOT, or grouping appears.
115        if self.at(TokenKind::Where) {
116            self.advance();
117            let cond = self.parse_condition()?;
118            match flatten_pure_and(&cond) {
119                Some(leaves) => filters = leaves,
120                None => condition = Some(cond),
121            }
122        }
123
124        // ORDER BY field
125        if self.at(TokenKind::OrderBy) {
126            self.advance();
127            // consume optional "BY"
128            if self.at(TokenKind::By) {
129                self.advance();
130            }
131            let field = self.parse_field()?;
132            // Optional direction; ASC (the default) or DESC.
133            let descending = if self.at(TokenKind::Desc) {
134                self.advance();
135                true
136            } else {
137                if self.at(TokenKind::Asc) {
138                    self.advance();
139                }
140                false
141            };
142            order_by = Some(OrderBy { field, descending });
143        }
144
145        // LIMIT n
146        if self.at(TokenKind::Limit) {
147            self.advance();
148            let tok = self.advance();
149            let n: usize = tok
150                .lexeme
151                .parse()
152                .map_err(|_| MenteError::Query(format!("invalid limit value: {}", tok.lexeme)))?;
153            limit = Some(n);
154        }
155
156        // AS OF <timestamp> — point-in-time temporal filter. Added as a ValidAt
157        // filter so it flows through the same pipeline as WHERE clauses.
158        if self.at(TokenKind::As) {
159            self.advance();
160            self.expect(TokenKind::Of)?;
161            let tok = self.advance();
162            let t: i64 = tok.lexeme.parse().map_err(|_| {
163                MenteError::Query(format!("invalid AS OF timestamp: {}", tok.lexeme))
164            })?;
165            let valid_at = Filter {
166                field: Field::ValidAt,
167                op: Operator::Eq,
168                value: Value::Integer(t),
169            };
170            // AS OF must AND with whatever WHERE produced. In the tree case, wrap
171            // it in; otherwise it is just another top-level AND leaf.
172            match condition.take() {
173                Some(c) => {
174                    condition = Some(Condition::And(vec![c, Condition::Leaf(valid_at)]));
175                }
176                None => filters.push(valid_at),
177            }
178        }
179
180        Ok(Statement::Recall(RecallStatement {
181            filters,
182            condition,
183            near,
184            limit,
185            order_by,
186        }))
187    }
188
189    fn parse_relate(&mut self) -> MenteResult<Statement> {
190        self.advance(); // RELATE
191
192        let source = self.parse_uuid()?;
193        self.expect(TokenKind::Arrow)?;
194        let target = self.parse_uuid()?;
195        self.expect(TokenKind::As)?;
196        let edge_type = self.parse_edge_type()?;
197
198        let mut weight = None;
199        if self.at(TokenKind::With) {
200            self.advance();
201            // expect "weight = <float>"
202            self.expect(TokenKind::Identifier)?; // "weight"
203            self.expect(TokenKind::Eq)?;
204            let tok = self.advance();
205            let w: f32 = tok
206                .lexeme
207                .parse()
208                .map_err(|_| MenteError::Query(format!("invalid weight value: {}", tok.lexeme)))?;
209            weight = Some(w);
210        }
211
212        Ok(Statement::Relate(RelateStatement {
213            source,
214            target,
215            edge_type,
216            weight,
217        }))
218    }
219
220    fn parse_forget(&mut self) -> MenteResult<Statement> {
221        self.advance(); // FORGET
222        let target = self.parse_uuid()?;
223        Ok(Statement::Forget(ForgetStatement { target }))
224    }
225
226    fn parse_consolidate(&mut self) -> MenteResult<Statement> {
227        self.advance(); // CONSOLIDATE
228        let mut filters = Vec::new();
229        if self.at(TokenKind::Where) {
230            self.advance();
231            filters = self.parse_filters()?;
232        }
233        Ok(Statement::Consolidate(ConsolidateStatement { filters }))
234    }
235
236    fn parse_traverse(&mut self) -> MenteResult<Statement> {
237        self.advance(); // TRAVERSE
238        let start = self.parse_uuid()?;
239
240        self.expect(TokenKind::Depth)?;
241        let tok = self.advance();
242        let depth: usize = tok
243            .lexeme
244            .parse()
245            .map_err(|_| MenteError::Query(format!("invalid depth value: {}", tok.lexeme)))?;
246
247        let mut edge_filter = None;
248        if self.at(TokenKind::Where) {
249            self.advance();
250            // edge_type = <type>
251            self.expect(TokenKind::EdgeType)?;
252            self.expect(TokenKind::Eq)?;
253            let et = self.parse_edge_type()?;
254            edge_filter = Some(vec![et]);
255        }
256
257        Ok(Statement::Traverse(TraverseStatement {
258            start,
259            depth,
260            edge_filter,
261        }))
262    }
263
264    fn parse_filters(&mut self) -> MenteResult<Vec<Filter>> {
265        let mut filters = vec![self.parse_filter()?];
266        while self.at(TokenKind::And) {
267            self.advance();
268            filters.push(self.parse_filter()?);
269        }
270        Ok(filters)
271    }
272
273    /// Parse a boolean WHERE expression with standard precedence:
274    /// OR binds loosest, then AND, then NOT, then a parenthesized group or a leaf
275    /// comparison. `a AND b OR c` parses as `(a AND b) OR c`.
276    fn parse_condition(&mut self) -> MenteResult<Condition> {
277        let mut node = self.parse_and_condition()?;
278        while self.at(TokenKind::Or) {
279            self.advance();
280            let rhs = self.parse_and_condition()?;
281            node = match node {
282                Condition::Or(mut v) => {
283                    v.push(rhs);
284                    Condition::Or(v)
285                }
286                other => Condition::Or(vec![other, rhs]),
287            };
288        }
289        Ok(node)
290    }
291
292    fn parse_and_condition(&mut self) -> MenteResult<Condition> {
293        let mut node = self.parse_not_condition()?;
294        while self.at(TokenKind::And) {
295            self.advance();
296            let rhs = self.parse_not_condition()?;
297            node = match node {
298                Condition::And(mut v) => {
299                    v.push(rhs);
300                    Condition::And(v)
301                }
302                other => Condition::And(vec![other, rhs]),
303            };
304        }
305        Ok(node)
306    }
307
308    fn parse_not_condition(&mut self) -> MenteResult<Condition> {
309        if self.at(TokenKind::Not) {
310            self.advance();
311            let inner = self.parse_not_condition()?;
312            return Ok(Condition::Not(Box::new(inner)));
313        }
314        self.parse_primary_condition()
315    }
316
317    fn parse_primary_condition(&mut self) -> MenteResult<Condition> {
318        if self.at(TokenKind::LParen) {
319            self.advance();
320            let inner = self.parse_condition()?;
321            self.expect(TokenKind::RParen)?;
322            return Ok(inner);
323        }
324        Ok(Condition::Leaf(self.parse_filter()?))
325    }
326
327    fn parse_filter(&mut self) -> MenteResult<Filter> {
328        let field = self.parse_field()?;
329        let op = self.parse_operator()?;
330        let value = if op == Operator::In {
331            self.parse_list_value(&field)?
332        } else {
333            self.parse_value(&field)?
334        };
335        Ok(Filter { field, op, value })
336    }
337
338    fn parse_field(&mut self) -> MenteResult<Field> {
339        let tok = self.advance();
340        match tok.kind {
341            TokenKind::Identifier if tok.lexeme.eq_ignore_ascii_case("content") => {
342                Ok(Field::Content)
343            }
344            TokenKind::Type => Ok(Field::Type),
345            TokenKind::Tag => Ok(Field::Tag),
346            TokenKind::Agent => Ok(Field::Agent),
347            TokenKind::Space => Ok(Field::Space),
348            TokenKind::Salience => Ok(Field::Salience),
349            TokenKind::Confidence => Ok(Field::Confidence),
350            TokenKind::Created => Ok(Field::Created),
351            TokenKind::Accessed => Ok(Field::Accessed),
352            _ => Err(MenteError::Query(format!(
353                "expected field name, found '{}' at position {}",
354                tok.lexeme, tok.position
355            ))),
356        }
357    }
358
359    fn parse_operator(&mut self) -> MenteResult<Operator> {
360        let tok = self.advance();
361        match tok.kind {
362            TokenKind::Eq => Ok(Operator::Eq),
363            TokenKind::Neq => Ok(Operator::Neq),
364            TokenKind::Gt => Ok(Operator::Gt),
365            TokenKind::Lt => Ok(Operator::Lt),
366            TokenKind::Gte => Ok(Operator::Gte),
367            TokenKind::Lte => Ok(Operator::Lte),
368            TokenKind::SimilarTo => Ok(Operator::SimilarTo),
369            TokenKind::In => Ok(Operator::In),
370            TokenKind::Contains => Ok(Operator::Contains),
371            _ => Err(MenteError::Query(format!(
372                "expected operator, found '{}' at position {}",
373                tok.lexeme, tok.position
374            ))),
375        }
376    }
377
378    /// Parse a bracketed, comma-separated list of scalar values for `IN`.
379    fn parse_list_value(&mut self, field: &Field) -> MenteResult<Value> {
380        self.expect(TokenKind::LBracket)?;
381        let mut items = Vec::new();
382        if !self.at(TokenKind::RBracket) {
383            loop {
384                items.push(self.parse_value(field)?);
385                if self.at(TokenKind::Comma) {
386                    self.advance();
387                } else {
388                    break;
389                }
390            }
391        }
392        self.expect(TokenKind::RBracket)?;
393        Ok(Value::List(items))
394    }
395
396    fn parse_value(&mut self, field: &Field) -> MenteResult<Value> {
397        // For Type field, parse as MemoryType
398        if *field == Field::Type {
399            return self.parse_memory_type_value();
400        }
401
402        let tok = self.advance();
403        match tok.kind {
404            TokenKind::StringLit => {
405                // Strip surrounding quotes
406                let inner = tok.lexeme[1..tok.lexeme.len() - 1].to_string();
407                // Check if this looks like a UUID inside quotes
408                if let Ok(uuid) = inner.parse::<MemoryId>() {
409                    return Ok(Value::Uuid(uuid.into()));
410                }
411                Ok(Value::Text(inner))
412            }
413            TokenKind::IntegerLit => {
414                let n: i64 = tok
415                    .lexeme
416                    .parse()
417                    .map_err(|_| MenteError::Query(format!("invalid integer: {}", tok.lexeme)))?;
418                Ok(Value::Integer(n))
419            }
420            TokenKind::FloatLit => {
421                let n: f64 = tok
422                    .lexeme
423                    .parse()
424                    .map_err(|_| MenteError::Query(format!("invalid float: {}", tok.lexeme)))?;
425                Ok(Value::Number(n))
426            }
427            TokenKind::UuidLit => {
428                let uuid: Uuid = tok
429                    .lexeme
430                    .parse()
431                    .map_err(|_| MenteError::Query(format!("invalid UUID: {}", tok.lexeme)))?;
432                Ok(Value::Uuid(uuid))
433            }
434            TokenKind::Identifier => {
435                let lower = tok.lexeme.to_lowercase();
436                match lower.as_str() {
437                    "true" => Ok(Value::Bool(true)),
438                    "false" => Ok(Value::Bool(false)),
439                    _ => Ok(Value::Text(tok.lexeme.clone())),
440                }
441            }
442            TokenKind::LBracket => {
443                // put back and parse as vector
444                self.pos -= 1;
445                let v = self.parse_vector()?;
446                Ok(Value::Vector(v))
447            }
448            _ => Err(MenteError::Query(format!(
449                "expected value, found '{}' at position {}",
450                tok.lexeme, tok.position
451            ))),
452        }
453    }
454
455    fn parse_memory_type_value(&mut self) -> MenteResult<Value> {
456        let tok = self.advance();
457        let name = match tok.kind {
458            TokenKind::Identifier | TokenKind::StringLit => {
459                if tok.kind == TokenKind::StringLit {
460                    tok.lexeme[1..tok.lexeme.len() - 1].to_string()
461                } else {
462                    tok.lexeme.clone()
463                }
464            }
465            _ => {
466                return Err(MenteError::Query(format!(
467                    "expected memory type, found '{}' at position {}",
468                    tok.lexeme, tok.position
469                )));
470            }
471        };
472
473        let mt = match name.to_lowercase().as_str() {
474            "episodic" => MemoryType::Episodic,
475            "semantic" => MemoryType::Semantic,
476            "procedural" => MemoryType::Procedural,
477            "antipattern" | "anti_pattern" => MemoryType::AntiPattern,
478            "reasoning" => MemoryType::Reasoning,
479            "correction" => MemoryType::Correction,
480            _ => {
481                return Err(MenteError::Query(format!("unknown memory type: {}", name)));
482            }
483        };
484        Ok(Value::MemoryType(mt))
485    }
486
487    fn parse_edge_type(&mut self) -> MenteResult<EdgeType> {
488        let tok = self.advance();
489        let name = match tok.kind {
490            TokenKind::Identifier | TokenKind::StringLit => {
491                if tok.kind == TokenKind::StringLit {
492                    tok.lexeme[1..tok.lexeme.len() - 1].to_string()
493                } else {
494                    tok.lexeme.clone()
495                }
496            }
497            _ => {
498                return Err(MenteError::Query(format!(
499                    "expected edge type, found '{}' at position {}",
500                    tok.lexeme, tok.position
501                )));
502            }
503        };
504
505        match name.to_lowercase().as_str() {
506            "caused" => Ok(EdgeType::Caused),
507            "before" => Ok(EdgeType::Before),
508            "related" => Ok(EdgeType::Related),
509            "contradicts" => Ok(EdgeType::Contradicts),
510            "supports" => Ok(EdgeType::Supports),
511            "supersedes" => Ok(EdgeType::Supersedes),
512            "derived" => Ok(EdgeType::Derived),
513            "partof" | "part_of" => Ok(EdgeType::PartOf),
514            _ => Err(MenteError::Query(format!("unknown edge type: {}", name))),
515        }
516    }
517
518    fn parse_uuid(&mut self) -> MenteResult<MemoryId> {
519        let tok = self.advance();
520        match tok.kind {
521            TokenKind::UuidLit => tok
522                .lexeme
523                .parse()
524                .map_err(|_| MenteError::Query(format!("invalid UUID: {}", tok.lexeme))),
525            TokenKind::StringLit => {
526                let inner = &tok.lexeme[1..tok.lexeme.len() - 1];
527                inner.parse().map_err(|_| {
528                    MenteError::Query(format!("invalid UUID in string: {}", tok.lexeme))
529                })
530            }
531            _ => Err(MenteError::Query(format!(
532                "expected UUID, found '{}' at position {}",
533                tok.lexeme, tok.position
534            ))),
535        }
536    }
537
538    fn parse_vector(&mut self) -> MenteResult<Vec<f32>> {
539        self.expect(TokenKind::LBracket)?;
540        let mut values = Vec::new();
541        if !self.at(TokenKind::RBracket) {
542            let tok = self.advance();
543            let v: f32 = tok.lexeme.parse().map_err(|_| {
544                MenteError::Query(format!("invalid float in vector: {}", tok.lexeme))
545            })?;
546            values.push(v);
547            while self.at(TokenKind::Comma) {
548                self.advance();
549                let tok = self.advance();
550                let v: f32 = tok.lexeme.parse().map_err(|_| {
551                    MenteError::Query(format!("invalid float in vector: {}", tok.lexeme))
552                })?;
553                values.push(v);
554            }
555        }
556        self.expect(TokenKind::RBracket)?;
557        Ok(values)
558    }
559}
560
561#[cfg(test)]
562mod tests {
563    use super::*;
564    use crate::lexer::tokenize;
565
566    #[test]
567    fn test_parse_recall_with_type_filter() {
568        let tokens = tokenize("RECALL memories WHERE type = episodic LIMIT 5").unwrap();
569        let stmt = Parser::parse(&tokens).unwrap();
570        match stmt {
571            Statement::Recall(r) => {
572                assert_eq!(r.filters.len(), 1);
573                assert_eq!(r.filters[0].field, Field::Type);
574                assert_eq!(r.filters[0].value, Value::MemoryType(MemoryType::Episodic));
575                assert_eq!(r.limit, Some(5));
576            }
577            _ => panic!("expected Recall"),
578        }
579    }
580
581    #[test]
582    fn test_pure_and_stays_on_flat_fast_path() {
583        // A pure AND clause must flatten to `filters` with `condition == None`, so
584        // the planner keeps its leaf-based index optimizations.
585        let tokens = tokenize("RECALL WHERE type = semantic AND tag = \"x\" LIMIT 5").unwrap();
586        match Parser::parse(&tokens).unwrap() {
587            Statement::Recall(r) => {
588                assert_eq!(r.filters.len(), 2, "both leaves flattened");
589                assert!(r.condition.is_none(), "pure AND must not build a tree");
590            }
591            _ => panic!("expected Recall"),
592        }
593    }
594
595    #[test]
596    fn test_order_by_direction() {
597        // DESC parses descending; ASC and no-direction parse ascending.
598        let cases = [
599            (
600                "RECALL WHERE type = semantic ORDER BY salience DESC LIMIT 5",
601                true,
602            ),
603            (
604                "RECALL WHERE type = semantic ORDER BY salience ASC LIMIT 5",
605                false,
606            ),
607            (
608                "RECALL WHERE type = semantic ORDER BY created LIMIT 5",
609                false,
610            ),
611        ];
612        for (q, descending) in cases {
613            let tokens = tokenize(q).unwrap();
614            match Parser::parse(&tokens).unwrap() {
615                Statement::Recall(r) => {
616                    let ob = r.order_by.expect("order_by present");
617                    assert_eq!(ob.descending, descending, "for query: {q}");
618                }
619                _ => panic!("expected Recall"),
620            }
621        }
622    }
623
624    #[test]
625    fn test_or_builds_condition_tree() {
626        let tokens = tokenize("RECALL WHERE type = semantic OR type = procedural LIMIT 5").unwrap();
627        match Parser::parse(&tokens).unwrap() {
628            Statement::Recall(r) => {
629                assert!(r.filters.is_empty(), "OR clause is carried by the tree");
630                match r.condition {
631                    Some(Condition::Or(branches)) => assert_eq!(branches.len(), 2),
632                    other => panic!("expected Or condition, got {other:?}"),
633                }
634            }
635            _ => panic!("expected Recall"),
636        }
637    }
638
639    #[test]
640    fn test_grouping_and_not_precedence() {
641        // (a OR b) AND NOT c  =>  And[ Or[a, b], Not(c) ]
642        let tokens = tokenize(
643            "RECALL WHERE (type = semantic OR type = procedural) AND NOT tag = \"x\" LIMIT 5",
644        )
645        .unwrap();
646        match Parser::parse(&tokens).unwrap() {
647            Statement::Recall(r) => match r.condition {
648                Some(Condition::And(parts)) => {
649                    assert_eq!(parts.len(), 2);
650                    assert!(matches!(parts[0], Condition::Or(_)), "left is the OR group");
651                    assert!(matches!(parts[1], Condition::Not(_)), "right is the NOT");
652                }
653                other => panic!("expected And condition, got {other:?}"),
654            },
655            _ => panic!("expected Recall"),
656        }
657    }
658
659    #[test]
660    fn test_parse_in_operator() {
661        let tokens =
662            tokenize("RECALL memories WHERE type IN [episodic, semantic] LIMIT 5").unwrap();
663        match Parser::parse(&tokens).unwrap() {
664            Statement::Recall(r) => {
665                assert_eq!(r.filters.len(), 1);
666                assert_eq!(r.filters[0].op, Operator::In);
667                assert_eq!(
668                    r.filters[0].value,
669                    Value::List(vec![
670                        Value::MemoryType(MemoryType::Episodic),
671                        Value::MemoryType(MemoryType::Semantic),
672                    ])
673                );
674            }
675            _ => panic!("expected Recall"),
676        }
677    }
678
679    #[test]
680    fn test_parse_in_operator_text_list() {
681        let tokens = tokenize("RECALL memories WHERE tag IN [\"work\", \"home\"]").unwrap();
682        match Parser::parse(&tokens).unwrap() {
683            Statement::Recall(r) => {
684                assert_eq!(r.filters[0].field, Field::Tag);
685                assert_eq!(r.filters[0].op, Operator::In);
686                assert_eq!(
687                    r.filters[0].value,
688                    Value::List(vec![Value::Text("work".into()), Value::Text("home".into()),])
689                );
690            }
691            _ => panic!("expected Recall"),
692        }
693    }
694
695    #[test]
696    fn test_parse_contains_operator() {
697        let tokens = tokenize("RECALL memories WHERE content CONTAINS \"coffee\"").unwrap();
698        match Parser::parse(&tokens).unwrap() {
699            Statement::Recall(r) => {
700                assert_eq!(r.filters[0].field, Field::Content);
701                assert_eq!(r.filters[0].op, Operator::Contains);
702                assert_eq!(r.filters[0].value, Value::Text("coffee".into()));
703            }
704            _ => panic!("expected Recall"),
705        }
706    }
707
708    #[test]
709    fn test_parse_recall_as_of() {
710        // AS OF <t> lowers to a ValidAt filter carrying the timestamp, on top of
711        // any WHERE filters, and coexists with LIMIT.
712        let tokens =
713            tokenize("RECALL memories WHERE type = semantic LIMIT 5 AS OF 1700000000").unwrap();
714        let stmt = Parser::parse(&tokens).unwrap();
715        match stmt {
716            Statement::Recall(r) => {
717                assert_eq!(r.limit, Some(5));
718                assert_eq!(r.filters.len(), 2);
719                let valid_at = r
720                    .filters
721                    .iter()
722                    .find(|f| f.field == Field::ValidAt)
723                    .expect("expected a ValidAt filter from AS OF");
724                assert_eq!(valid_at.value, Value::Integer(1_700_000_000));
725            }
726            _ => panic!("expected Recall"),
727        }
728    }
729
730    #[test]
731    fn test_parse_recall_as_of_without_where() {
732        let tokens = tokenize("RECALL memories AS OF 42").unwrap();
733        let stmt = Parser::parse(&tokens).unwrap();
734        match stmt {
735            Statement::Recall(r) => {
736                assert_eq!(r.filters.len(), 1);
737                assert_eq!(r.filters[0].field, Field::ValidAt);
738                assert_eq!(r.filters[0].value, Value::Integer(42));
739            }
740            _ => panic!("expected Recall"),
741        }
742    }
743
744    #[test]
745    fn test_parse_recall_similar_to() {
746        let tokens =
747            tokenize(r#"RECALL memories WHERE content ~> "database migration" LIMIT 10"#).unwrap();
748        let stmt = Parser::parse(&tokens).unwrap();
749        match stmt {
750            Statement::Recall(r) => {
751                assert_eq!(r.filters.len(), 1);
752                assert_eq!(r.filters[0].op, Operator::SimilarTo);
753                assert_eq!(r.limit, Some(10));
754            }
755            _ => panic!("expected Recall"),
756        }
757    }
758
759    #[test]
760    fn test_parse_recall_near() {
761        let tokens = tokenize("RECALL memories NEAR [0.1, 0.2, 0.3] LIMIT 10").unwrap();
762        let stmt = Parser::parse(&tokens).unwrap();
763        match stmt {
764            Statement::Recall(r) => {
765                assert_eq!(r.near, Some(vec![0.1, 0.2, 0.3]));
766                assert_eq!(r.limit, Some(10));
767            }
768            _ => panic!("expected Recall"),
769        }
770    }
771
772    #[test]
773    fn test_parse_relate() {
774        let tokens = tokenize(
775            "RELATE 550e8400-e29b-41d4-a716-446655440000 -> 660e8400-e29b-41d4-a716-446655440000 AS caused WITH weight = 0.9"
776        ).unwrap();
777        let stmt = Parser::parse(&tokens).unwrap();
778        match stmt {
779            Statement::Relate(r) => {
780                assert_eq!(r.edge_type, EdgeType::Caused);
781                assert_eq!(r.weight, Some(0.9));
782            }
783            _ => panic!("expected Relate"),
784        }
785    }
786
787    #[test]
788    fn test_parse_forget() {
789        let tokens = tokenize("FORGET 550e8400-e29b-41d4-a716-446655440000").unwrap();
790        let stmt = Parser::parse(&tokens).unwrap();
791        match stmt {
792            Statement::Forget(f) => {
793                assert_eq!(
794                    f.target,
795                    "550e8400-e29b-41d4-a716-446655440000"
796                        .parse::<MemoryId>()
797                        .unwrap()
798                );
799            }
800            _ => panic!("expected Forget"),
801        }
802    }
803
804    #[test]
805    fn test_parse_consolidate() {
806        let tokens =
807            tokenize(r#"CONSOLIDATE WHERE type = episodic AND accessed < "2024-01-01""#).unwrap();
808        let stmt = Parser::parse(&tokens).unwrap();
809        match stmt {
810            Statement::Consolidate(c) => {
811                assert_eq!(c.filters.len(), 2);
812            }
813            _ => panic!("expected Consolidate"),
814        }
815    }
816
817    #[test]
818    fn test_parse_traverse() {
819        let tokens = tokenize(
820            "TRAVERSE 550e8400-e29b-41d4-a716-446655440000 DEPTH 3 WHERE edge_type = caused",
821        )
822        .unwrap();
823        let stmt = Parser::parse(&tokens).unwrap();
824        match stmt {
825            Statement::Traverse(t) => {
826                assert_eq!(t.depth, 3);
827                assert_eq!(t.edge_filter, Some(vec![EdgeType::Caused]));
828            }
829            _ => panic!("expected Traverse"),
830        }
831    }
832}