use super::error::{GqlError, GqlResult};
#[derive(Debug, Clone, PartialEq)]
pub struct Token {
pub kind: TokenKind,
pub lexeme: String,
pub position: usize,
}
#[derive(Debug, Clone, PartialEq)]
pub enum TokenKind {
Match,
Where,
Return,
Filter,
OrderBy,
Limit,
Offset,
Graph,
Next,
Insert,
Update,
Delete,
Set,
Remove,
Create,
Merge,
With,
Optional,
Union,
Distinct,
As,
Asc,
Desc,
Assign,
Equal,
NotEqual,
LessThan,
LessEqual,
GreaterThan,
GreaterEqual,
Plus,
Minus,
Multiply,
Divide,
Modulo,
And,
Or,
Not,
In,
Contains,
StartsWith,
EndsWith,
LeftParen,
RightParen,
LeftBracket,
RightBracket,
LeftBrace,
RightBrace,
Comma,
Dot,
DotDot,
Colon,
Semicolon,
Arrow,
Pipe,
Integer(i64),
Float(f64),
String(String),
Boolean(bool),
Null,
Identifier(String),
Parameter(String),
Eof,
Newline,
}
pub struct Lexer {
input: Vec<char>,
current: usize,
position: usize,
}
impl Lexer {
pub fn new(input: &str) -> Self {
Self {
input: input.chars().collect(),
current: 0,
position: 0,
}
}
pub fn tokenize(&mut self) -> GqlResult<Vec<Token>> {
let mut tokens = Vec::new();
while !self.is_at_end() {
let _start_pos = self.position;
match self.scan_token() {
Ok(Some(token)) => tokens.push(token),
Ok(None) => {} Err(err) => return Err(err),
}
}
tokens.push(Token {
kind: TokenKind::Eof,
lexeme: String::new(),
position: self.position,
});
Ok(tokens)
}
fn scan_token(&mut self) -> GqlResult<Option<Token>> {
let start_pos = self.position;
let c = self.advance();
match c {
' ' | '\r' | '\t' => Ok(None),
'\n' => Ok(Some(Token {
kind: TokenKind::Newline,
lexeme: "\n".to_string(),
position: start_pos,
})),
'(' => Ok(Some(self.make_token(TokenKind::LeftParen, start_pos))),
')' => Ok(Some(self.make_token(TokenKind::RightParen, start_pos))),
'[' => Ok(Some(self.make_token(TokenKind::LeftBracket, start_pos))),
']' => Ok(Some(self.make_token(TokenKind::RightBracket, start_pos))),
'{' => Ok(Some(self.make_token(TokenKind::LeftBrace, start_pos))),
'}' => Ok(Some(self.make_token(TokenKind::RightBrace, start_pos))),
',' => Ok(Some(self.make_token(TokenKind::Comma, start_pos))),
'.' => {
if self.match_char('.') {
Ok(Some(self.make_token(TokenKind::DotDot, start_pos)))
} else {
Ok(Some(self.make_token(TokenKind::Dot, start_pos)))
}
}
':' => Ok(Some(self.make_token(TokenKind::Colon, start_pos))),
';' => Ok(Some(self.make_token(TokenKind::Semicolon, start_pos))),
'|' => Ok(Some(self.make_token(TokenKind::Pipe, start_pos))),
'+' => Ok(Some(self.make_token(TokenKind::Plus, start_pos))),
'*' => Ok(Some(self.make_token(TokenKind::Multiply, start_pos))),
'/' => Ok(Some(self.make_token(TokenKind::Divide, start_pos))),
'%' => Ok(Some(self.make_token(TokenKind::Modulo, start_pos))),
'=' => {
if self.match_char('=') {
Ok(Some(self.make_token(TokenKind::Equal, start_pos)))
} else {
Ok(Some(self.make_token(TokenKind::Assign, start_pos)))
}
}
'!' => {
if self.match_char('=') {
Ok(Some(self.make_token(TokenKind::NotEqual, start_pos)))
} else {
Err(GqlError::LexError {
message: "Unexpected character '!'".to_string(),
position: start_pos,
})
}
}
'<' => {
if self.match_char('=') {
Ok(Some(self.make_token(TokenKind::LessEqual, start_pos)))
} else if self.match_char('>') {
Ok(Some(self.make_token(TokenKind::NotEqual, start_pos)))
} else {
Ok(Some(self.make_token(TokenKind::LessThan, start_pos)))
}
}
'>' => {
if self.match_char('=') {
Ok(Some(self.make_token(TokenKind::GreaterEqual, start_pos)))
} else {
Ok(Some(self.make_token(TokenKind::GreaterThan, start_pos)))
}
}
'-' => {
if self.match_char('>') {
Ok(Some(self.make_token(TokenKind::Arrow, start_pos)))
} else {
Ok(Some(self.make_token(TokenKind::Minus, start_pos)))
}
}
'"' | '\'' => self.scan_string(c, start_pos),
'$' => self.scan_parameter(start_pos),
'0'..='9' => self.scan_number(start_pos),
'a'..='z' | 'A'..='Z' | '_' => self.scan_identifier(start_pos),
_ => Err(GqlError::LexError {
message: format!("Unexpected character '{c}'"),
position: start_pos,
}),
}
}
fn scan_string(&mut self, quote: char, start_pos: usize) -> GqlResult<Option<Token>> {
let mut value = String::new();
while !self.is_at_end() && self.peek() != quote {
if self.peek() == '\\' {
self.advance(); match self.advance() {
'n' => value.push('\n'),
'r' => value.push('\r'),
't' => value.push('\t'),
'\\' => value.push('\\'),
'\'' => value.push('\''),
'"' => value.push('"'),
c => {
return Err(GqlError::LexError {
message: format!("Invalid escape sequence: \\{c}"),
position: self.position - 1,
});
}
}
} else {
value.push(self.advance());
}
}
if self.is_at_end() {
return Err(GqlError::LexError {
message: "Unterminated string".to_string(),
position: start_pos,
});
}
self.advance();
Ok(Some(Token {
kind: TokenKind::String(value),
lexeme: self.lexeme_from(start_pos),
position: start_pos,
}))
}
fn scan_parameter(&mut self, start_pos: usize) -> GqlResult<Option<Token>> {
while self.peek().is_alphanumeric() || self.peek() == '_' {
self.advance();
}
let lexeme = self.lexeme_from(start_pos);
let param_name = lexeme[1..].to_string();
Ok(Some(Token {
kind: TokenKind::Parameter(param_name),
lexeme,
position: start_pos,
}))
}
fn scan_number(&mut self, start_pos: usize) -> GqlResult<Option<Token>> {
while self.peek().is_ascii_digit() {
self.advance();
}
if self.peek() == '.' && self.peek_next().is_ascii_digit() {
self.advance();
while self.peek().is_ascii_digit() {
self.advance();
}
let lexeme = self.lexeme_from(start_pos);
let value = lexeme.parse::<f64>().map_err(|_| GqlError::LexError {
message: format!("Invalid float literal: {lexeme}"),
position: start_pos,
})?;
Ok(Some(Token {
kind: TokenKind::Float(value),
lexeme,
position: start_pos,
}))
} else {
let lexeme = self.lexeme_from(start_pos);
let value = lexeme.parse::<i64>().map_err(|_| GqlError::LexError {
message: format!("Invalid integer literal: {lexeme}"),
position: start_pos,
})?;
Ok(Some(Token {
kind: TokenKind::Integer(value),
lexeme,
position: start_pos,
}))
}
}
fn scan_identifier(&mut self, start_pos: usize) -> GqlResult<Option<Token>> {
while self.peek().is_alphanumeric() || self.peek() == '_' {
self.advance();
}
let lexeme = self.lexeme_from(start_pos);
let kind = self.keyword_or_identifier(&lexeme);
Ok(Some(Token {
kind,
lexeme,
position: start_pos,
}))
}
fn keyword_or_identifier(&self, text: &str) -> TokenKind {
match text.to_uppercase().as_str() {
"MATCH" => TokenKind::Match,
"WHERE" => TokenKind::Where,
"RETURN" => TokenKind::Return,
"FILTER" => TokenKind::Filter,
"ORDER" => {
TokenKind::Identifier(text.to_string())
}
"BY" => TokenKind::Identifier(text.to_string()),
"LIMIT" => TokenKind::Limit,
"OFFSET" => TokenKind::Offset,
"GRAPH" => TokenKind::Graph,
"NEXT" => TokenKind::Next,
"INSERT" => TokenKind::Insert,
"UPDATE" => TokenKind::Update,
"DELETE" => TokenKind::Delete,
"SET" => TokenKind::Set,
"REMOVE" => TokenKind::Remove,
"CREATE" => TokenKind::Create,
"MERGE" => TokenKind::Merge,
"WITH" => TokenKind::With,
"OPTIONAL" => TokenKind::Optional,
"UNION" => TokenKind::Union,
"DISTINCT" => TokenKind::Distinct,
"AS" => TokenKind::As,
"ASC" => TokenKind::Asc,
"DESC" => TokenKind::Desc,
"AND" => TokenKind::And,
"OR" => TokenKind::Or,
"NOT" => TokenKind::Not,
"IN" => TokenKind::In,
"CONTAINS" => TokenKind::Contains,
"STARTS" => TokenKind::Identifier(text.to_string()), "ENDS" => TokenKind::Identifier(text.to_string()), "TRUE" => TokenKind::Boolean(true),
"FALSE" => TokenKind::Boolean(false),
"NULL" => TokenKind::Null,
_ => TokenKind::Identifier(text.to_string()),
}
}
fn make_token(&self, kind: TokenKind, start_pos: usize) -> Token {
Token {
kind,
lexeme: self.lexeme_from(start_pos),
position: start_pos,
}
}
fn lexeme_from(&self, start_pos: usize) -> String {
self.input[start_pos..self.current].iter().collect()
}
fn advance(&mut self) -> char {
if self.is_at_end() {
'\0'
} else {
let c = self.input[self.current];
self.current += 1;
self.position += 1;
c
}
}
fn peek(&self) -> char {
if self.is_at_end() {
'\0'
} else {
self.input[self.current]
}
}
fn peek_next(&self) -> char {
if self.current + 1 >= self.input.len() {
'\0'
} else {
self.input[self.current + 1]
}
}
fn match_char(&mut self, expected: char) -> bool {
if self.is_at_end() || self.input[self.current] != expected {
false
} else {
self.current += 1;
self.position += 1;
true
}
}
fn is_at_end(&self) -> bool {
self.current >= self.input.len()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_basic_tokens() {
let mut lexer = Lexer::new("MATCH (n) RETURN n");
let tokens = lexer.tokenize().unwrap();
assert_eq!(tokens[0].kind, TokenKind::Match);
assert_eq!(tokens[1].kind, TokenKind::LeftParen);
assert_eq!(tokens[2].kind, TokenKind::Identifier("n".to_string()));
assert_eq!(tokens[3].kind, TokenKind::RightParen);
assert_eq!(tokens[4].kind, TokenKind::Return);
assert_eq!(tokens[5].kind, TokenKind::Identifier("n".to_string()));
assert_eq!(tokens[6].kind, TokenKind::Eof);
}
#[test]
fn test_string_literals() {
let mut lexer = Lexer::new("\"hello world\" 'test'");
let tokens = lexer.tokenize().unwrap();
assert_eq!(tokens[0].kind, TokenKind::String("hello world".to_string()));
assert_eq!(tokens[1].kind, TokenKind::String("test".to_string()));
}
#[test]
#[allow(clippy::approx_constant)]
fn test_numbers() {
let mut lexer = Lexer::new("42 3.14 0");
let tokens = lexer.tokenize().unwrap();
assert_eq!(tokens[0].kind, TokenKind::Integer(42));
assert_eq!(tokens[1].kind, TokenKind::Float(3.14));
assert_eq!(tokens[2].kind, TokenKind::Integer(0));
}
#[test]
fn test_operators() {
let mut lexer = Lexer::new("= == <> != < <= > >= -> + - * / %");
let tokens = lexer.tokenize().unwrap();
let expected = vec![
TokenKind::Assign,
TokenKind::Equal,
TokenKind::NotEqual,
TokenKind::NotEqual,
TokenKind::LessThan,
TokenKind::LessEqual,
TokenKind::GreaterThan,
TokenKind::GreaterEqual,
TokenKind::Arrow,
TokenKind::Plus,
TokenKind::Minus,
TokenKind::Multiply,
TokenKind::Divide,
TokenKind::Modulo,
];
for (i, expected_kind) in expected.iter().enumerate() {
assert_eq!(tokens[i].kind, *expected_kind);
}
}
}