1use 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(); 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 if self.at(TokenKind::Near) {
84 self.advance();
85 near = Some(self.parse_vector()?);
86 }
87
88 if self.at(TokenKind::Where) {
90 self.advance();
91 filters = self.parse_filters()?;
92 }
93
94 if self.at(TokenKind::OrderBy) {
96 self.advance();
97 if self.at(TokenKind::By) {
99 self.advance();
100 }
101 let field = self.parse_field()?;
102 let descending = false; order_by = Some(OrderBy { field, descending });
104 }
105
106 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 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(); 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 self.expect(TokenKind::Identifier)?; 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(); let target = self.parse_uuid()?;
175 Ok(Statement::Forget(ForgetStatement { target }))
176 }
177
178 fn parse_consolidate(&mut self) -> MenteResult<Statement> {
179 self.advance(); 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(); 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 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 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 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 let inner = tok.lexeme[1..tok.lexeme.len() - 1].to_string();
305 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 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 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}