Skip to main content

inillucent_sql/
lexer.rs

1//! The zero-copy lexer: SQL bytes in, tokens with spans out.
2//!
3//! Invariant: every input byte belongs to exactly one token or trivia span,
4//! spans are ordered, non-overlapping and inside the source, and no token owns
5//! a byte. Token text is always a slice of the original SQL, which is what
6//! makes `prepare` free of identifier allocation and what lets an error point
7//! at an exact offset in the caller's own string.
8//!
9//! An unterminated quote or comment reports the offset of the byte that opened
10//! it, not the end of input. That is the difference between "there is a problem
11//! at character 4093" and "there is a problem somewhere", and it is the reason
12//! the opening offset is carried down rather than recomputed.
13
14use crate::keyword::{self, Keyword};
15
16/// A half-open byte range in the source SQL.
17#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash, Default)]
18pub struct Span {
19    /// The first byte of the span.
20    pub start: u32,
21    /// One past the last byte of the span.
22    pub end: u32,
23}
24
25impl Span {
26    /// Returns a span covering `start..end`.
27    pub fn new(start: usize, end: usize) -> Span {
28        Span {
29            start: start.min(u32::MAX as usize) as u32,
30            end: end.min(u32::MAX as usize) as u32,
31        }
32    }
33
34    /// Returns an empty span at one offset, used for end-of-input.
35    pub fn at(offset: usize) -> Span {
36        Span::new(offset, offset)
37    }
38
39    /// Returns the span covering both spans and everything between them.
40    pub fn to(self, other: Span) -> Span {
41        Span {
42            start: self.start.min(other.start),
43            end: self.end.max(other.end),
44        }
45    }
46
47    /// Returns the length of the span in bytes.
48    pub fn len(self) -> usize {
49        self.end.saturating_sub(self.start) as usize
50    }
51
52    /// Returns whether the span covers no bytes.
53    pub fn is_empty(self) -> bool {
54        self.end <= self.start
55    }
56
57    /// Returns the source bytes the span covers.
58    pub fn slice(self, source: &[u8]) -> &[u8] {
59        source
60            .get(self.start as usize..self.end as usize)
61            .unwrap_or(&[])
62    }
63}
64
65/// The quoting form an identifier was written with.
66///
67/// SQLite treats the four forms differently once semantics begin: a
68/// double-quoted word falls back to a string literal when it resolves to no
69/// name and the connection permits it, and the other three never do.
70#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
71pub enum QuoteForm {
72    /// `plain`, with no quoting at all.
73    Bare,
74    /// `"quoted"`, which may fall back to a string literal.
75    Double,
76    /// `[quoted]`, the MS-Access form.
77    Bracket,
78    /// `` `quoted` ``, the MySQL form.
79    Backtick,
80}
81
82/// The kind of a token.
83#[derive(Clone, Copy, Debug, PartialEq, Eq)]
84pub enum TokenKind {
85    /// A word: a keyword when `keyword` is set, otherwise an identifier.
86    Identifier {
87        /// The keyword this word spells, when it spells one.
88        keyword: Option<Keyword>,
89        /// How the identifier was quoted.
90        quote: QuoteForm,
91    },
92    /// A `'string'` literal.
93    String,
94    /// An `x'..'` blob literal.
95    Blob,
96    /// An integer literal, in decimal or hexadecimal.
97    Integer,
98    /// A floating-point literal.
99    Float,
100    /// A bound parameter.
101    Parameter,
102    /// Punctuation or an operator.
103    Punctuator(Punctuator),
104    /// End of input.
105    EndOfInput,
106}
107
108/// Every punctuation and operator token the pinned release accepts.
109#[derive(Clone, Copy, Debug, PartialEq, Eq)]
110pub enum Punctuator {
111    /// `(`
112    LeftParen,
113    /// `)`
114    RightParen,
115    /// `,`
116    Comma,
117    /// `;`
118    Semicolon,
119    /// `.`
120    Dot,
121    /// `+`
122    Plus,
123    /// `-`
124    Minus,
125    /// `*`
126    Star,
127    /// `/`
128    Slash,
129    /// `%`
130    Percent,
131    /// `=` or `==`
132    Equal,
133    /// `<>` or `!=`
134    NotEqual,
135    /// `<`
136    Less,
137    /// `<=`
138    LessEqual,
139    /// `>`
140    Greater,
141    /// `>=`
142    GreaterEqual,
143    /// `<<`
144    ShiftLeft,
145    /// `>>`
146    ShiftRight,
147    /// `&`
148    BitAnd,
149    /// `|`
150    BitOr,
151    /// `~`
152    BitNot,
153    /// `||`
154    Concat,
155    /// `->`
156    Arrow,
157    /// `->>`
158    DoubleArrow,
159    /// `<->`, pgvector's Euclidean distance.
160    L2Distance,
161    /// `<=>`, pgvector's cosine distance.
162    CosineDistance,
163    /// `<#>`, pgvector's negative inner product.
164    NegativeInnerProduct,
165    /// `<+>`, pgvector's taxicab distance.
166    L1Distance,
167    /// `<~>`, pgvector's Hamming distance.
168    HammingDistance,
169    /// `<%>`, pgvector's Jaccard distance.
170    JaccardDistance,
171}
172
173impl Punctuator {
174    /// Returns the canonical spelling, for diagnostics.
175    pub fn as_str(self) -> &'static str {
176        match self {
177            Punctuator::LeftParen => "(",
178            Punctuator::RightParen => ")",
179            Punctuator::Comma => ",",
180            Punctuator::Semicolon => ";",
181            Punctuator::Dot => ".",
182            Punctuator::Plus => "+",
183            Punctuator::Minus => "-",
184            Punctuator::Star => "*",
185            Punctuator::Slash => "/",
186            Punctuator::Percent => "%",
187            Punctuator::Equal => "=",
188            Punctuator::NotEqual => "<>",
189            Punctuator::Less => "<",
190            Punctuator::LessEqual => "<=",
191            Punctuator::Greater => ">",
192            Punctuator::GreaterEqual => ">=",
193            Punctuator::ShiftLeft => "<<",
194            Punctuator::ShiftRight => ">>",
195            Punctuator::BitAnd => "&",
196            Punctuator::BitOr => "|",
197            Punctuator::BitNot => "~",
198            Punctuator::Concat => "||",
199            Punctuator::Arrow => "->",
200            Punctuator::L2Distance => "<->",
201            Punctuator::CosineDistance => "<=>",
202            Punctuator::NegativeInnerProduct => "<#>",
203            Punctuator::L1Distance => "<+>",
204            Punctuator::HammingDistance => "<~>",
205            Punctuator::JaccardDistance => "<%>",
206            Punctuator::DoubleArrow => "->>",
207        }
208    }
209}
210
211/// One token: what it is and where it came from.
212#[derive(Clone, Copy, Debug, PartialEq, Eq)]
213pub struct Token {
214    /// What kind of token this is.
215    pub kind: TokenKind,
216    /// The bytes of the source it covers.
217    pub span: Span,
218}
219
220impl Token {
221    /// Returns the source text of the token.
222    pub fn text(self, source: &[u8]) -> &[u8] {
223        self.span.slice(source)
224    }
225
226    /// Returns the keyword this token spells, if it spells one.
227    pub fn keyword(self) -> Option<Keyword> {
228        match self.kind {
229            TokenKind::Identifier { keyword, .. } => keyword,
230            _ => None,
231        }
232    }
233
234    /// Returns whether the token is a punctuator of the given kind.
235    pub fn is(self, punctuator: Punctuator) -> bool {
236        self.kind == TokenKind::Punctuator(punctuator)
237    }
238}
239
240/// Why lexing stopped.
241#[derive(Clone, Copy, Debug, PartialEq, Eq)]
242pub enum LexErrorKind {
243    /// A quote or bracket was opened and never closed.
244    UnterminatedQuote,
245    /// A block comment was opened and never closed.
246    UnterminatedComment,
247    /// A byte that begins no token.
248    UnrecognisedByte,
249    /// A blob literal whose body is not an even number of hex digits.
250    MalformedBlob,
251    /// A numeric literal SQLite does not accept in this form.
252    MalformedNumber,
253    /// A parameter name that is empty or out of range.
254    MalformedParameter,
255}
256
257impl LexErrorKind {
258    /// Returns a stable one-line description.
259    pub fn message(self) -> &'static str {
260        match self {
261            LexErrorKind::UnterminatedQuote => "unrecognized token: unterminated quoted name",
262            LexErrorKind::UnterminatedComment => "unrecognized token: unterminated comment",
263            LexErrorKind::UnrecognisedByte => "unrecognized token",
264            LexErrorKind::MalformedBlob => "unrecognized token: malformed blob literal",
265            LexErrorKind::MalformedNumber => "unrecognized token: malformed numeric literal",
266            LexErrorKind::MalformedParameter => "unrecognized token: malformed parameter",
267        }
268    }
269}
270
271/// A lexing failure, with the offset of the byte that caused it.
272#[derive(Clone, Copy, Debug, PartialEq, Eq)]
273pub struct LexError {
274    /// Why it failed.
275    pub kind: LexErrorKind,
276    /// The offset the diagnostic points at.
277    pub offset: u32,
278}
279
280/// The scanner. It holds the source and a cursor and nothing else.
281#[derive(Clone, Debug)]
282pub struct Lexer<'a> {
283    source: &'a [u8],
284    offset: usize,
285}
286
287impl<'a> Lexer<'a> {
288    /// Returns a lexer positioned at the start of the source.
289    pub fn new(source: &'a [u8]) -> Lexer<'a> {
290        Lexer { source, offset: 0 }
291    }
292
293    /// Returns a lexer positioned at a byte offset in the source.
294    pub fn at(source: &'a [u8], offset: usize) -> Lexer<'a> {
295        Lexer {
296            source,
297            offset: offset.min(source.len()),
298        }
299    }
300
301    /// Returns the current byte offset.
302    pub fn offset(&self) -> usize {
303        self.offset
304    }
305
306    /// Returns the source being scanned.
307    pub fn source(&self) -> &'a [u8] {
308        self.source
309    }
310
311    /// Returns the byte at an offset, if there is one.
312    fn byte(&self, offset: usize) -> Option<u8> {
313        self.source.get(offset).copied()
314    }
315
316    /// Skips whitespace and both comment forms, returning an error for an
317    /// unterminated block comment at the offset that opened it.
318    fn skip_trivia(&mut self) -> Result<(), LexError> {
319        loop {
320            match self.byte(self.offset) {
321                Some(byte) if is_space(byte) => self.offset += 1,
322                Some(b'-') if self.byte(self.offset + 1) == Some(b'-') => {
323                    self.offset += 2;
324                    while let Some(byte) = self.byte(self.offset) {
325                        self.offset += 1;
326                        if byte == b'\n' {
327                            break;
328                        }
329                    }
330                }
331                Some(b'/') if self.byte(self.offset + 1) == Some(b'*') => {
332                    let opened = self.offset;
333                    self.offset += 2;
334                    loop {
335                        match self.byte(self.offset) {
336                            None => {
337                                // SQLite accepts an unterminated block comment
338                                // at the very end of input, treating it as
339                                // closed. It is the one place a missing
340                                // terminator is not an error.
341                                return Ok(());
342                            }
343                            Some(b'*') if self.byte(self.offset + 1) == Some(b'/') => {
344                                self.offset += 2;
345                                break;
346                            }
347                            Some(_) => self.offset += 1,
348                        }
349                    }
350                    let _ = opened;
351                }
352                _ => return Ok(()),
353            }
354        }
355    }
356
357    /// Scans and returns the next token, skipping trivia first.
358    pub fn next_token(&mut self) -> Result<Token, LexError> {
359        self.skip_trivia()?;
360        let start = self.offset;
361        let Some(byte) = self.byte(start) else {
362            return Ok(Token {
363                kind: TokenKind::EndOfInput,
364                span: Span::at(start),
365            });
366        };
367        match byte {
368            b'\'' => self.scan_quoted(start, b'\'', TokenKind::String),
369            b'"' => self.scan_quoted(
370                start,
371                b'"',
372                TokenKind::Identifier {
373                    keyword: None,
374                    quote: QuoteForm::Double,
375                },
376            ),
377            b'`' => self.scan_quoted(
378                start,
379                b'`',
380                TokenKind::Identifier {
381                    keyword: None,
382                    quote: QuoteForm::Backtick,
383                },
384            ),
385            b'[' => self.scan_bracket(start),
386            b'0'..=b'9' => self.scan_number(start),
387            b'.' if self.byte(start + 1).is_some_and(is_digit) => self.scan_number(start),
388            b'?' | b':' | b'@' | b'$' => self.scan_parameter(start),
389            byte if is_identifier_start(byte) => self.scan_word(start),
390            _ => self.scan_punctuator(start),
391        }
392    }
393
394    /// Scans a `'`, `"` or backtick delimited run, where the delimiter is
395    /// escaped by doubling it.
396    fn scan_quoted(
397        &mut self,
398        start: usize,
399        delimiter: u8,
400        kind: TokenKind,
401    ) -> Result<Token, LexError> {
402        let mut cursor = start + 1;
403        loop {
404            match self.byte(cursor) {
405                None => {
406                    return Err(LexError {
407                        kind: LexErrorKind::UnterminatedQuote,
408                        offset: start as u32,
409                    })
410                }
411                Some(byte) if byte == delimiter => {
412                    if self.byte(cursor + 1) == Some(delimiter) {
413                        cursor += 2;
414                        continue;
415                    }
416                    cursor += 1;
417                    break;
418                }
419                Some(_) => cursor += 1,
420            }
421        }
422        self.offset = cursor;
423        Ok(Token {
424            kind,
425            span: Span::new(start, cursor),
426        })
427    }
428
429    /// Scans a `[bracketed]` identifier, which has no escape at all.
430    fn scan_bracket(&mut self, start: usize) -> Result<Token, LexError> {
431        let mut cursor = start + 1;
432        loop {
433            match self.byte(cursor) {
434                None => {
435                    return Err(LexError {
436                        kind: LexErrorKind::UnterminatedQuote,
437                        offset: start as u32,
438                    })
439                }
440                Some(b']') => {
441                    cursor += 1;
442                    break;
443                }
444                Some(_) => cursor += 1,
445            }
446        }
447        self.offset = cursor;
448        Ok(Token {
449            kind: TokenKind::Identifier {
450                keyword: None,
451                quote: QuoteForm::Bracket,
452            },
453            span: Span::new(start, cursor),
454        })
455    }
456
457    /// Scans a bare word, which may be a keyword or the `x'..'` blob prefix.
458    fn scan_word(&mut self, start: usize) -> Result<Token, LexError> {
459        let mut cursor = start;
460        while self.byte(cursor).is_some_and(is_identifier_part) {
461            cursor += 1;
462        }
463        let word = self.source.get(start..cursor).unwrap_or(&[]);
464        if word.len() == 1
465            && word
466                .first()
467                .is_some_and(|byte| byte.eq_ignore_ascii_case(&b'x'))
468            && self.byte(cursor) == Some(b'\'')
469        {
470            return self.scan_blob(start, cursor);
471        }
472        self.offset = cursor;
473        Ok(Token {
474            kind: TokenKind::Identifier {
475                keyword: keyword::lookup(word),
476                quote: QuoteForm::Bare,
477            },
478            span: Span::new(start, cursor),
479        })
480    }
481
482    /// Scans the `'..'` body of a blob literal, which must be an even number of
483    /// hexadecimal digits.
484    fn scan_blob(&mut self, start: usize, quote: usize) -> Result<Token, LexError> {
485        let mut cursor = quote + 1;
486        let body = cursor;
487        loop {
488            match self.byte(cursor) {
489                None => {
490                    return Err(LexError {
491                        kind: LexErrorKind::UnterminatedQuote,
492                        offset: start as u32,
493                    })
494                }
495                Some(b'\'') => break,
496                Some(byte) if byte.is_ascii_hexdigit() => cursor += 1,
497                Some(_) => {
498                    return Err(LexError {
499                        kind: LexErrorKind::MalformedBlob,
500                        offset: start as u32,
501                    })
502                }
503            }
504        }
505        if !(cursor - body).is_multiple_of(2) {
506            return Err(LexError {
507                kind: LexErrorKind::MalformedBlob,
508                offset: start as u32,
509            });
510        }
511        self.offset = cursor + 1;
512        Ok(Token {
513            kind: TokenKind::Blob,
514            span: Span::new(start, cursor + 1),
515        })
516    }
517
518    /// Scans a numeric literal in every form the pinned release accepts.
519    fn scan_number(&mut self, start: usize) -> Result<Token, LexError> {
520        if self.byte(start) == Some(b'0')
521            && self
522                .byte(start + 1)
523                .is_some_and(|byte| byte.eq_ignore_ascii_case(&b'x'))
524        {
525            return self.scan_hex_number(start);
526        }
527        let mut cursor = start;
528        let mut float = false;
529        cursor = self.scan_digits(cursor);
530        if self.byte(cursor) == Some(b'.') {
531            float = true;
532            cursor = self.scan_digits(cursor + 1);
533        }
534        if self
535            .byte(cursor)
536            .is_some_and(|byte| byte.eq_ignore_ascii_case(&b'e'))
537        {
538            let mut lookahead = cursor + 1;
539            if matches!(self.byte(lookahead), Some(b'+') | Some(b'-')) {
540                lookahead += 1;
541            }
542            if self.byte(lookahead).is_some_and(is_digit) {
543                float = true;
544                cursor = self.scan_digits(lookahead);
545            }
546        }
547        // A digit run that runs straight into a word is not two tokens; SQLite
548        // rejects `123abc` rather than lexing an integer and an identifier.
549        if self.byte(cursor).is_some_and(is_identifier_part) {
550            return Err(LexError {
551                kind: LexErrorKind::MalformedNumber,
552                offset: start as u32,
553            });
554        }
555        self.offset = cursor;
556        Ok(Token {
557            kind: if float {
558                TokenKind::Float
559            } else {
560                TokenKind::Integer
561            },
562            span: Span::new(start, cursor),
563        })
564    }
565
566    /// Scans a `0x` literal, which has no fractional or exponent part.
567    fn scan_hex_number(&mut self, start: usize) -> Result<Token, LexError> {
568        let mut cursor = start + 2;
569        let digits = cursor;
570        while let Some(byte) = self.byte(cursor) {
571            // SQLite 3.46 and later accept a `_` between two hex digits.
572            let separator = byte == b'_'
573                && cursor > digits
574                && self
575                    .byte(cursor.wrapping_sub(1))
576                    .is_some_and(|b| b.is_ascii_hexdigit())
577                && self.byte(cursor + 1).is_some_and(|b| b.is_ascii_hexdigit());
578            if !(byte.is_ascii_hexdigit() || separator) {
579                break;
580            }
581            cursor += 1;
582        }
583        if cursor == digits || self.byte(cursor).is_some_and(is_identifier_part) {
584            return Err(LexError {
585                kind: LexErrorKind::MalformedNumber,
586                offset: start as u32,
587            });
588        }
589        self.offset = cursor;
590        Ok(Token {
591            kind: TokenKind::Integer,
592            span: Span::new(start, cursor),
593        })
594    }
595
596    /// Advances over a run of decimal digits, allowing SQLite's `_` separators.
597    fn scan_digits(&mut self, from: usize) -> usize {
598        let mut cursor = from;
599        while let Some(byte) = self.byte(cursor) {
600            if is_digit(byte) {
601                cursor += 1;
602                continue;
603            }
604            // A separator is only a separator between two digits.
605            if byte == b'_'
606                && cursor > from
607                && self.byte(cursor + 1).is_some_and(is_digit)
608                && self.byte(cursor.wrapping_sub(1)).is_some_and(is_digit)
609            {
610                cursor += 1;
611                continue;
612            }
613            break;
614        }
615        cursor
616    }
617
618    /// Scans `?`, `?NNN`, `:name`, `@name` and `$name`.
619    fn scan_parameter(&mut self, start: usize) -> Result<Token, LexError> {
620        let sigil = self.byte(start).unwrap_or(b'?');
621        let mut cursor = start + 1;
622        if sigil == b'?' {
623            cursor = self.scan_digits(cursor);
624            self.offset = cursor;
625            return Ok(Token {
626                kind: TokenKind::Parameter,
627                span: Span::new(start, cursor),
628            });
629        }
630        while self.byte(cursor).is_some_and(is_identifier_part) {
631            cursor += 1;
632        }
633        // `$name` accepts a bracketed or quoted suffix in SQLite's TCL variable
634        // syntax; the parenthesised form is the one that reaches SQL.
635        if sigil == b'$' && self.byte(cursor) == Some(b'(') {
636            while let Some(byte) = self.byte(cursor) {
637                cursor += 1;
638                if byte == b')' {
639                    break;
640                }
641            }
642        }
643        if cursor == start + 1 {
644            // A bare `:` or `@` is not a parameter. `:` is punctuation SQLite
645            // has no use for, so the byte is unrecognised.
646            return Err(LexError {
647                kind: LexErrorKind::MalformedParameter,
648                offset: start as u32,
649            });
650        }
651        self.offset = cursor;
652        Ok(Token {
653            kind: TokenKind::Parameter,
654            span: Span::new(start, cursor),
655        })
656    }
657
658    /// Scans punctuation and operators, longest form first.
659    fn scan_punctuator(&mut self, start: usize) -> Result<Token, LexError> {
660        let one = self.byte(start).unwrap_or(0);
661        let two = self.byte(start + 1);
662        let three = self.byte(start + 2);
663        let (punctuator, length) = match (one, two, three) {
664            (b'-', Some(b'>'), Some(b'>')) => (Punctuator::DoubleArrow, 3),
665            // **Before the two-byte forms, because `<=>` starts with `<=`.**
666            // These are pgvector's distance operators, and the longest-form-
667            // first rule is the only thing that keeps `v <=> q` from lexing as
668            // `v <= (> q)`.
669            (b'<', Some(b'-'), Some(b'>')) => (Punctuator::L2Distance, 3),
670            (b'<', Some(b'='), Some(b'>')) => (Punctuator::CosineDistance, 3),
671            (b'<', Some(b'#'), Some(b'>')) => (Punctuator::NegativeInnerProduct, 3),
672            (b'<', Some(b'+'), Some(b'>')) => (Punctuator::L1Distance, 3),
673            (b'<', Some(b'~'), Some(b'>')) => (Punctuator::HammingDistance, 3),
674            (b'<', Some(b'%'), Some(b'>')) => (Punctuator::JaccardDistance, 3),
675            (b'-', Some(b'>'), _) => (Punctuator::Arrow, 2),
676            (b'|', Some(b'|'), _) => (Punctuator::Concat, 2),
677            (b'<', Some(b'<'), _) => (Punctuator::ShiftLeft, 2),
678            (b'>', Some(b'>'), _) => (Punctuator::ShiftRight, 2),
679            (b'<', Some(b'='), _) => (Punctuator::LessEqual, 2),
680            (b'>', Some(b'='), _) => (Punctuator::GreaterEqual, 2),
681            (b'<', Some(b'>'), _) => (Punctuator::NotEqual, 2),
682            (b'!', Some(b'='), _) => (Punctuator::NotEqual, 2),
683            (b'=', Some(b'='), _) => (Punctuator::Equal, 2),
684            (b'(', _, _) => (Punctuator::LeftParen, 1),
685            (b')', _, _) => (Punctuator::RightParen, 1),
686            (b',', _, _) => (Punctuator::Comma, 1),
687            (b';', _, _) => (Punctuator::Semicolon, 1),
688            (b'.', _, _) => (Punctuator::Dot, 1),
689            (b'+', _, _) => (Punctuator::Plus, 1),
690            (b'-', _, _) => (Punctuator::Minus, 1),
691            (b'*', _, _) => (Punctuator::Star, 1),
692            (b'/', _, _) => (Punctuator::Slash, 1),
693            (b'%', _, _) => (Punctuator::Percent, 1),
694            (b'=', _, _) => (Punctuator::Equal, 1),
695            (b'<', _, _) => (Punctuator::Less, 1),
696            (b'>', _, _) => (Punctuator::Greater, 1),
697            (b'&', _, _) => (Punctuator::BitAnd, 1),
698            (b'|', _, _) => (Punctuator::BitOr, 1),
699            (b'~', _, _) => (Punctuator::BitNot, 1),
700            _ => {
701                return Err(LexError {
702                    kind: LexErrorKind::UnrecognisedByte,
703                    offset: start as u32,
704                })
705            }
706        };
707        self.offset = start + length;
708        Ok(Token {
709            kind: TokenKind::Punctuator(punctuator),
710            span: Span::new(start, start + length),
711        })
712    }
713}
714
715/// Returns whether the byte is SQL whitespace.
716pub fn is_space(byte: u8) -> bool {
717    matches!(byte, b' ' | b'\t' | b'\n' | b'\r' | 0x0b | 0x0c)
718}
719
720/// Returns whether the byte is an ASCII decimal digit.
721pub fn is_digit(byte: u8) -> bool {
722    byte.is_ascii_digit()
723}
724
725/// Returns whether the byte may begin a bare identifier.
726///
727/// SQLite treats every byte at or above 0x80 as an identifier character, which
728/// is how it accepts UTF-8 names without decoding them.
729pub fn is_identifier_start(byte: u8) -> bool {
730    byte.is_ascii_alphabetic() || byte == b'_' || byte >= 0x80
731}
732
733/// Returns whether the byte may continue a bare identifier.
734pub fn is_identifier_part(byte: u8) -> bool {
735    is_identifier_start(byte) || byte.is_ascii_digit() || byte == b'$'
736}
737
738/// Returns the text SQLite quotes in `unrecognized token: "..."` for a lexing failure.
739///
740/// SQLite names the bytes its tokenizer consumed before it gave up, and that length is
741/// not the same for every kind. An unterminated quote runs to the end of the input,
742/// because the tokenizer looks for the closing quote until it runs out of text. A
743/// malformed blob runs to the closing quote. A malformed number runs to the end of the
744/// identifier characters that follow it, so `123abc` and `0b11` are named whole. A bare
745/// `@`, `:` or `$`, and a byte that begins no token, are named alone. Measured against
746/// the pinned shell for each of these.
747///
748/// @param source - the SQL text the lexer was given
749/// @param error - the failure the lexer reported
750pub fn illegal_token_text(source: &[u8], error: LexError) -> &[u8] {
751    let start = (error.offset as usize).min(source.len());
752    let end = match error.kind {
753        LexErrorKind::UnterminatedQuote | LexErrorKind::UnterminatedComment => source.len(),
754        LexErrorKind::MalformedBlob => blob_token_end(source, start),
755        LexErrorKind::MalformedNumber => number_token_end(source, start),
756        LexErrorKind::MalformedParameter | LexErrorKind::UnrecognisedByte => start + 1,
757    };
758    source.get(start..end.min(source.len())).unwrap_or(&[])
759}
760
761/// Returns where SQLite's tokenizer stops on a blob literal whose body is not valid.
762///
763/// @param source - the SQL text
764/// @param start - the offset of the `x` that opens the literal
765fn blob_token_end(source: &[u8], start: usize) -> usize {
766    let mut cursor = start + 2;
767    while source.get(cursor).is_some_and(u8::is_ascii_hexdigit) {
768        cursor += 1;
769    }
770    while source.get(cursor).is_some_and(|byte| *byte != b'\'') {
771        cursor += 1;
772    }
773    (cursor + 1).min(source.len())
774}
775
776/// Returns where SQLite's tokenizer stops on a numeric literal that runs into a word.
777///
778/// The digits, fraction and exponent are read the way a valid number is, and then every
779/// identifier character after them is part of the same bad token.
780///
781/// @param source - the SQL text
782/// @param start - the offset of the first digit, or of the `.` of `.5x`
783fn number_token_end(source: &[u8], start: usize) -> usize {
784    let at = |index: usize| source.get(index).copied();
785    let mut cursor = start;
786    let hex = at(start) == Some(b'0')
787        && at(start + 1).is_some_and(|byte| byte.eq_ignore_ascii_case(&b'x'))
788        && at(start + 2).is_some_and(|byte| byte.is_ascii_hexdigit());
789    if hex {
790        cursor += 2;
791        while at(cursor).is_some_and(|byte| byte.is_ascii_hexdigit()) {
792            cursor += 1;
793        }
794    } else {
795        cursor = skip_digits_and_separators(source, cursor);
796        if at(cursor) == Some(b'.') {
797            cursor = skip_digits_and_separators(source, cursor + 1);
798        }
799        let sign = usize::from(matches!(at(cursor + 1), Some(b'+') | Some(b'-')));
800        let exponent = at(cursor).is_some_and(|byte| byte.eq_ignore_ascii_case(&b'e'))
801            && at(cursor + 1 + sign).is_some_and(|byte| byte.is_ascii_digit());
802        if exponent {
803            cursor = skip_digits_and_separators(source, cursor + 1 + sign);
804        }
805    }
806    while at(cursor).is_some_and(is_identifier_part) {
807        cursor += 1;
808    }
809    cursor
810}
811
812/// Advances over decimal digits, allowing a `_` that sits between two digits.
813///
814/// @param source - the SQL text
815/// @param from - where to start
816fn skip_digits_and_separators(source: &[u8], from: usize) -> usize {
817    let mut cursor = from;
818    loop {
819        let byte = source.get(cursor).copied();
820        let between = byte == Some(b'_')
821            && cursor > from
822            && source.get(cursor + 1).is_some_and(u8::is_ascii_digit)
823            && source.get(cursor - 1).is_some_and(u8::is_ascii_digit);
824        if byte.is_some_and(|byte| byte.is_ascii_digit()) || between {
825            cursor += 1;
826        } else {
827            return cursor;
828        }
829    }
830}
831
832/// Returns the unquoted text of an identifier token, undoubling escapes.
833///
834/// The common case is a bare word, which borrows. Only a quoted name that
835/// actually contains a doubled delimiter has to allocate.
836pub fn identifier_text<'a>(source: &'a [u8], token: Token) -> std::borrow::Cow<'a, [u8]> {
837    let raw = token.text(source);
838    let TokenKind::Identifier { quote, .. } = token.kind else {
839        return std::borrow::Cow::Borrowed(raw);
840    };
841    match quote {
842        QuoteForm::Bare => std::borrow::Cow::Borrowed(raw),
843        QuoteForm::Bracket => {
844            std::borrow::Cow::Borrowed(raw.get(1..raw.len().saturating_sub(1)).unwrap_or(&[]))
845        }
846        QuoteForm::Double => unquote(raw, b'"'),
847        QuoteForm::Backtick => unquote(raw, b'`'),
848    }
849}
850
851/// Returns the body of a `'string'` token with doubled quotes undoubled.
852pub fn string_text(source: &[u8], token: Token) -> std::borrow::Cow<'_, [u8]> {
853    unquote(token.text(source), b'\'')
854}
855
856/// Strips the delimiters and undoubles the escapes of a quoted run.
857fn unquote(raw: &[u8], delimiter: u8) -> std::borrow::Cow<'_, [u8]> {
858    let body = raw.get(1..raw.len().saturating_sub(1)).unwrap_or(&[]);
859    if !body.contains(&delimiter) {
860        return std::borrow::Cow::Borrowed(body);
861    }
862    let mut out = Vec::with_capacity(body.len());
863    let mut index = 0;
864    while let Some(byte) = body.get(index).copied() {
865        out.push(byte);
866        index += if byte == delimiter && body.get(index + 1) == Some(&delimiter) {
867            2
868        } else {
869            1
870        };
871    }
872    std::borrow::Cow::Owned(out)
873}
874
875/// Returns the bytes of a blob literal token, decoded from its hex digits.
876pub fn blob_bytes(source: &[u8], token: Token) -> Vec<u8> {
877    let raw = token.text(source);
878    let body = raw.get(2..raw.len().saturating_sub(1)).unwrap_or(&[]);
879    let mut out = Vec::with_capacity(body.len() / 2);
880    let mut index = 0;
881    while let (Some(high), Some(low)) = (body.get(index), body.get(index + 1)) {
882        let high = (*high as char).to_digit(16).unwrap_or(0) as u8;
883        let low = (*low as char).to_digit(16).unwrap_or(0) as u8;
884        out.push((high << 4) | low);
885        index += 2;
886    }
887    out
888}
889
890/// Converts a byte offset into a one-based line and column, scanning lazily.
891pub fn line_and_column(source: &[u8], offset: u32) -> (u32, u32) {
892    let limit = (offset as usize).min(source.len());
893    let mut line = 1u32;
894    let mut column = 1u32;
895    for byte in source.get(..limit).unwrap_or(&[]) {
896        if *byte == b'\n' {
897            line = line.saturating_add(1);
898            column = 1;
899        } else {
900            column = column.saturating_add(1);
901        }
902    }
903    (line, column)
904}
905
906#[cfg(test)]
907mod tests {
908    use super::*;
909
910    /// Collects every token of a source, for the tests below.
911    fn tokens(source: &str) -> Result<Vec<Token>, LexError> {
912        let bytes = source.as_bytes();
913        let mut lexer = Lexer::new(bytes);
914        let mut out = Vec::new();
915        loop {
916            let token = lexer.next_token()?;
917            if token.kind == TokenKind::EndOfInput {
918                return Ok(out);
919            }
920            out.push(token);
921        }
922    }
923
924    /// Every byte of the source belongs to exactly one token or to trivia, and
925    /// the spans are ordered and inside the source. This is the lexer's first
926    /// invariant and it is checkable directly.
927    #[test]
928    fn spans_are_ordered_disjoint_and_inside_the_source() {
929        let source = "SELECT a, 'x' /* c */ FROM t -- tail\nWHERE b=1;";
930        let found = tokens(source).expect("it lexes");
931        let mut previous_end = 0u32;
932        for token in &found {
933            assert!(token.span.start >= previous_end, "{token:?}");
934            assert!(token.span.end <= source.len() as u32, "{token:?}");
935            assert!(token.span.end > token.span.start, "{token:?}");
936            previous_end = token.span.end;
937        }
938    }
939
940    /// A keyword is an identifier token carrying a keyword, so the parser can
941    /// choose per position whether to accept it as a name.
942    #[test]
943    fn a_keyword_is_an_identifier_carrying_a_keyword() {
944        let found = tokens("select key").expect("it lexes");
945        assert_eq!(
946            found.first().and_then(|t| t.keyword()),
947            Some(Keyword::SELECT)
948        );
949        assert_eq!(found.get(1).and_then(|t| t.keyword()), Some(Keyword::KEY));
950        assert!(found
951            .get(1)
952            .and_then(|t| t.keyword())
953            .is_some_and(Keyword::may_fall_back));
954    }
955
956    /// The four quoting forms are distinguished, because only one of them may
957    /// later become a string literal.
958    #[test]
959    fn the_four_identifier_quote_forms_are_distinguished() {
960        let found = tokens("a \"b\" [c] `d`").expect("it lexes");
961        let forms: Vec<QuoteForm> = found
962            .iter()
963            .filter_map(|token| match token.kind {
964                TokenKind::Identifier { quote, .. } => Some(quote),
965                _ => None,
966            })
967            .collect();
968        assert_eq!(
969            forms,
970            vec![
971                QuoteForm::Bare,
972                QuoteForm::Double,
973                QuoteForm::Bracket,
974                QuoteForm::Backtick
975            ]
976        );
977    }
978
979    /// A doubled quote inside a string is one quote, and the borrow is only
980    /// given up when there is one.
981    #[test]
982    fn doubled_quotes_are_undoubled() {
983        let source = b"'it''s'";
984        let mut lexer = Lexer::new(source);
985        let token = lexer.next_token().expect("it lexes");
986        assert_eq!(token.kind, TokenKind::String);
987        assert_eq!(string_text(source, token).as_ref(), b"it's");
988
989        let plain = b"'plain'";
990        let mut lexer = Lexer::new(plain);
991        let token = lexer.next_token().expect("it lexes");
992        assert!(matches!(
993            string_text(plain, token),
994            std::borrow::Cow::Borrowed(_)
995        ));
996    }
997
998    /// An unterminated quote reports the byte that opened it, not the end of
999    /// input, because that is the offset a caller can act on.
1000    #[test]
1001    fn an_unterminated_quote_reports_its_opening_byte() {
1002        let mut lexer = Lexer::new(b"SELECT 'abc");
1003        assert_eq!(
1004            lexer.next_token().map(|t| t.kind),
1005            Ok(TokenKind::Identifier {
1006                keyword: Some(Keyword::SELECT),
1007                quote: QuoteForm::Bare
1008            })
1009        );
1010        assert_eq!(
1011            lexer.next_token(),
1012            Err(LexError {
1013                kind: LexErrorKind::UnterminatedQuote,
1014                offset: 7
1015            })
1016        );
1017    }
1018
1019    /// Numeric forms: decimal, leading dot, exponent, hexadecimal, and the
1020    /// underscore separators the pinned release accepts.
1021    #[test]
1022    fn every_numeric_form_lexes() {
1023        for (source, kind) in [
1024            ("1", TokenKind::Integer),
1025            ("1_000", TokenKind::Integer),
1026            ("0x1f", TokenKind::Integer),
1027            ("0XFF", TokenKind::Integer),
1028            ("1.5", TokenKind::Float),
1029            (".5", TokenKind::Float),
1030            ("1.", TokenKind::Float),
1031            ("1e10", TokenKind::Float),
1032            ("1E+10", TokenKind::Float),
1033            ("1.5e-3", TokenKind::Float),
1034        ] {
1035            let found = tokens(source).expect(source);
1036            assert_eq!(found.first().map(|t| t.kind), Some(kind), "{source}");
1037            assert_eq!(found.len(), 1, "{source}");
1038        }
1039    }
1040
1041    /// A number that runs into a word is one bad token, not two good ones.
1042    #[test]
1043    fn a_number_glued_to_a_word_is_rejected() {
1044        assert_eq!(
1045            tokens("123abc").map(|_| ()),
1046            Err(LexError {
1047                kind: LexErrorKind::MalformedNumber,
1048                offset: 0
1049            })
1050        );
1051    }
1052
1053    /// Blob literals must be an even number of hex digits, and the prefix is
1054    /// case-insensitive.
1055    #[test]
1056    fn blob_literals_decode() {
1057        let source = b"X'48690a'";
1058        let mut lexer = Lexer::new(source);
1059        let token = lexer.next_token().expect("it lexes");
1060        assert_eq!(token.kind, TokenKind::Blob);
1061        assert_eq!(blob_bytes(source, token), vec![0x48, 0x69, 0x0a]);
1062        assert!(tokens("x'abc'").is_err());
1063        assert!(tokens("x'zz'").is_err());
1064    }
1065
1066    /// Every parameter form lexes as one token.
1067    #[test]
1068    fn every_parameter_form_lexes() {
1069        for source in ["?", "?12", ":name", "@name", "$name"] {
1070            let found = tokens(source).expect(source);
1071            assert_eq!(found.len(), 1, "{source}");
1072            assert_eq!(found.first().map(|t| t.kind), Some(TokenKind::Parameter));
1073        }
1074        assert!(tokens(":").is_err());
1075    }
1076
1077    /// Both comment forms are trivia, and a line comment ends at the newline.
1078    #[test]
1079    fn comments_are_trivia() {
1080        let found = tokens("1 -- comment\n+ /* block */ 2").expect("it lexes");
1081        assert_eq!(found.len(), 3);
1082        assert!(found.get(1).is_some_and(|t| t.is(Punctuator::Plus)));
1083    }
1084
1085    /// An unterminated block comment at end of input is accepted, which is what
1086    /// the pinned release does.
1087    #[test]
1088    fn an_unterminated_block_comment_at_end_of_input_is_accepted() {
1089        let found = tokens("SELECT 1 /* trailing").expect("it lexes");
1090        assert_eq!(found.len(), 2);
1091    }
1092
1093    /// Operators lex longest-first, so `->>` never becomes `->` and `>`.
1094    #[test]
1095    fn operators_lex_longest_first() {
1096        let found = tokens("a->>b->c||d<<e").expect("it lexes");
1097        let punctuators: Vec<&'static str> = found
1098            .iter()
1099            .filter_map(|token| match token.kind {
1100                TokenKind::Punctuator(punctuator) => Some(punctuator.as_str()),
1101                _ => None,
1102            })
1103            .collect();
1104        assert_eq!(punctuators, vec!["->>", "->", "||", "<<"]);
1105    }
1106
1107    /// Line and column are derived from an offset only when asked for, and
1108    /// count from one.
1109    #[test]
1110    fn line_and_column_count_from_one() {
1111        let source = b"SELECT\n  1";
1112        assert_eq!(line_and_column(source, 0), (1, 1));
1113        assert_eq!(line_and_column(source, 9), (2, 3));
1114    }
1115
1116    /// Bytes above 0x7f are identifier characters, so a UTF-8 name lexes as one
1117    /// token without the lexer decoding anything.
1118    #[test]
1119    fn high_bytes_are_identifier_characters() {
1120        let found = tokens("naïve").expect("it lexes");
1121        assert_eq!(found.len(), 1);
1122    }
1123}