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