1use 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 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 len(&self) -> usize {
59 self.kind.len() - 1
60 }
61
62 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 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 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 prefixed_string_continuation: Option<PrefixedStringKind>,
118}
119
120#[derive(Clone, Copy)]
121enum PrefixedStringKind {
122 Bit,
123 Byte,
124 Escape,
125}
126
127fn is_empty_quoted_ident(token_text: &str, uescape: bool) -> bool {
128 let inner = if uescape {
129 token_text
130 .strip_prefix(['u', 'U'])
131 .and_then(|s| s.strip_prefix('&'))
132 } else {
133 Some(token_text)
134 };
135 inner == Some("\"\"")
136}
137
138impl<'a> Converter<'a> {
139 fn new(text: &'a str) -> Self {
140 Self {
141 res: LexedStr {
142 text,
143 kind: Vec::new(),
144 start: Vec::new(),
145 error: Vec::new(),
146 },
147 offset: 0,
148 prefixed_string_continuation: None,
149 }
150 }
151
152 fn finalize_with_eof(mut self) -> LexedStr<'a> {
153 self.res.push(SyntaxKind::EOF, self.offset);
154 self.res
155 }
156
157 fn push(&mut self, kind: SyntaxKind, len: usize, err: Option<(&str, ops::Range<u32>)>) {
158 let token_start = self.offset as u32;
159 self.res.push(kind, self.offset);
160 self.offset += len;
161
162 if let Some((msg, err_range)) = err {
163 self.res.error.push(LexError {
164 msg: msg.to_owned(),
165 range: token_start + err_range.start..token_start + err_range.end,
166 });
167 }
168 }
169
170 fn extend_token(&mut self, kind: &squawk_lexer::TokenKind, token_text: &str) {
171 if !matches!(
172 kind,
173 squawk_lexer::TokenKind::Whitespace
174 | squawk_lexer::TokenKind::LineComment
175 | squawk_lexer::TokenKind::BlockComment { .. }
176 | squawk_lexer::TokenKind::Literal { .. }
177 ) {
178 self.prefixed_string_continuation = None;
179 }
180
181 let mut err = "";
186 let mut err_range: Option<ops::Range<u32>> = None;
187
188 let syntax_kind = {
189 match kind {
190 squawk_lexer::TokenKind::LineComment => SyntaxKind::COMMENT,
191 squawk_lexer::TokenKind::BlockComment { terminated } => {
192 if !terminated {
193 err = "Missing trailing `*/` symbols to terminate the block comment";
194 }
195 SyntaxKind::COMMENT
196 }
197
198 squawk_lexer::TokenKind::Whitespace => SyntaxKind::WHITESPACE,
199 squawk_lexer::TokenKind::Ident => {
200 SyntaxKind::from_keyword(token_text).unwrap_or(SyntaxKind::IDENT)
201 }
202 squawk_lexer::TokenKind::Literal { kind, .. } => {
203 self.extend_literal(token_text, kind);
204 return;
205 }
206 squawk_lexer::TokenKind::Semi => SyntaxKind::SEMICOLON,
207 squawk_lexer::TokenKind::Comma => SyntaxKind::COMMA,
208 squawk_lexer::TokenKind::Dot => SyntaxKind::DOT,
209 squawk_lexer::TokenKind::OpenParen => SyntaxKind::L_PAREN,
210 squawk_lexer::TokenKind::CloseParen => SyntaxKind::R_PAREN,
211 squawk_lexer::TokenKind::OpenBracket => SyntaxKind::L_BRACK,
212 squawk_lexer::TokenKind::CloseBracket => SyntaxKind::R_BRACK,
213 squawk_lexer::TokenKind::OpenCurly => SyntaxKind::L_CURLY,
214 squawk_lexer::TokenKind::CloseCurly => SyntaxKind::R_CURLY,
215 squawk_lexer::TokenKind::At => SyntaxKind::AT,
216 squawk_lexer::TokenKind::Pound => SyntaxKind::POUND,
217 squawk_lexer::TokenKind::Tilde => SyntaxKind::TILDE,
218 squawk_lexer::TokenKind::Question => SyntaxKind::QUESTION,
219 squawk_lexer::TokenKind::Colon => SyntaxKind::COLON,
220 squawk_lexer::TokenKind::Eq => SyntaxKind::EQ,
221 squawk_lexer::TokenKind::Bang => SyntaxKind::BANG,
222 squawk_lexer::TokenKind::Lt => SyntaxKind::L_ANGLE,
223 squawk_lexer::TokenKind::Gt => SyntaxKind::R_ANGLE,
224 squawk_lexer::TokenKind::Minus => SyntaxKind::MINUS,
225 squawk_lexer::TokenKind::And => SyntaxKind::AMP,
226 squawk_lexer::TokenKind::Or => SyntaxKind::PIPE,
227 squawk_lexer::TokenKind::Plus => SyntaxKind::PLUS,
228 squawk_lexer::TokenKind::Star => SyntaxKind::STAR,
229 squawk_lexer::TokenKind::Slash => SyntaxKind::SLASH,
230 squawk_lexer::TokenKind::Caret => SyntaxKind::CARET,
231 squawk_lexer::TokenKind::Percent => SyntaxKind::PERCENT,
232 squawk_lexer::TokenKind::Unknown => SyntaxKind::ERROR,
233 squawk_lexer::TokenKind::Eof => SyntaxKind::EOF,
234 squawk_lexer::TokenKind::Backtick => SyntaxKind::BACKTICK,
235 squawk_lexer::TokenKind::PositionalParam {
236 trailing_junk_start,
237 } => {
238 let digits = &token_text[1..*trailing_junk_start as usize];
239 if digits.is_empty() {
240 err = "missing parameter number";
241 err_range = Some(0..1);
242 } else if digits
243 .parse::<i32>()
244 .is_err_and(|err| matches!(err.kind(), IntErrorKind::PosOverflow))
245 {
246 err = "parameter number too large";
247 err_range = Some(0..*trailing_junk_start);
248 } else if (*trailing_junk_start as usize) < token_text.len() {
249 err = "trailing junk after positional parameter";
250 err_range = Some(*trailing_junk_start..token_text.len() as u32);
251 }
252 SyntaxKind::POSITIONAL_PARAM
253 }
254 squawk_lexer::TokenKind::QuotedIdent {
255 terminated,
256 uescape,
257 } => {
258 if !terminated {
259 err = "Missing trailing \" to terminate the quoted identifier"
260 } else if is_empty_quoted_ident(token_text, *uescape) {
261 err = "empty delimited identifier";
262 }
263 SyntaxKind::IDENT
264 }
265 }
266 };
267
268 let err = if err.is_empty() { None } else { Some(err) };
269 let err = err.map(|msg| (msg, err_range.unwrap_or(0..token_text.len() as u32)));
270 self.push(syntax_kind, token_text.len(), err);
271 }
272
273 fn extend_literal(&mut self, token_text: &str, kind: &squawk_lexer::LiteralKind) {
274 let mut err: Option<String> = None;
275 let mut err_range: Option<ops::Range<u32>> = None;
276 let continuation = self.prefixed_string_continuation.take();
277
278 let syntax_kind = match *kind {
279 squawk_lexer::LiteralKind::Int {
280 empty_int,
281 base,
282 trailing_junk_start,
283 } => {
284 if empty_int {
285 err = Some("Missing digits after the integer base prefix".into());
286 } else {
287 if matches!(base, squawk_lexer::Base::Binary | squawk_lexer::Base::Octal) {
288 let prefix_len = 2u32;
289 let digits = &token_text[prefix_len as usize..trailing_junk_start as usize];
290 let base = base as u32;
291 let token_start = self.offset as u32;
292 for (i, c) in digits.char_indices() {
293 if c != '_' && c.to_digit(base).is_none() {
294 let start = token_start + prefix_len + i as u32;
295 let end = start + c.len_utf8() as u32;
296 self.res.error.push(LexError {
297 msg: format!("invalid digit for a base {base} literal"),
298 range: start..end,
299 });
300 }
301 }
302 }
303 if (trailing_junk_start as usize) < token_text.len() {
304 err = Some("trailing junk after numeric literal".into());
305 err_range = Some(trailing_junk_start..token_text.len() as u32);
306 }
307 }
308 SyntaxKind::INT_NUMBER
309 }
310 squawk_lexer::LiteralKind::Numeric {
311 empty_exponent_start,
312 trailing_junk_start,
313 } => {
314 if let Some(exponent_start) = empty_exponent_start {
315 err = Some("Missing digits after the exponent symbol".into());
316 err_range = Some(exponent_start..exponent_start + 1);
317 } else if (trailing_junk_start as usize) < token_text.len() {
318 err = Some("trailing junk after numeric literal".into());
319 err_range = Some(trailing_junk_start..token_text.len() as u32);
320 }
321 SyntaxKind::NUMERIC_NUMBER
322 }
323 squawk_lexer::LiteralKind::Str { terminated } => {
324 if !terminated {
325 err =
326 Some("Missing trailing `'` symbol to terminate the string literal".into());
327 } else if let Some(kind) = continuation {
328 self.validate_prefixed_string_content(token_text, 1, kind);
329 self.prefixed_string_continuation = Some(kind);
330 }
331 SyntaxKind::STRING
332 }
333 squawk_lexer::LiteralKind::NationalStr { terminated } => {
334 if !terminated {
335 err = Some(
336 "Missing trailing `'` symbol to terminate the national character string literal"
337 .into(),
338 );
339 }
340 SyntaxKind::NATIONAL_STRING
341 }
342 squawk_lexer::LiteralKind::ByteStr { terminated } => {
343 if !terminated {
344 err = Some(
345 "Missing trailing `'` symbol to terminate the hex bit string literal"
346 .into(),
347 );
348 } else {
349 self.validate_prefixed_string_content(token_text, 2, PrefixedStringKind::Byte);
350 self.prefixed_string_continuation = Some(PrefixedStringKind::Byte);
351 }
352 SyntaxKind::BYTE_STRING
353 }
354 squawk_lexer::LiteralKind::BitStr { terminated } => {
355 if !terminated {
356 err = Some(
357 "Missing trailing `'` symbol to terminate the bit string literal".into(),
358 );
359 } else {
360 self.validate_prefixed_string_content(token_text, 2, PrefixedStringKind::Bit);
361 self.prefixed_string_continuation = Some(PrefixedStringKind::Bit);
362 }
363 SyntaxKind::BIT_STRING
364 }
365 squawk_lexer::LiteralKind::DollarQuotedString { terminated } => {
366 if !terminated {
367 err = Some("Unterminated dollar quoted string literal".into());
369 }
370 SyntaxKind::DOLLAR_QUOTED_STRING
371 }
372 squawk_lexer::LiteralKind::UnicodeEscStr { terminated } => {
373 if !terminated {
374 err = Some(
375 "Missing trailing `'` symbol to terminate the unicode escape string literal"
376 .into(),
377 );
378 }
379 SyntaxKind::UNICODE_ESC_STRING
381 }
382 squawk_lexer::LiteralKind::EscStr { terminated } => {
383 if !terminated {
384 err = Some(
385 "Missing trailing `'` symbol to terminate the escape string literal".into(),
386 );
387 } else {
388 self.validate_prefixed_string_content(
389 token_text,
390 2,
391 PrefixedStringKind::Escape,
392 );
393 self.prefixed_string_continuation = Some(PrefixedStringKind::Escape);
394 }
395 SyntaxKind::ESC_STRING
396 }
397 };
398
399 let err = err
400 .as_deref()
401 .map(|msg| (msg, err_range.unwrap_or(0..token_text.len() as u32)));
402 self.push(syntax_kind, token_text.len(), err);
403 }
404
405 fn validate_prefixed_string_content(
406 &mut self,
407 token_text: &str,
408 inner_start: usize,
409 kind: PrefixedStringKind,
410 ) {
411 let inner = &token_text[inner_start..token_text.len() - 1];
412 match kind {
413 PrefixedStringKind::Bit => {
414 for (i, c) in inner.char_indices() {
415 if !matches!(c, '0' | '1') {
416 self.push_content_error(
417 format!(r#""{c}" is not a valid binary digit"#),
418 inner_start + i,
419 c.len_utf8(),
420 );
421 }
422 }
423 }
424 PrefixedStringKind::Byte => {
425 for (i, c) in inner.char_indices() {
426 if !c.is_ascii_hexdigit() {
427 self.push_content_error(
428 format!(r#""{c}" is not a valid hexadecimal digit"#),
429 inner_start + i,
430 c.len_utf8(),
431 );
432 }
433 }
434 }
435 PrefixedStringKind::Escape => {
436 let mut chars = inner.char_indices().peekable();
437 while let Some((escape_start, c)) = chars.next() {
438 if c != '\\' {
439 continue;
440 }
441 let Some((next_pos, next_c)) = chars.next() else {
442 break;
443 };
444 let (required, example) = match next_c {
445 'u' => (4usize, r"\uXXXX"),
446 'U' => (8usize, r"\UXXXXXXXX"),
447 _ => continue,
448 };
449 let mut end = next_pos + next_c.len_utf8();
450 let mut got_all = true;
451 for _ in 0..required {
452 match chars.peek() {
453 Some(&(i, ch)) if ch.is_ascii_hexdigit() => {
454 end = i + ch.len_utf8();
455 chars.next();
456 }
457 _ => {
458 got_all = false;
459 break;
460 }
461 }
462 }
463 if !got_all {
464 self.push_content_error(
465 format!("Unicode escape requires {required} hex digits: {example}"),
466 inner_start + escape_start,
467 end - escape_start,
468 );
469 }
470 }
471 }
472 }
473 }
474
475 fn push_content_error(&mut self, msg: String, start: usize, len: usize) {
476 let token_start = self.offset as u32;
477 let start = token_start + start as u32;
478 self.res.error.push(LexError {
479 msg,
480 range: start..start + len as u32,
481 });
482 }
483}
484
485#[cfg(test)]
486mod tests {
487 use annotate_snippets::{AnnotationKind, Level, Renderer, Snippet, renderer::DecorStyle};
488 use insta::{assert_debug_snapshot, assert_snapshot};
489
490 use super::LexedStr;
491
492 fn lex(text: &str) -> String {
493 let lexed = LexedStr::new(text);
494 let renderer = Renderer::plain().decor_style(DecorStyle::Unicode);
495 let mut res = String::new();
496
497 for (range, msg) in lexed.errors() {
498 let span = range.start as usize..range.end as usize;
499 let group = Level::ERROR.primary_title(msg).element(
500 Snippet::source(text)
501 .fold(true)
502 .annotation(AnnotationKind::Primary.span(span)),
503 );
504 res.push_str(&renderer.render(&[group]).to_string());
505 res.push('\n');
506 }
507
508 res
509 }
510
511 fn lex_errors(text: &str) -> Vec<(std::ops::Range<u32>, String)> {
512 LexedStr::new(text)
513 .errors()
514 .map(|(range, msg)| (range.clone(), msg.to_owned()))
515 .collect()
516 }
517
518 #[test]
519 fn prefixed_string_content_errors() {
520 assert_debug_snapshot!(lex_errors("B'102' X'1G' E'\\u00'"), @r#"
521 [
522 (
523 4..5,
524 "\"2\" is not a valid binary digit",
525 ),
526 (
527 10..11,
528 "\"G\" is not a valid hexadecimal digit",
529 ),
530 (
531 15..19,
532 "Unicode escape requires 4 hex digits: \\uXXXX",
533 ),
534 ]
535 "#);
536 }
537
538 #[test]
539 fn prefixed_string_continuations_use_the_initial_string_kind() {
540 assert_debug_snapshot!(lex_errors("B'0'\n'2'"), @r#"
541 [
542 (
543 6..7,
544 "\"2\" is not a valid binary digit",
545 ),
546 ]
547 "#);
548 assert_debug_snapshot!(lex_errors("X'F'\n'G'"), @r#"
549 [
550 (
551 6..7,
552 "\"G\" is not a valid hexadecimal digit",
553 ),
554 ]
555 "#);
556 assert_debug_snapshot!(lex_errors("E'ok'\n'\\u0'"), @r#"
557 [
558 (
559 7..10,
560 "Unicode escape requires 4 hex digits: \\uXXXX",
561 ),
562 ]
563 "#);
564 }
565
566 #[test]
567 fn prefixed_string_continuation_state_resets() {
568 assert_debug_snapshot!(
569 lex_errors("B'01' || '2'; X'0F' N'x' 'G'; E'ok' $tag$x$tag$ '\\u0'"),
570 @"[]"
571 );
572 assert_debug_snapshot!(lex_errors("B'01' 2 '2'; X'0F' 1.5 'G'"), @"[]");
573 }
574
575 #[test]
576 fn empty_int_error() {
577 assert_snapshot!(lex("select 0x;"), @"
578 error: Missing digits after the integer base prefix
579 ╭▸
580 1 │ select 0x;
581 ╰╴ ━━
582 ");
583 }
584
585 #[test]
586 fn empty_int_with_trailing_ident_error() {
587 assert_snapshot!(lex("select 0xg;"), @"
588 error: trailing junk after numeric literal
589 ╭▸
590 1 │ select 0xg;
591 ╰╴ ━
592 ");
593 }
594
595 #[test]
596 fn invalid_octal_digits_error() {
597 assert_snapshot!(lex("select 0o999;"), @"
598 error: invalid digit for a base 8 literal
599 ╭▸
600 1 │ select 0o999;
601 ╰╴ ━
602 error: invalid digit for a base 8 literal
603 ╭▸
604 1 │ select 0o999;
605 ╰╴ ━
606 error: invalid digit for a base 8 literal
607 ╭▸
608 1 │ select 0o999;
609 ╰╴ ━
610 ");
611 }
612
613 #[test]
614 fn invalid_binary_digits_error() {
615 assert_snapshot!(lex("select 0b234;"), @"
616 error: invalid digit for a base 2 literal
617 ╭▸
618 1 │ select 0b234;
619 ╰╴ ━
620 error: invalid digit for a base 2 literal
621 ╭▸
622 1 │ select 0b234;
623 ╰╴ ━
624 error: invalid digit for a base 2 literal
625 ╭▸
626 1 │ select 0b234;
627 ╰╴ ━
628 ");
629 }
630
631 #[test]
632 fn invalid_octal_digits_after_valid_error() {
633 assert_snapshot!(lex("select 0o7889;"), @"
634 error: invalid digit for a base 8 literal
635 ╭▸
636 1 │ select 0o7889;
637 ╰╴ ━
638 error: invalid digit for a base 8 literal
639 ╭▸
640 1 │ select 0o7889;
641 ╰╴ ━
642 error: invalid digit for a base 8 literal
643 ╭▸
644 1 │ select 0o7889;
645 ╰╴ ━
646 ");
647 }
648
649 #[test]
650 fn empty_exponent_error() {
651 assert_snapshot!(lex("select 1e;"), @"
652 error: Missing digits after the exponent symbol
653 ╭▸
654 1 │ select 1e;
655 ╰╴ ━
656 ");
657 }
658
659 #[test]
660 fn unterminated_string_error() {
661 assert_snapshot!(lex("select 'hello;"), @"
662 error: Missing trailing `'` symbol to terminate the string literal
663 ╭▸
664 1 │ select 'hello;
665 ╰╴ ━━━━━━━
666 ");
667 }
668
669 #[test]
670 fn unterminated_hex_bit_string_error() {
671 assert_snapshot!(lex("select X'1F;"), @"
672 error: Missing trailing `'` symbol to terminate the hex bit string literal
673 ╭▸
674 1 │ select X'1F;
675 ╰╴ ━━━━━
676 ");
677 }
678
679 #[test]
680 fn unterminated_bit_string_error() {
681 assert_snapshot!(lex("select B'101;"), @"
682 error: Missing trailing `'` symbol to terminate the bit string literal
683 ╭▸
684 1 │ select B'101;
685 ╰╴ ━━━━━━
686 ");
687 }
688
689 #[test]
690 fn unterminated_dollar_quoted_string_error() {
691 assert_snapshot!(lex("select $tag$hello;"), @"
692 error: Unterminated dollar quoted string literal
693 ╭▸
694 1 │ select $tag$hello;
695 ╰╴ ━━━━━━━━━━━
696 ");
697 }
698
699 #[test]
700 fn unterminated_unicode_escape_string_error() {
701 assert_snapshot!(lex("select U&'hello;"), @"
702 error: Missing trailing `'` symbol to terminate the unicode escape string literal
703 ╭▸
704 1 │ select U&'hello;
705 ╰╴ ━━━━━━━━━
706 ");
707 }
708
709 #[test]
710 fn unterminated_escape_string_error() {
711 assert_snapshot!(lex("select E'hello;"), @"
712 error: Missing trailing `'` symbol to terminate the escape string literal
713 ╭▸
714 1 │ select E'hello;
715 ╰╴ ━━━━━━━━
716 ");
717 }
718}