Skip to main content

mentedb_query/
lexer.rs

1//! Hand-written lexer for MQL.
2
3use mentedb_core::error::{MenteError, MenteResult};
4
5#[derive(Debug, Clone, PartialEq)]
6pub struct Token {
7    pub kind: TokenKind,
8    pub lexeme: String,
9    pub position: usize,
10}
11
12#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
13pub enum TokenKind {
14    // Statements
15    Recall,
16    Relate,
17    Forget,
18    Consolidate,
19    Traverse,
20
21    // Clauses
22    Where,
23    And,
24    Or,
25    Not,
26    In,
27    Contains,
28    Near,
29    Within,
30    Limit,
31    OrderBy,
32    Asc,
33    Desc,
34    As,
35    Of,
36    From,
37    To,
38    With,
39
40    // Keywords
41    Agent,
42    Space,
43    Type,
44    Tag,
45    Salience,
46    Confidence,
47    Created,
48    Accessed,
49    Depth,
50    Hops,
51    Memories,
52    By,
53    EdgeType,
54
55    // Operators
56    Eq,        // =
57    Neq,       // !=
58    Gt,        // >
59    Lt,        // <
60    Gte,       // >=
61    Lte,       // <=
62    SimilarTo, // ~>
63    Arrow,     // ->
64
65    // Punctuation
66    LParen,
67    RParen,
68    LBracket,
69    RBracket,
70    Comma,
71    Dot,
72    Colon,
73    Semicolon,
74
75    // Literals
76    StringLit,
77    IntegerLit,
78    FloatLit,
79    Identifier,
80    UuidLit,
81
82    Eof,
83}
84
85pub fn tokenize(input: &str) -> MenteResult<Vec<Token>> {
86    let mut tokens = Vec::new();
87    let bytes = input.as_bytes();
88    let len = bytes.len();
89    let mut pos = 0;
90
91    while pos < len {
92        // Skip whitespace
93        if bytes[pos].is_ascii_whitespace() {
94            pos += 1;
95            continue;
96        }
97
98        let start = pos;
99
100        // String literal
101        if bytes[pos] == b'"' {
102            pos += 1;
103            while pos < len && bytes[pos] != b'"' {
104                if bytes[pos] == b'\\' {
105                    pos += 1; // skip escaped char
106                }
107                pos += 1;
108            }
109            if pos >= len {
110                return Err(MenteError::Query("unterminated string literal".into()));
111            }
112            pos += 1; // closing quote
113            let lexeme = input[start..pos].to_string();
114            tokens.push(Token {
115                kind: TokenKind::StringLit,
116                lexeme,
117                position: start,
118            });
119            continue;
120        }
121
122        // Two-char operators
123        if pos + 1 < len {
124            let two = &input[start..start + 2];
125            let kind = match two {
126                "!=" => Some(TokenKind::Neq),
127                ">=" => Some(TokenKind::Gte),
128                "<=" => Some(TokenKind::Lte),
129                "~>" => Some(TokenKind::SimilarTo),
130                "->" => Some(TokenKind::Arrow),
131                _ => None,
132            };
133            if let Some(k) = kind {
134                tokens.push(Token {
135                    kind: k,
136                    lexeme: two.to_string(),
137                    position: start,
138                });
139                pos += 2;
140                continue;
141            }
142        }
143
144        // Single-char operators/punctuation
145        let single = match bytes[pos] {
146            b'=' => Some(TokenKind::Eq),
147            b'>' => Some(TokenKind::Gt),
148            b'<' => Some(TokenKind::Lt),
149            b'(' => Some(TokenKind::LParen),
150            b')' => Some(TokenKind::RParen),
151            b'[' => Some(TokenKind::LBracket),
152            b']' => Some(TokenKind::RBracket),
153            b',' => Some(TokenKind::Comma),
154            b'.' => Some(TokenKind::Dot),
155            b':' => Some(TokenKind::Colon),
156            b';' => Some(TokenKind::Semicolon),
157            _ => None,
158        };
159        if let Some(k) = single {
160            tokens.push(Token {
161                kind: k,
162                lexeme: input[start..start + 1].to_string(),
163                position: start,
164            });
165            pos += 1;
166            continue;
167        }
168
169        // Try UUID first: if we see a hex digit, speculatively scan for UUID pattern
170        if bytes[pos].is_ascii_hexdigit() {
171            let saved = pos;
172            // Consume alphanumeric + hyphens to check for UUID
173            while pos < len
174                && (bytes[pos].is_ascii_alphanumeric() || bytes[pos] == b'_' || bytes[pos] == b'-')
175            {
176                pos += 1;
177            }
178            let candidate = &input[saved..pos];
179            if is_uuid_like(candidate) {
180                tokens.push(Token {
181                    kind: TokenKind::UuidLit,
182                    lexeme: candidate.to_string(),
183                    position: start,
184                });
185                continue;
186            }
187            // Not a UUID — reset and fall through to number/identifier parsing
188            pos = saved;
189        }
190
191        // Numbers (may start with - for negative)
192        if bytes[pos].is_ascii_digit()
193            || (bytes[pos] == b'-' && pos + 1 < len && bytes[pos + 1].is_ascii_digit())
194        {
195            if bytes[pos] == b'-' {
196                pos += 1;
197            }
198            while pos < len && bytes[pos].is_ascii_digit() {
199                pos += 1;
200            }
201            let mut is_float = false;
202            if pos < len && bytes[pos] == b'.' && pos + 1 < len && bytes[pos + 1].is_ascii_digit() {
203                is_float = true;
204                pos += 1;
205                while pos < len && bytes[pos].is_ascii_digit() {
206                    pos += 1;
207                }
208            }
209            let lexeme = input[start..pos].to_string();
210            let kind = if is_float {
211                TokenKind::FloatLit
212            } else {
213                TokenKind::IntegerLit
214            };
215            tokens.push(Token {
216                kind,
217                lexeme,
218                position: start,
219            });
220            continue;
221        }
222
223        // Identifiers, keywords
224        if bytes[pos].is_ascii_alphanumeric() || bytes[pos] == b'_' {
225            while pos < len
226                && (bytes[pos].is_ascii_alphanumeric() || bytes[pos] == b'_' || bytes[pos] == b'-')
227            {
228                pos += 1;
229            }
230            let lexeme = input[start..pos].to_string();
231
232            let kind = match lexeme.to_lowercase().as_str() {
233                "recall" => TokenKind::Recall,
234                "relate" => TokenKind::Relate,
235                "forget" => TokenKind::Forget,
236                "consolidate" => TokenKind::Consolidate,
237                "traverse" => TokenKind::Traverse,
238                "where" => TokenKind::Where,
239                "and" => TokenKind::And,
240                "or" => TokenKind::Or,
241                "not" => TokenKind::Not,
242                "in" => TokenKind::In,
243                "contains" => TokenKind::Contains,
244                "near" => TokenKind::Near,
245                "within" => TokenKind::Within,
246                "limit" => TokenKind::Limit,
247                "order" => TokenKind::OrderBy,
248                "asc" => TokenKind::Asc,
249                "desc" => TokenKind::Desc,
250                "as" => TokenKind::As,
251                "of" => TokenKind::Of,
252                "from" => TokenKind::From,
253                "to" => TokenKind::To,
254                "with" => TokenKind::With,
255                "agent" => TokenKind::Agent,
256                "space" => TokenKind::Space,
257                "type" => TokenKind::Type,
258                "tag" => TokenKind::Tag,
259                "salience" => TokenKind::Salience,
260                "confidence" => TokenKind::Confidence,
261                "created" => TokenKind::Created,
262                "accessed" => TokenKind::Accessed,
263                "depth" => TokenKind::Depth,
264                "hops" => TokenKind::Hops,
265                "memories" => TokenKind::Memories,
266                "by" => TokenKind::By,
267                "edge_type" => TokenKind::EdgeType,
268                _ => TokenKind::Identifier,
269            };
270            tokens.push(Token {
271                kind,
272                lexeme,
273                position: start,
274            });
275            continue;
276        }
277
278        return Err(MenteError::Query(format!(
279            "unexpected character '{}' at position {}",
280            bytes[pos] as char, pos
281        )));
282    }
283
284    tokens.push(Token {
285        kind: TokenKind::Eof,
286        lexeme: String::new(),
287        position: pos,
288    });
289    Ok(tokens)
290}
291
292fn is_uuid_like(s: &str) -> bool {
293    // UUID format: 8-4-4-4-12 hex chars (with dashes)
294    if s.len() != 36 {
295        return false;
296    }
297    let parts: Vec<&str> = s.split('-').collect();
298    if parts.len() != 5 {
299        return false;
300    }
301    let expected_lens = [8, 4, 4, 4, 12];
302    for (part, &expected) in parts.iter().zip(&expected_lens) {
303        if part.len() != expected || !part.chars().all(|c| c.is_ascii_hexdigit()) {
304            return false;
305        }
306    }
307    true
308}
309
310#[cfg(test)]
311mod tests {
312    use super::*;
313
314    #[test]
315    fn test_recall_statement_tokens() {
316        let tokens = tokenize("RECALL memories WHERE type = episodic LIMIT 10").unwrap();
317        assert_eq!(tokens[0].kind, TokenKind::Recall);
318        assert_eq!(tokens[1].kind, TokenKind::Memories);
319        assert_eq!(tokens[2].kind, TokenKind::Where);
320        assert_eq!(tokens[3].kind, TokenKind::Type);
321        assert_eq!(tokens[4].kind, TokenKind::Eq);
322        assert_eq!(tokens[5].kind, TokenKind::Identifier);
323        assert_eq!(tokens[5].lexeme, "episodic");
324        assert_eq!(tokens[6].kind, TokenKind::Limit);
325        assert_eq!(tokens[7].kind, TokenKind::IntegerLit);
326        assert_eq!(tokens[8].kind, TokenKind::Eof);
327    }
328
329    #[test]
330    fn test_string_literal() {
331        let tokens = tokenize(r#"content ~> "database migration""#).unwrap();
332        assert_eq!(tokens[0].kind, TokenKind::Identifier);
333        assert_eq!(tokens[1].kind, TokenKind::SimilarTo);
334        assert_eq!(tokens[2].kind, TokenKind::StringLit);
335        assert_eq!(tokens[2].lexeme, r#""database migration""#);
336    }
337
338    #[test]
339    fn test_operators() {
340        let tokens = tokenize("= != > < >= <= ~> ->").unwrap();
341        let kinds: Vec<TokenKind> = tokens.iter().map(|t| t.kind).collect();
342        assert_eq!(
343            kinds,
344            vec![
345                TokenKind::Eq,
346                TokenKind::Neq,
347                TokenKind::Gt,
348                TokenKind::Lt,
349                TokenKind::Gte,
350                TokenKind::Lte,
351                TokenKind::SimilarTo,
352                TokenKind::Arrow,
353                TokenKind::Eof,
354            ]
355        );
356    }
357
358    #[test]
359    fn test_uuid_token() {
360        let tokens = tokenize("550e8400-e29b-41d4-a716-446655440000").unwrap();
361        assert_eq!(tokens[0].kind, TokenKind::UuidLit);
362    }
363
364    #[test]
365    fn test_float_literal() {
366        let tokens = tokenize("0.1 42 3.14").unwrap();
367        assert_eq!(tokens[0].kind, TokenKind::FloatLit);
368        assert_eq!(tokens[1].kind, TokenKind::IntegerLit);
369        assert_eq!(tokens[2].kind, TokenKind::FloatLit);
370    }
371
372    #[test]
373    fn test_vector_literal() {
374        let tokens = tokenize("[0.1, 0.2, 0.3]").unwrap();
375        assert_eq!(tokens[0].kind, TokenKind::LBracket);
376        assert_eq!(tokens[1].kind, TokenKind::FloatLit);
377        assert_eq!(tokens[2].kind, TokenKind::Comma);
378        assert_eq!(tokens[5].kind, TokenKind::FloatLit);
379        assert_eq!(tokens[6].kind, TokenKind::RBracket);
380    }
381
382    #[test]
383    fn test_punctuation() {
384        let tokens = tokenize("( ) [ ] , . : ;").unwrap();
385        let kinds: Vec<TokenKind> = tokens.iter().map(|t| t.kind).collect();
386        assert_eq!(
387            kinds,
388            vec![
389                TokenKind::LParen,
390                TokenKind::RParen,
391                TokenKind::LBracket,
392                TokenKind::RBracket,
393                TokenKind::Comma,
394                TokenKind::Dot,
395                TokenKind::Colon,
396                TokenKind::Semicolon,
397                TokenKind::Eof,
398            ]
399        );
400    }
401}