use std::fmt;
mod keywords;
use keywords::keyword_lookup;
#[derive(Debug, Clone, PartialEq)]
#[non_exhaustive]
pub struct Token<'a> {
pub(crate) kind: TokenKind,
pub(crate) text: &'a str,
pub(crate) span: Span,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub struct Span {
pub(crate) start: usize,
pub(crate) end: usize,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub enum TokenKind {
Identifier,
StringLiteral,
IntegerLiteral,
FloatLiteral,
Parameter,
Select,
From,
Where,
And,
Or,
Not,
As,
In,
Between,
Like,
Escape,
Order,
By,
Asc,
Desc,
Top,
Distinct,
Value,
Group,
Having,
Join,
Cross,
Inner,
Exists,
Array,
Null,
True,
False,
Undefined,
Offset,
Limit,
Udf,
Is,
Let,
Left,
Right,
Set,
Over,
Rank,
For,
Plus,
Minus,
Star,
Slash,
Percent,
Tilde,
Ampersand,
Pipe,
Caret,
Eq,
NotEq,
Lt,
Gt,
LtEq,
GtEq,
LeftShift,
RightShift,
ZeroFillRightShift,
StringConcat,
Coalesce,
Question,
Colon,
Bang,
LParen,
RParen,
LBracket,
RBracket,
LBrace,
RBrace,
Dot,
Comma,
Eof,
ErrUnterminatedString,
ErrUnterminatedQuotedIdentifier,
ErrUnterminatedBlockComment,
}
impl fmt::Display for TokenKind {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let s = match self {
Self::Identifier => "identifier",
Self::StringLiteral => "string",
Self::IntegerLiteral => "integer",
Self::FloatLiteral => "float",
Self::Parameter => "parameter",
Self::ErrUnterminatedString => "unterminated string literal",
Self::ErrUnterminatedQuotedIdentifier => "unterminated quoted identifier",
Self::ErrUnterminatedBlockComment => "unterminated block comment",
Self::Select => "SELECT",
Self::From => "FROM",
Self::Where => "WHERE",
Self::And => "AND",
Self::Or => "OR",
Self::Not => "NOT",
Self::As => "AS",
Self::In => "IN",
Self::Between => "BETWEEN",
Self::Like => "LIKE",
Self::Escape => "ESCAPE",
Self::Order => "ORDER",
Self::By => "BY",
Self::Asc => "ASC",
Self::Desc => "DESC",
Self::Top => "TOP",
Self::Distinct => "DISTINCT",
Self::Value => "VALUE",
Self::Group => "GROUP",
Self::Having => "HAVING",
Self::Join => "JOIN",
Self::Cross => "CROSS",
Self::Inner => "INNER",
Self::Exists => "EXISTS",
Self::Array => "ARRAY",
Self::Null => "null",
Self::True => "true",
Self::False => "false",
Self::Undefined => "undefined",
Self::Offset => "OFFSET",
Self::Limit => "LIMIT",
Self::Udf => "udf",
Self::Is => "IS",
Self::Let => "LET",
Self::Left => "LEFT",
Self::Right => "RIGHT",
Self::Set => "SET",
Self::Over => "OVER",
Self::Rank => "RANK",
Self::For => "FOR",
Self::Plus => "+",
Self::Minus => "-",
Self::Star => "*",
Self::Slash => "/",
Self::Percent => "%",
Self::Tilde => "~",
Self::Ampersand => "&",
Self::Pipe => "|",
Self::Caret => "^",
Self::Eq => "=",
Self::NotEq => "!=",
Self::Lt => "<",
Self::Gt => ">",
Self::LtEq => "<=",
Self::GtEq => ">=",
Self::LeftShift => "<<",
Self::RightShift => ">>",
Self::ZeroFillRightShift => ">>>",
Self::StringConcat => "||",
Self::Coalesce => "??",
Self::Question => "?",
Self::Colon => ":",
Self::Bang => "!",
Self::LParen => "(",
Self::RParen => ")",
Self::LBracket => "[",
Self::RBracket => "]",
Self::LBrace => "{",
Self::RBrace => "}",
Self::Dot => ".",
Self::Comma => ",",
Self::Eof => "EOF",
};
write!(f, "{s}")
}
}
pub struct Lexer<'a> {
source: &'a str,
bytes: &'a [u8],
pos: usize,
pending_block_comment_error: Option<usize>,
}
impl<'a> Lexer<'a> {
pub(crate) fn new(source: &'a str) -> Self {
Self {
source,
bytes: source.as_bytes(),
pos: 0,
pending_block_comment_error: None,
}
}
pub(crate) fn next_token(&mut self) -> Token<'a> {
self.skip_whitespace_and_comments();
if let Some(err_start) = self.pending_block_comment_error.take() {
return Token {
kind: TokenKind::ErrUnterminatedBlockComment,
text: &self.source[err_start..self.pos],
span: Span {
start: err_start,
end: self.pos,
},
};
}
if self.pos >= self.bytes.len() {
return Token {
kind: TokenKind::Eof,
text: "",
span: Span {
start: self.pos,
end: self.pos,
},
};
}
let start = self.pos;
let ch = self.bytes[self.pos];
match ch {
b'\'' => self.scan_string_literal(start),
b'"' => self.scan_quoted_identifier(start),
b'@' => self.scan_parameter(start),
b'0'..=b'9' => self.scan_number(start),
b'a'..=b'z' | b'A'..=b'Z' | b'_' => self.scan_identifier(start),
b'(' => self.single_char_token(start, TokenKind::LParen),
b')' => self.single_char_token(start, TokenKind::RParen),
b'[' => self.single_char_token(start, TokenKind::LBracket),
b']' => self.single_char_token(start, TokenKind::RBracket),
b'{' => self.single_char_token(start, TokenKind::LBrace),
b'}' => self.single_char_token(start, TokenKind::RBrace),
b'.' => self.single_char_token(start, TokenKind::Dot),
b',' => self.single_char_token(start, TokenKind::Comma),
b'+' => self.single_char_token(start, TokenKind::Plus),
b'-' => self.single_char_token(start, TokenKind::Minus),
b'*' => self.single_char_token(start, TokenKind::Star),
b'/' => self.single_char_token(start, TokenKind::Slash),
b'%' => self.single_char_token(start, TokenKind::Percent),
b'~' => self.single_char_token(start, TokenKind::Tilde),
b'^' => self.single_char_token(start, TokenKind::Caret),
b'=' => self.single_char_token(start, TokenKind::Eq),
b':' => self.single_char_token(start, TokenKind::Colon),
b'!' => {
self.pos += 1;
if self.peek() == Some(b'=') {
self.pos += 1;
self.make_token(start, TokenKind::NotEq)
} else {
self.make_token(start, TokenKind::Bang)
}
}
b'<' => {
self.pos += 1;
match self.peek() {
Some(b'=') => {
self.pos += 1;
self.make_token(start, TokenKind::LtEq)
}
Some(b'<') => {
self.pos += 1;
self.make_token(start, TokenKind::LeftShift)
}
Some(b'>') => {
self.pos += 1;
self.make_token(start, TokenKind::NotEq)
}
_ => self.make_token(start, TokenKind::Lt),
}
}
b'>' => {
self.pos += 1;
match self.peek() {
Some(b'=') => {
self.pos += 1;
self.make_token(start, TokenKind::GtEq)
}
Some(b'>') => {
self.pos += 1;
if self.peek() == Some(b'>') {
self.pos += 1;
self.make_token(start, TokenKind::ZeroFillRightShift)
} else {
self.make_token(start, TokenKind::RightShift)
}
}
_ => self.make_token(start, TokenKind::Gt),
}
}
b'&' => {
self.pos += 1;
if self.peek() == Some(b'&') {
self.pos += 1;
self.make_token(start, TokenKind::And)
} else {
self.make_token(start, TokenKind::Ampersand)
}
}
b'|' => {
self.pos += 1;
match self.peek() {
Some(b'|') => {
self.pos += 1;
self.make_token(start, TokenKind::StringConcat)
}
_ => self.make_token(start, TokenKind::Pipe),
}
}
b'?' => {
self.pos += 1;
if self.peek() == Some(b'?') {
self.pos += 1;
self.make_token(start, TokenKind::Coalesce)
} else {
self.make_token(start, TokenKind::Question)
}
}
_ => {
let mut next_pos = self.pos + 1;
while next_pos < self.bytes.len() && !self.source.is_char_boundary(next_pos) {
next_pos += 1;
}
self.pos = next_pos;
self.make_token(start, TokenKind::Identifier)
}
}
}
pub fn tokenize(source: &'a str) -> Vec<Token<'a>> {
let mut lexer = Lexer::new(source);
let mut tokens = Vec::new();
loop {
let tok = lexer.next_token();
if tok.kind == TokenKind::Eof {
break;
}
tokens.push(tok);
}
tokens
}
fn peek(&self) -> Option<u8> {
self.bytes.get(self.pos).copied()
}
fn skip_whitespace_and_comments(&mut self) {
loop {
while self.pos < self.bytes.len() && self.bytes[self.pos].is_ascii_whitespace() {
self.pos += 1;
}
if self.pos + 1 < self.bytes.len()
&& self.bytes[self.pos] == b'-'
&& self.bytes[self.pos + 1] == b'-'
{
self.pos += 2;
while self.pos < self.bytes.len() && self.bytes[self.pos] != b'\n' {
self.pos += 1;
}
continue;
}
if self.pos + 1 < self.bytes.len()
&& self.bytes[self.pos] == b'/'
&& self.bytes[self.pos + 1] == b'*'
{
let comment_start = self.pos;
self.pos += 2;
while self.pos + 1 < self.bytes.len()
&& !(self.bytes[self.pos] == b'*' && self.bytes[self.pos + 1] == b'/')
{
self.pos += 1;
}
if self.pos + 1 < self.bytes.len() {
self.pos += 2; } else {
self.pos = self.bytes.len();
self.pending_block_comment_error = Some(comment_start);
return;
}
continue;
}
break;
}
}
fn scan_string_literal(&mut self, start: usize) -> Token<'a> {
self.pos += 1; while self.pos < self.bytes.len() {
if self.bytes[self.pos] == b'\'' {
if self.pos + 1 < self.bytes.len() && self.bytes[self.pos + 1] == b'\'' {
self.pos += 2;
} else {
self.pos += 1; return self.make_token(start, TokenKind::StringLiteral);
}
} else {
self.pos += 1;
}
}
self.make_token(start, TokenKind::ErrUnterminatedString)
}
fn scan_quoted_identifier(&mut self, start: usize) -> Token<'a> {
self.pos += 1; while self.pos < self.bytes.len() && self.bytes[self.pos] != b'"' {
self.pos += 1;
}
if self.pos < self.bytes.len() {
self.pos += 1; self.make_token(start, TokenKind::Identifier)
} else {
self.make_token(start, TokenKind::ErrUnterminatedQuotedIdentifier)
}
}
fn scan_parameter(&mut self, start: usize) -> Token<'a> {
self.pos += 1; while self.pos < self.bytes.len() && is_ident_char(self.bytes[self.pos]) {
self.pos += 1;
}
self.make_token(start, TokenKind::Parameter)
}
fn scan_number(&mut self, start: usize) -> Token<'a> {
let mut is_float = false;
while self.pos < self.bytes.len() && self.bytes[self.pos].is_ascii_digit() {
self.pos += 1;
}
if self.pos < self.bytes.len() && self.bytes[self.pos] == b'.' {
if self.pos + 1 < self.bytes.len() && self.bytes[self.pos + 1].is_ascii_digit() {
is_float = true;
self.pos += 1; while self.pos < self.bytes.len() && self.bytes[self.pos].is_ascii_digit() {
self.pos += 1;
}
}
}
if self.pos < self.bytes.len()
&& (self.bytes[self.pos] == b'e' || self.bytes[self.pos] == b'E')
{
is_float = true;
self.pos += 1;
if self.pos < self.bytes.len()
&& (self.bytes[self.pos] == b'+' || self.bytes[self.pos] == b'-')
{
self.pos += 1;
}
while self.pos < self.bytes.len() && self.bytes[self.pos].is_ascii_digit() {
self.pos += 1;
}
}
if is_float {
self.make_token(start, TokenKind::FloatLiteral)
} else {
self.make_token(start, TokenKind::IntegerLiteral)
}
}
fn scan_identifier(&mut self, start: usize) -> Token<'a> {
while self.pos < self.bytes.len() && is_ident_char(self.bytes[self.pos]) {
self.pos += 1;
}
let text = &self.source[start..self.pos];
let kind = keyword_lookup(text);
Token {
kind,
text,
span: Span {
start,
end: self.pos,
},
}
}
fn single_char_token(&mut self, start: usize, kind: TokenKind) -> Token<'a> {
self.pos += 1;
self.make_token(start, kind)
}
fn make_token(&self, start: usize, kind: TokenKind) -> Token<'a> {
Token {
kind,
text: &self.source[start..self.pos],
span: Span {
start,
end: self.pos,
},
}
}
}
fn is_ident_char(b: u8) -> bool {
b.is_ascii_alphanumeric() || b == b'_'
}
pub(crate) fn extract_string_content(token_text: &str) -> String {
let inner = if token_text.len() >= 2
&& token_text.starts_with(char::from(b'\''))
&& token_text.ends_with(char::from(b'\''))
{
&token_text[1..token_text.len() - 1]
} else {
token_text
};
inner.replace("''", "'")
}
pub(crate) fn extract_identifier(token_text: &str) -> &str {
if token_text.starts_with('"') && token_text.ends_with('"') && token_text.len() >= 2 {
&token_text[1..token_text.len() - 1]
} else {
token_text
}
}
pub(crate) fn extract_parameter_name(token_text: &str) -> &str {
if let Some(stripped) = token_text.strip_prefix('@') {
stripped
} else {
token_text
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn simple_select() {
let tokens = Lexer::tokenize("SELECT * FROM c");
assert_eq!(tokens.len(), 4);
assert_eq!(tokens[0].kind, TokenKind::Select);
assert_eq!(tokens[1].kind, TokenKind::Star);
assert_eq!(tokens[2].kind, TokenKind::From);
assert_eq!(tokens[3].kind, TokenKind::Identifier);
assert_eq!(tokens[3].text, "c");
}
#[test]
fn string_literal() {
let tokens = Lexer::tokenize("'hello world'");
assert_eq!(tokens.len(), 1);
assert_eq!(tokens[0].kind, TokenKind::StringLiteral);
assert_eq!(extract_string_content(tokens[0].text), "hello world");
}
#[test]
fn escaped_string() {
let tokens = Lexer::tokenize("'it''s'");
assert_eq!(tokens.len(), 1);
assert_eq!(extract_string_content(tokens[0].text), "it's");
}
#[test]
fn unterminated_string_yields_error_token() {
let tokens = Lexer::tokenize("'unclosed");
assert_eq!(tokens.len(), 1);
assert_eq!(tokens[0].kind, TokenKind::ErrUnterminatedString);
}
#[test]
fn unterminated_string_with_trailing_input_yields_error_token() {
let tokens = Lexer::tokenize("SELECT 'unclosed FROM c");
assert_eq!(tokens.first().map(|t| t.kind), Some(TokenKind::Select));
assert!(
tokens
.iter()
.any(|t| t.kind == TokenKind::ErrUnterminatedString),
"expected an ErrUnterminatedString token; got {:?}",
tokens.iter().map(|t| t.kind).collect::<Vec<_>>()
);
}
#[test]
fn unterminated_quoted_identifier_yields_error_token() {
let tokens = Lexer::tokenize("SELECT \"unclosed FROM c");
assert_eq!(tokens.first().map(|t| t.kind), Some(TokenKind::Select));
assert!(
tokens
.iter()
.any(|t| t.kind == TokenKind::ErrUnterminatedQuotedIdentifier),
"expected ErrUnterminatedQuotedIdentifier; got {:?}",
tokens.iter().map(|t| t.kind).collect::<Vec<_>>()
);
}
#[test]
fn unterminated_block_comment_yields_error_token() {
let tokens = Lexer::tokenize("SELECT /* unclosed");
assert_eq!(tokens.first().map(|t| t.kind), Some(TokenKind::Select));
assert!(
tokens
.iter()
.any(|t| t.kind == TokenKind::ErrUnterminatedBlockComment),
"expected ErrUnterminatedBlockComment; got {:?}",
tokens.iter().map(|t| t.kind).collect::<Vec<_>>()
);
}
#[test]
fn non_ascii_character_respects_char_boundary() {
let tokens = Lexer::tokenize("\u{00e9}"); assert_eq!(tokens.len(), 1, "expected one token, got {:?}", tokens);
assert_eq!(tokens[0].text.len(), 2);
assert_eq!(tokens[0].text, "\u{00e9}");
}
#[test]
fn numbers() {
let tokens = Lexer::tokenize("42 3.14 1e10 2.5E-3");
assert_eq!(tokens[0].kind, TokenKind::IntegerLiteral);
assert_eq!(tokens[1].kind, TokenKind::FloatLiteral);
assert_eq!(tokens[2].kind, TokenKind::FloatLiteral);
assert_eq!(tokens[3].kind, TokenKind::FloatLiteral);
}
#[test]
fn parameters() {
let tokens = Lexer::tokenize("@p1 @customer_id");
assert_eq!(tokens[0].kind, TokenKind::Parameter);
assert_eq!(extract_parameter_name(tokens[0].text), "p1");
assert_eq!(tokens[1].kind, TokenKind::Parameter);
assert_eq!(extract_parameter_name(tokens[1].text), "customer_id");
}
#[test]
fn operators() {
let tokens = Lexer::tokenize("!= <= >= << >> >>> || ??");
assert_eq!(tokens[0].kind, TokenKind::NotEq);
assert_eq!(tokens[1].kind, TokenKind::LtEq);
assert_eq!(tokens[2].kind, TokenKind::GtEq);
assert_eq!(tokens[3].kind, TokenKind::LeftShift);
assert_eq!(tokens[4].kind, TokenKind::RightShift);
assert_eq!(tokens[5].kind, TokenKind::ZeroFillRightShift);
assert_eq!(tokens[6].kind, TokenKind::StringConcat);
assert_eq!(tokens[7].kind, TokenKind::Coalesce);
}
#[test]
fn keywords_case_insensitive() {
let tokens = Lexer::tokenize("select FROM Where");
assert_eq!(tokens[0].kind, TokenKind::Select);
assert_eq!(tokens[1].kind, TokenKind::From);
assert_eq!(tokens[2].kind, TokenKind::Where);
}
#[test]
fn line_comment() {
let tokens = Lexer::tokenize("SELECT -- this is a comment\n* FROM c");
assert_eq!(tokens.len(), 4);
assert_eq!(tokens[0].kind, TokenKind::Select);
assert_eq!(tokens[1].kind, TokenKind::Star);
}
#[test]
fn block_comment() {
let tokens = Lexer::tokenize("SELECT /* comment */ * FROM c");
assert_eq!(tokens.len(), 4);
assert_eq!(tokens[0].kind, TokenKind::Select);
assert_eq!(tokens[1].kind, TokenKind::Star);
}
#[test]
fn full_query_tokenization() {
let tokens = Lexer::tokenize(
"SELECT c.name, c.age FROM c WHERE c.pk = 'hello' AND c.age > 21 ORDER BY c.age DESC",
);
assert!(tokens.len() > 10);
assert_eq!(tokens[0].kind, TokenKind::Select);
assert_eq!(tokens.last().unwrap().kind, TokenKind::Desc);
}
}