Skip to main content

squawk_parser/
lexed_str.rs

1// based on https://github.com/rust-lang/rust-analyzer/blob/d8887c0758bbd2d5f752d5bd405d4491e90e7ed6/crates/parser/src/lexed_str.rs
2
3use std::{num::IntErrorKind, ops};
4
5use squawk_lexer::tokenize;
6
7use crate::SyntaxKind;
8
9pub struct LexedStr<'a> {
10    text: &'a str,
11    kind: Vec<SyntaxKind>,
12    start: Vec<u32>,
13    error: Vec<LexError>,
14}
15
16struct LexError {
17    msg: String,
18    range: ops::Range<u32>,
19}
20
21impl<'a> LexedStr<'a> {
22    // TODO: rust-analyzer has an edition thing to specify things that are only
23    // available in certain version, we can do that later
24    pub fn new(text: &'a str) -> LexedStr<'a> {
25        let mut conv = Converter::new(text);
26
27        for token in tokenize(&text[conv.offset..]) {
28            let token_text = &text[conv.offset..][..token.len as usize];
29
30            conv.extend_token(&token.kind, token_text);
31        }
32
33        conv.finalize_with_eof()
34    }
35
36    // pub(crate) fn single_token(text: &'a str) -> Option<(SyntaxKind, Option<String>)> {
37    //     if text.is_empty() {
38    //         return None;
39    //     }
40
41    //     let token = tokenize(text).next()?;
42    //     if token.len as usize != text.len() {
43    //         return None;
44    //     }
45
46    //     let mut conv = Converter::new(text);
47    //     conv.extend_token(&token.kind, text);
48    //     match &*conv.res.kind {
49    //         [kind] => Some((*kind, conv.res.error.pop().map(|it| it.msg))),
50    //         _ => None,
51    //     }
52    // }
53
54    // pub(crate) fn as_str(&self) -> &str {
55    //     self.text
56    // }
57
58    pub(crate) fn len(&self) -> usize {
59        self.kind.len() - 1
60    }
61
62    // pub(crate) fn is_empty(&self) -> bool {
63    //     self.len() == 0
64    // }
65
66    pub(crate) fn kind(&self, i: usize) -> SyntaxKind {
67        assert!(i < self.len());
68        self.kind[i]
69    }
70
71    pub(crate) fn range_text(&self, r: ops::Range<usize>) -> &str {
72        assert!(r.start < r.end && r.end <= self.len());
73        let lo = self.start[r.start] as usize;
74        let hi = self.start[r.end] as usize;
75        &self.text[lo..hi]
76    }
77
78    // Naming is hard.
79    pub fn text_range(&self, i: usize) -> ops::Range<usize> {
80        assert!(i < self.len());
81        let lo = self.start[i] as usize;
82        let hi = self.start[i + 1] as usize;
83        lo..hi
84    }
85    pub fn text_start(&self, i: usize) -> usize {
86        assert!(i <= self.len());
87        self.start[i] as usize
88    }
89    // pub(crate) fn text_len(&self, i: usize) -> usize {
90    //     assert!(i < self.len());
91    //     let r = self.text_range(i);
92    //     r.end - r.start
93    // }
94
95    // pub(crate) fn error(&self, i: usize) -> Option<&str> {
96    //     assert!(i < self.len());
97    //     let err = self
98    //         .error
99    //         .binary_search_by_key(&(i as u32), |i| i.token)
100    //         .ok()?;
101    //     Some(self.error[err].msg.as_str())
102    // }
103
104    pub fn errors(&self) -> impl Iterator<Item = (&ops::Range<u32>, &str)> + '_ {
105        self.error.iter().map(|it| (&it.range, it.msg.as_str()))
106    }
107
108    fn push(&mut self, kind: SyntaxKind, offset: usize) {
109        self.kind.push(kind);
110        self.start.push(offset as u32);
111    }
112}
113
114struct Converter<'a> {
115    res: LexedStr<'a>,
116    offset: usize,
117}
118
119fn is_empty_quoted_ident(token_text: &str, uescape: bool) -> bool {
120    let inner = if uescape {
121        token_text
122            .strip_prefix(['u', 'U'])
123            .and_then(|s| s.strip_prefix('&'))
124    } else {
125        Some(token_text)
126    };
127    inner == Some("\"\"")
128}
129
130impl<'a> Converter<'a> {
131    fn new(text: &'a str) -> Self {
132        Self {
133            res: LexedStr {
134                text,
135                kind: Vec::new(),
136                start: Vec::new(),
137                error: Vec::new(),
138            },
139            offset: 0,
140        }
141    }
142
143    fn finalize_with_eof(mut self) -> LexedStr<'a> {
144        self.res.push(SyntaxKind::EOF, self.offset);
145        self.res
146    }
147
148    fn push(&mut self, kind: SyntaxKind, len: usize, err: Option<(&str, ops::Range<u32>)>) {
149        let token_start = self.offset as u32;
150        self.res.push(kind, self.offset);
151        self.offset += len;
152
153        if let Some((msg, err_range)) = err {
154            self.res.error.push(LexError {
155                msg: msg.to_owned(),
156                range: token_start + err_range.start..token_start + err_range.end,
157            });
158        }
159    }
160
161    fn extend_token(&mut self, kind: &squawk_lexer::TokenKind, token_text: &str) {
162        // A note on an intended tradeoff:
163        // We drop some useful information here (see patterns with double dots `..`)
164        // Storing that info in `SyntaxKind` is not possible due to its layout requirements of
165        // being `u16` that come from `rowan::SyntaxKind`.
166        let mut err = "";
167        let mut err_range: Option<ops::Range<u32>> = None;
168
169        let syntax_kind = {
170            match kind {
171                squawk_lexer::TokenKind::LineComment => SyntaxKind::COMMENT,
172                squawk_lexer::TokenKind::BlockComment { terminated } => {
173                    if !terminated {
174                        err = "Missing trailing `*/` symbols to terminate the block comment";
175                    }
176                    SyntaxKind::COMMENT
177                }
178
179                squawk_lexer::TokenKind::Whitespace => SyntaxKind::WHITESPACE,
180                squawk_lexer::TokenKind::Ident => {
181                    SyntaxKind::from_keyword(token_text).unwrap_or(SyntaxKind::IDENT)
182                }
183                squawk_lexer::TokenKind::Literal { kind, .. } => {
184                    self.extend_literal(token_text, kind);
185                    return;
186                }
187                squawk_lexer::TokenKind::Semi => SyntaxKind::SEMICOLON,
188                squawk_lexer::TokenKind::Comma => SyntaxKind::COMMA,
189                squawk_lexer::TokenKind::Dot => SyntaxKind::DOT,
190                squawk_lexer::TokenKind::OpenParen => SyntaxKind::L_PAREN,
191                squawk_lexer::TokenKind::CloseParen => SyntaxKind::R_PAREN,
192                squawk_lexer::TokenKind::OpenBracket => SyntaxKind::L_BRACK,
193                squawk_lexer::TokenKind::CloseBracket => SyntaxKind::R_BRACK,
194                squawk_lexer::TokenKind::OpenCurly => SyntaxKind::L_CURLY,
195                squawk_lexer::TokenKind::CloseCurly => SyntaxKind::R_CURLY,
196                squawk_lexer::TokenKind::At => SyntaxKind::AT,
197                squawk_lexer::TokenKind::Pound => SyntaxKind::POUND,
198                squawk_lexer::TokenKind::Tilde => SyntaxKind::TILDE,
199                squawk_lexer::TokenKind::Question => SyntaxKind::QUESTION,
200                squawk_lexer::TokenKind::Colon => SyntaxKind::COLON,
201                squawk_lexer::TokenKind::Eq => SyntaxKind::EQ,
202                squawk_lexer::TokenKind::Bang => SyntaxKind::BANG,
203                squawk_lexer::TokenKind::Lt => SyntaxKind::L_ANGLE,
204                squawk_lexer::TokenKind::Gt => SyntaxKind::R_ANGLE,
205                squawk_lexer::TokenKind::Minus => SyntaxKind::MINUS,
206                squawk_lexer::TokenKind::And => SyntaxKind::AMP,
207                squawk_lexer::TokenKind::Or => SyntaxKind::PIPE,
208                squawk_lexer::TokenKind::Plus => SyntaxKind::PLUS,
209                squawk_lexer::TokenKind::Star => SyntaxKind::STAR,
210                squawk_lexer::TokenKind::Slash => SyntaxKind::SLASH,
211                squawk_lexer::TokenKind::Caret => SyntaxKind::CARET,
212                squawk_lexer::TokenKind::Percent => SyntaxKind::PERCENT,
213                squawk_lexer::TokenKind::Unknown => SyntaxKind::ERROR,
214                squawk_lexer::TokenKind::Eof => SyntaxKind::EOF,
215                squawk_lexer::TokenKind::Backtick => SyntaxKind::BACKTICK,
216                squawk_lexer::TokenKind::PositionalParam {
217                    trailing_junk_start,
218                } => {
219                    let digits = &token_text[1..*trailing_junk_start as usize];
220                    if digits.is_empty() {
221                        err = "missing parameter number";
222                        err_range = Some(0..1);
223                    } else if digits
224                        .parse::<i32>()
225                        .is_err_and(|err| matches!(err.kind(), IntErrorKind::PosOverflow))
226                    {
227                        err = "parameter number too large";
228                        err_range = Some(0..*trailing_junk_start);
229                    } else if (*trailing_junk_start as usize) < token_text.len() {
230                        err = "trailing junk after positional parameter";
231                        err_range = Some(*trailing_junk_start..token_text.len() as u32);
232                    }
233                    SyntaxKind::POSITIONAL_PARAM
234                }
235                squawk_lexer::TokenKind::QuotedIdent {
236                    terminated,
237                    uescape,
238                } => {
239                    if !terminated {
240                        err = "Missing trailing \" to terminate the quoted identifier"
241                    } else if is_empty_quoted_ident(token_text, *uescape) {
242                        err = "empty delimited identifier";
243                    }
244                    SyntaxKind::IDENT
245                }
246            }
247        };
248
249        let err = if err.is_empty() { None } else { Some(err) };
250        let err = err.map(|msg| (msg, err_range.unwrap_or(0..token_text.len() as u32)));
251        self.push(syntax_kind, token_text.len(), err);
252    }
253
254    fn extend_literal(&mut self, token_text: &str, kind: &squawk_lexer::LiteralKind) {
255        let mut err: Option<String> = None;
256        let mut err_range: Option<ops::Range<u32>> = None;
257
258        let syntax_kind = match *kind {
259            squawk_lexer::LiteralKind::Int {
260                empty_int,
261                base,
262                trailing_junk_start,
263            } => {
264                if empty_int {
265                    err = Some("Missing digits after the integer base prefix".into());
266                } else {
267                    if matches!(base, squawk_lexer::Base::Binary | squawk_lexer::Base::Octal) {
268                        let prefix_len = 2u32;
269                        let digits = &token_text[prefix_len as usize..trailing_junk_start as usize];
270                        let base = base as u32;
271                        let token_start = self.offset as u32;
272                        for (i, c) in digits.char_indices() {
273                            if c != '_' && c.to_digit(base).is_none() {
274                                let start = token_start + prefix_len + i as u32;
275                                let end = start + c.len_utf8() as u32;
276                                self.res.error.push(LexError {
277                                    msg: format!("invalid digit for a base {base} literal"),
278                                    range: start..end,
279                                });
280                            }
281                        }
282                    }
283                    if (trailing_junk_start as usize) < token_text.len() {
284                        err = Some("trailing junk after numeric literal".into());
285                        err_range = Some(trailing_junk_start..token_text.len() as u32);
286                    }
287                }
288                SyntaxKind::INT_NUMBER
289            }
290            squawk_lexer::LiteralKind::Numeric {
291                empty_exponent_start,
292                trailing_junk_start,
293            } => {
294                if let Some(exponent_start) = empty_exponent_start {
295                    err = Some("Missing digits after the exponent symbol".into());
296                    err_range = Some(exponent_start..exponent_start + 1);
297                } else if (trailing_junk_start as usize) < token_text.len() {
298                    err = Some("trailing junk after numeric literal".into());
299                    err_range = Some(trailing_junk_start..token_text.len() as u32);
300                }
301                SyntaxKind::NUMERIC_NUMBER
302            }
303            squawk_lexer::LiteralKind::Str { terminated } => {
304                if !terminated {
305                    err =
306                        Some("Missing trailing `'` symbol to terminate the string literal".into());
307                }
308                SyntaxKind::STRING
309            }
310            squawk_lexer::LiteralKind::NationalStr { terminated } => {
311                if !terminated {
312                    err = Some(
313                        "Missing trailing `'` symbol to terminate the national character string literal"
314                            .into(),
315                    );
316                }
317                SyntaxKind::NATIONAL_STRING
318            }
319            squawk_lexer::LiteralKind::ByteStr { terminated } => {
320                if !terminated {
321                    err = Some(
322                        "Missing trailing `'` symbol to terminate the hex bit string literal"
323                            .into(),
324                    );
325                }
326                // digit validation in squawk_syntax
327                SyntaxKind::BYTE_STRING
328            }
329            squawk_lexer::LiteralKind::BitStr { terminated } => {
330                if !terminated {
331                    err = Some(
332                        "Missing trailing `'` symbol to terminate the bit string literal".into(),
333                    );
334                }
335                // digit validation in squawk_syntax
336                SyntaxKind::BIT_STRING
337            }
338            squawk_lexer::LiteralKind::DollarQuotedString { terminated } => {
339                if !terminated {
340                    // TODO: we could be fancier and say the ending string we're looking for
341                    err = Some("Unterminated dollar quoted string literal".into());
342                }
343                SyntaxKind::DOLLAR_QUOTED_STRING
344            }
345            squawk_lexer::LiteralKind::UnicodeEscStr { terminated } => {
346                if !terminated {
347                    err = Some(
348                        "Missing trailing `'` symbol to terminate the unicode escape string literal"
349                            .into(),
350                    );
351                }
352                // validated in squawk_syntax
353                SyntaxKind::UNICODE_ESC_STRING
354            }
355            squawk_lexer::LiteralKind::EscStr { terminated } => {
356                if !terminated {
357                    err = Some(
358                        "Missing trailing `'` symbol to terminate the escape string literal".into(),
359                    );
360                }
361                // unicode escape sequences validated in squawk_syntax
362                SyntaxKind::ESC_STRING
363            }
364        };
365
366        let err = err
367            .as_deref()
368            .map(|msg| (msg, err_range.unwrap_or(0..token_text.len() as u32)));
369        self.push(syntax_kind, token_text.len(), err);
370    }
371}
372
373#[cfg(test)]
374mod tests {
375    use annotate_snippets::{AnnotationKind, Level, Renderer, Snippet, renderer::DecorStyle};
376    use insta::assert_snapshot;
377
378    use super::LexedStr;
379
380    fn lex(text: &str) -> String {
381        let lexed = LexedStr::new(text);
382        let renderer = Renderer::plain().decor_style(DecorStyle::Unicode);
383        let mut res = String::new();
384
385        for (range, msg) in lexed.errors() {
386            let span = range.start as usize..range.end as usize;
387            let group = Level::ERROR.primary_title(msg).element(
388                Snippet::source(text)
389                    .fold(true)
390                    .annotation(AnnotationKind::Primary.span(span)),
391            );
392            res.push_str(&renderer.render(&[group]).to_string());
393            res.push('\n');
394        }
395
396        res
397    }
398
399    #[test]
400    fn empty_int_error() {
401        assert_snapshot!(lex("select 0x;"), @"
402        error: Missing digits after the integer base prefix
403          ╭▸ 
404        1 │ select 0x;
405          ╰╴       ━━
406        ");
407    }
408
409    #[test]
410    fn empty_int_with_trailing_ident_error() {
411        assert_snapshot!(lex("select 0xg;"), @"
412        error: trailing junk after numeric literal
413          ╭▸ 
414        1 │ select 0xg;
415          ╰╴         ━
416        ");
417    }
418
419    #[test]
420    fn invalid_octal_digits_error() {
421        assert_snapshot!(lex("select 0o999;"), @"
422        error: invalid digit for a base 8 literal
423          ╭▸ 
424        1 │ select 0o999;
425          ╰╴         ━
426        error: invalid digit for a base 8 literal
427          ╭▸ 
428        1 │ select 0o999;
429          ╰╴          ━
430        error: invalid digit for a base 8 literal
431          ╭▸ 
432        1 │ select 0o999;
433          ╰╴           ━
434        ");
435    }
436
437    #[test]
438    fn invalid_binary_digits_error() {
439        assert_snapshot!(lex("select 0b234;"), @"
440        error: invalid digit for a base 2 literal
441          ╭▸ 
442        1 │ select 0b234;
443          ╰╴         ━
444        error: invalid digit for a base 2 literal
445          ╭▸ 
446        1 │ select 0b234;
447          ╰╴          ━
448        error: invalid digit for a base 2 literal
449          ╭▸ 
450        1 │ select 0b234;
451          ╰╴           ━
452        ");
453    }
454
455    #[test]
456    fn invalid_octal_digits_after_valid_error() {
457        assert_snapshot!(lex("select 0o7889;"), @"
458        error: invalid digit for a base 8 literal
459          ╭▸ 
460        1 │ select 0o7889;
461          ╰╴          ━
462        error: invalid digit for a base 8 literal
463          ╭▸ 
464        1 │ select 0o7889;
465          ╰╴           ━
466        error: invalid digit for a base 8 literal
467          ╭▸ 
468        1 │ select 0o7889;
469          ╰╴            ━
470        ");
471    }
472
473    #[test]
474    fn empty_exponent_error() {
475        assert_snapshot!(lex("select 1e;"), @"
476        error: Missing digits after the exponent symbol
477          ╭▸ 
478        1 │ select 1e;
479          ╰╴        ━
480        ");
481    }
482
483    #[test]
484    fn unterminated_string_error() {
485        assert_snapshot!(lex("select 'hello;"), @"
486        error: Missing trailing `'` symbol to terminate the string literal
487          ╭▸ 
488        1 │ select 'hello;
489          ╰╴       ━━━━━━━
490        ");
491    }
492
493    #[test]
494    fn unterminated_hex_bit_string_error() {
495        assert_snapshot!(lex("select X'1F;"), @"
496        error: Missing trailing `'` symbol to terminate the hex bit string literal
497          ╭▸ 
498        1 │ select X'1F;
499          ╰╴       ━━━━━
500        ");
501    }
502
503    #[test]
504    fn unterminated_bit_string_error() {
505        assert_snapshot!(lex("select B'101;"), @"
506        error: Missing trailing `'` symbol to terminate the bit string literal
507          ╭▸ 
508        1 │ select B'101;
509          ╰╴       ━━━━━━
510        ");
511    }
512
513    #[test]
514    fn unterminated_dollar_quoted_string_error() {
515        assert_snapshot!(lex("select $tag$hello;"), @"
516        error: Unterminated dollar quoted string literal
517          ╭▸ 
518        1 │ select $tag$hello;
519          ╰╴       ━━━━━━━━━━━
520        ");
521    }
522
523    #[test]
524    fn unterminated_unicode_escape_string_error() {
525        assert_snapshot!(lex("select U&'hello;"), @"
526        error: Missing trailing `'` symbol to terminate the unicode escape string literal
527          ╭▸ 
528        1 │ select U&'hello;
529          ╰╴       ━━━━━━━━━
530        ");
531    }
532
533    #[test]
534    fn unterminated_escape_string_error() {
535        assert_snapshot!(lex("select E'hello;"), @"
536        error: Missing trailing `'` symbol to terminate the escape string literal
537          ╭▸ 
538        1 │ select E'hello;
539          ╰╴       ━━━━━━━━
540        ");
541    }
542}