1use 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 Recall,
16 Relate,
17 Forget,
18 Consolidate,
19 Traverse,
20
21 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 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 Eq, Neq, Gt, Lt, Gte, Lte, SimilarTo, Arrow, LParen,
67 RParen,
68 LBracket,
69 RBracket,
70 Comma,
71 Dot,
72 Colon,
73 Semicolon,
74
75 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 if bytes[pos].is_ascii_whitespace() {
94 pos += 1;
95 continue;
96 }
97
98 let start = pos;
99
100 if bytes[pos] == b'"' {
102 pos += 1;
103 while pos < len && bytes[pos] != b'"' {
104 if bytes[pos] == b'\\' {
105 pos += 1; }
107 pos += 1;
108 }
109 if pos >= len {
110 return Err(MenteError::Query("unterminated string literal".into()));
111 }
112 pos += 1; 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 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 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 if bytes[pos].is_ascii_hexdigit() {
171 let saved = pos;
172 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 pos = saved;
189 }
190
191 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 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 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}