Skip to main content

basalt/sql/
lexer.rs

1/// Tokenizer for the Basalt SQL dialect.
2#[derive(Debug, Clone, PartialEq)]
3pub enum Token {
4    Ident(String),
5    /// Quoted identifier: "col name" or [col name]
6    QuotedIdent(String),
7    Integer(i128),
8    Real(f64),
9    Str(String),
10    // punctuation / operators
11    LParen,
12    RParen,
13    Comma,
14    Semi,
15    Star,
16    Plus,
17    Minus,
18    Slash,
19    Percent,
20    Eq,
21    NotEq, // != or <>
22    Lt,
23    LtEq,
24    Gt,
25    GtEq,
26    Dot,
27    Eof,
28}
29
30#[derive(Debug, Clone)]
31pub struct TokenSpan {
32    pub token: Token,
33    pub offset: usize,
34}
35
36#[derive(Debug)]
37pub struct LexError {
38    pub message: String,
39    pub offset: usize,
40}
41
42pub fn lex(input: &str) -> Result<Vec<TokenSpan>, LexError> {
43    let bytes = input.as_bytes();
44    let mut i = 0usize;
45    let mut out = Vec::new();
46    while i < bytes.len() {
47        let b = bytes[i];
48        match b {
49            b' ' | b'\t' | b'\r' | b'\n' => i += 1,
50            b'-' if i + 1 < bytes.len() && bytes[i + 1] == b'-' => {
51                // line comment
52                while i < bytes.len() && bytes[i] != b'\n' {
53                    i += 1;
54                }
55            }
56            b'/' if i + 1 < bytes.len() && bytes[i + 1] == b'*' => {
57                let mut closed = false;
58                i += 2;
59                while i + 1 < bytes.len() {
60                    if bytes[i] == b'*' && bytes[i + 1] == b'/' {
61                        closed = true;
62                        i += 2;
63                        break;
64                    }
65                    i += 1;
66                }
67                if !closed {
68                    return Err(LexError {
69                        message: "unterminated block comment".into(),
70                        offset: i,
71                    });
72                }
73            }
74            b'(' => {
75                out.push_tok(Token::LParen, i);
76                i += 1;
77            }
78            b')' => {
79                out.push_tok(Token::RParen, i);
80                i += 1;
81            }
82            b',' => {
83                out.push_tok(Token::Comma, i);
84                i += 1;
85            }
86            b';' => {
87                out.push_tok(Token::Semi, i);
88                i += 1;
89            }
90            b'*' => {
91                out.push_tok(Token::Star, i);
92                i += 1;
93            }
94            b'+' => {
95                out.push_tok(Token::Plus, i);
96                i += 1;
97            }
98            b'-' => {
99                out.push_tok(Token::Minus, i);
100                i += 1;
101            }
102            b'/' => {
103                out.push_tok(Token::Slash, i);
104                i += 1;
105            }
106            b'%' => {
107                out.push_tok(Token::Percent, i);
108                i += 1;
109            }
110            b'.' if i + 1 < bytes.len() && bytes[i + 1].is_ascii_digit() => {
111                let (token, end) = scan_number(input, i)?;
112                out.push_tok(token, i);
113                i = end;
114            }
115            b'.' => {
116                out.push_tok(Token::Dot, i);
117                i += 1;
118            }
119            b'=' => {
120                out.push_tok(Token::Eq, i);
121                i += 1;
122            }
123            b'!' => {
124                if i + 1 < bytes.len() && bytes[i + 1] == b'=' {
125                    out.push_tok(Token::NotEq, i);
126                    i += 2;
127                } else {
128                    return Err(LexError {
129                        message: "unexpected '!'".into(),
130                        offset: i,
131                    });
132                }
133            }
134            b'<' => {
135                if i + 1 < bytes.len() && bytes[i + 1] == b'=' {
136                    out.push_tok(Token::LtEq, i);
137                    i += 2;
138                } else if i + 1 < bytes.len() && bytes[i + 1] == b'>' {
139                    out.push_tok(Token::NotEq, i);
140                    i += 2;
141                } else {
142                    out.push_tok(Token::Lt, i);
143                    i += 1;
144                }
145            }
146            b'>' => {
147                if i + 1 < bytes.len() && bytes[i + 1] == b'=' {
148                    out.push_tok(Token::GtEq, i);
149                    i += 2;
150                } else {
151                    out.push_tok(Token::Gt, i);
152                    i += 1;
153                }
154            }
155            b'\'' => {
156                // single-quoted string with '' escape
157                let start = i;
158                i += 1;
159                let mut s = String::new();
160                loop {
161                    if i >= bytes.len() {
162                        return Err(LexError {
163                            message: "unterminated string literal".into(),
164                            offset: start,
165                        });
166                    }
167                    if bytes[i] == b'\'' {
168                        if i + 1 < bytes.len() && bytes[i + 1] == b'\'' {
169                            s.push('\'');
170                            i += 2;
171                        } else {
172                            i += 1;
173                            break;
174                        }
175                    } else {
176                        let ch_len = utf8_len(bytes[i]);
177                        s.push_str(std::str::from_utf8(&bytes[i..i + ch_len]).map_err(|_| {
178                            LexError {
179                                message: "invalid UTF-8".into(),
180                                offset: i,
181                            }
182                        })?);
183                        i += ch_len;
184                    }
185                }
186                out.push_tok(Token::Str(s), start);
187            }
188            b'"' => {
189                let start = i;
190                i += 1;
191                let mut s = String::new();
192                loop {
193                    if i >= bytes.len() {
194                        return Err(LexError {
195                            message: "unterminated quoted identifier".into(),
196                            offset: start,
197                        });
198                    }
199                    if bytes[i] == b'"' {
200                        if i + 1 < bytes.len() && bytes[i + 1] == b'"' {
201                            s.push('"');
202                            i += 2;
203                        } else {
204                            i += 1;
205                            break;
206                        }
207                    } else {
208                        let ch_len = utf8_len(bytes[i]);
209                        s.push_str(std::str::from_utf8(&bytes[i..i + ch_len]).map_err(|_| {
210                            LexError {
211                                message: "invalid UTF-8".into(),
212                                offset: i,
213                            }
214                        })?);
215                        i += ch_len;
216                    }
217                }
218                out.push_tok(Token::QuotedIdent(s), start);
219            }
220            b'[' => {
221                let start = i;
222                i += 1;
223                let mut s = String::new();
224                loop {
225                    if i >= bytes.len() {
226                        return Err(LexError {
227                            message: "unterminated bracketed identifier".into(),
228                            offset: start,
229                        });
230                    }
231                    if bytes[i] == b']' {
232                        if i + 1 < bytes.len() && bytes[i + 1] == b']' {
233                            s.push(']');
234                            i += 2;
235                        } else {
236                            i += 1;
237                            break;
238                        }
239                    } else {
240                        let ch_len = utf8_len(bytes[i]);
241                        s.push_str(std::str::from_utf8(&bytes[i..i + ch_len]).map_err(|_| {
242                            LexError {
243                                message: "invalid UTF-8".into(),
244                                offset: i,
245                            }
246                        })?);
247                        i += ch_len;
248                    }
249                }
250                out.push_tok(Token::QuotedIdent(s), start);
251            }
252            b'0'..=b'9' => {
253                let (token, end) = scan_number(input, i)?;
254                out.push_tok(token, i);
255                i = end;
256            }
257            _ if b.is_ascii_alphabetic() || b == b'_' => {
258                let start = i;
259                while i < bytes.len() && (bytes[i].is_ascii_alphanumeric() || bytes[i] == b'_') {
260                    i += 1;
261                }
262                out.push_tok(Token::Ident(input[start..i].to_string()), start);
263            }
264            _ => {
265                let ch_len = utf8_len(b);
266                // allow unicode identifiers
267                let ch = input[i..].chars().next().unwrap();
268                if ch.is_alphabetic() {
269                    let start = i;
270                    i += ch_len;
271                    while i < bytes.len() {
272                        let c = input[i..].chars().next().unwrap();
273                        if c.is_alphanumeric() || c == '_' {
274                            i += c.len_utf8();
275                        } else {
276                            break;
277                        }
278                    }
279                    out.push_tok(Token::Ident(input[start..i].to_string()), start);
280                } else {
281                    return Err(LexError {
282                        message: format!("unexpected character {:?}", ch),
283                        offset: i,
284                    });
285                }
286            }
287        }
288    }
289    out.push_tok(Token::Eof, bytes.len());
290    Ok(out)
291}
292
293fn scan_number(input: &str, start: usize) -> Result<(Token, usize), LexError> {
294    let bytes = input.as_bytes();
295    let mut i = start;
296    let mut is_real = false;
297
298    if bytes[i] == b'.' {
299        is_real = true;
300        i += 1;
301        while i < bytes.len() && bytes[i].is_ascii_digit() {
302            i += 1;
303        }
304    } else {
305        while i < bytes.len() && bytes[i].is_ascii_digit() {
306            i += 1;
307        }
308        if bytes.get(i) == Some(&b'.') {
309            is_real = true;
310            i += 1;
311            while i < bytes.len() && bytes[i].is_ascii_digit() {
312                i += 1;
313            }
314        }
315    }
316
317    if bytes
318        .get(i)
319        .is_some_and(|byte| *byte == b'e' || *byte == b'E')
320    {
321        is_real = true;
322        i += 1;
323        if bytes
324            .get(i)
325            .is_some_and(|byte| *byte == b'+' || *byte == b'-')
326        {
327            i += 1;
328        }
329        let exponent_start = i;
330        while i < bytes.len() && bytes[i].is_ascii_digit() {
331            i += 1;
332        }
333        if i == exponent_start {
334            return Err(LexError {
335                message: "malformed number exponent".into(),
336                offset: start,
337            });
338        }
339    }
340
341    let identifier_follows = input
342        .get(i..)
343        .and_then(|remaining| remaining.chars().next())
344        .is_some_and(|character| character.is_alphanumeric() || character == '_');
345    if bytes.get(i) == Some(&b'.') || identifier_follows {
346        return Err(LexError {
347            message: "malformed number".into(),
348            offset: start,
349        });
350    }
351
352    let text = &input[start..i];
353    if is_real {
354        let f: f64 = text.parse().map_err(|_| LexError {
355            message: "malformed number".into(),
356            offset: start,
357        })?;
358        if !f.is_finite() {
359            return Err(LexError {
360                message: "real number out of range".into(),
361                offset: start,
362            });
363        }
364        Ok((Token::Real(f), i))
365    } else {
366        let n: i128 = text.parse().map_err(|_| LexError {
367            message: "integer out of range".into(),
368            offset: start,
369        })?;
370        Ok((Token::Integer(n), i))
371    }
372}
373
374trait PushToken {
375    fn push_tok(&mut self, t: Token, off: usize);
376}
377impl PushToken for Vec<TokenSpan> {
378    fn push_tok(&mut self, t: Token, off: usize) {
379        self.push(TokenSpan {
380            token: t,
381            offset: off,
382        });
383    }
384}
385
386fn utf8_len(b: u8) -> usize {
387    if b < 0x80 {
388        1
389    } else if b >> 5 == 0b110 {
390        2
391    } else if b >> 4 == 0b1110 {
392        3
393    } else {
394        4
395    }
396}
397
398#[cfg(test)]
399mod tests {
400    use super::*;
401
402    #[test]
403    fn escaped_quoted_identifiers() {
404        let tokens = lex("SELECT \"a\"\"b\", [c]]d] FROM t").unwrap();
405        assert!(
406            tokens
407                .iter()
408                .any(|span| span.token == Token::QuotedIdent("a\"b".into()))
409        );
410        assert!(
411            tokens
412                .iter()
413                .any(|span| span.token == Token::QuotedIdent("c]d".into()))
414        );
415    }
416
417    #[test]
418    fn rejects_truncated_escaped_quoted_identifier() {
419        assert!(lex("\"\"\"").is_err());
420    }
421
422    #[test]
423    fn scans_decimal_and_exponent_literals() {
424        let tokens = lex("SELECT .5, 1., 1e3, 1.25E-2").unwrap();
425        assert!(tokens.iter().any(|span| span.token == Token::Real(0.5)));
426        assert!(tokens.iter().any(|span| span.token == Token::Real(1.0)));
427        assert!(tokens.iter().any(|span| span.token == Token::Real(1000.0)));
428        assert!(tokens.iter().any(|span| span.token == Token::Real(0.0125)));
429    }
430
431    #[test]
432    fn rejects_malformed_numbers() {
433        assert!(lex("1.2.3").is_err());
434        assert!(lex("1e").is_err());
435        assert!(lex("1foo").is_err());
436        assert!(lex("1e309").is_err());
437    }
438}