use crate::syntax::SyntaxKind;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct LexError {
pub message: String,
pub offset: Option<usize>,
}
impl std::fmt::Display for LexError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self.offset {
Some(o) => write!(f, "{} (at byte {o})", self.message),
None => f.write_str(&self.message),
}
}
}
impl std::error::Error for LexError {}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Token<'a> {
pub kind: SyntaxKind,
pub text: &'a str,
}
const BOM_STR: &str = "\u{FEFF}";
pub fn validate_utf8(bytes: &[u8]) -> Result<&str, LexError> {
match std::str::from_utf8(bytes) {
Ok(s) => Ok(s),
Err(e) => Err(LexError {
message: "input is not valid UTF-8".to_string(),
offset: Some(e.valid_up_to()),
}),
}
}
#[must_use]
pub fn tokenize(src: &str) -> Vec<Token<'_>> {
let mut lexer = Lexer::new(src);
let mut tokens = Vec::new();
while let Some(tok) = lexer.next_token() {
tokens.push(tok);
}
debug_assert_eq!(
tokens.iter().map(|t| t.text.len()).sum::<usize>(),
src.len(),
"lexer must cover every source byte exactly once (INV-1)"
);
tokens
}
struct Lexer<'a> {
src: &'a str,
pos: usize,
at_start: bool,
}
impl<'a> Lexer<'a> {
fn new(src: &'a str) -> Self {
Self {
src,
pos: 0,
at_start: true,
}
}
#[inline]
fn rest(&self) -> &'a str {
&self.src[self.pos..]
}
#[inline]
fn peek(&self) -> Option<char> {
self.rest().chars().next()
}
#[inline]
fn peek_nth(&self, n: usize) -> Option<char> {
self.rest().chars().nth(n)
}
#[inline]
fn emit(&self, kind: SyntaxKind, start: usize) -> Token<'a> {
Token {
kind,
text: &self.src[start..self.pos],
}
}
fn next_token(&mut self) -> Option<Token<'a>> {
if self.pos >= self.src.len() {
return None;
}
let start = self.pos;
if self.at_start && self.rest().starts_with(BOM_STR) {
self.pos += BOM_STR.len();
self.at_start = false;
return Some(self.emit(SyntaxKind::Bom, start));
}
self.at_start = false;
let c = self.peek().expect("non-empty rest has a char");
let token = match c {
c if is_whitespace(c) => self.lex_whitespace(start),
'/' if self.peek_nth(1) == Some('/') => self.lex_line_comment(start),
'/' if self.peek_nth(1) == Some('*') => self.lex_block_comment(start),
'"' => self.lex_string(start),
'r' if matches!(self.peek_nth(1), Some('"') | Some('#')) => self.lex_raw_string(start),
'\'' => self.lex_char(start),
'0'..='9' => self.lex_number(start, false),
'+' | '-' if matches!(self.peek_nth(1), Some('0'..='9') | Some('.')) => {
self.lex_number(start, true)
}
'.' if matches!(self.peek_nth(1), Some('0'..='9')) => self.lex_number(start, false),
c if is_ident_start(c) => self.lex_ident_or_keyword(start),
'(' => self.bump_punct(SyntaxKind::LParen, start),
')' => self.bump_punct(SyntaxKind::RParen, start),
'[' => self.bump_punct(SyntaxKind::LBracket, start),
']' => self.bump_punct(SyntaxKind::RBracket, start),
'{' => self.bump_punct(SyntaxKind::LBrace, start),
'}' => self.bump_punct(SyntaxKind::RBrace, start),
':' => self.bump_punct(SyntaxKind::Colon, start),
',' => self.bump_punct(SyntaxKind::Comma, start),
'#' => self.bump_punct(SyntaxKind::Hash, start),
'!' => self.bump_punct(SyntaxKind::Bang, start),
_ => {
self.pos += c.len_utf8();
self.emit(SyntaxKind::LexError, start)
}
};
Some(token)
}
#[inline]
fn bump_punct(&mut self, kind: SyntaxKind, start: usize) -> Token<'a> {
self.pos += 1;
self.emit(kind, start)
}
fn lex_whitespace(&mut self, start: usize) -> Token<'a> {
while let Some(c) = self.peek() {
if is_whitespace(c) {
self.pos += c.len_utf8();
} else {
break;
}
}
self.emit(SyntaxKind::Whitespace, start)
}
fn lex_line_comment(&mut self, start: usize) -> Token<'a> {
self.pos += 2;
while let Some(c) = self.peek() {
if c == '\n' {
break;
}
self.pos += c.len_utf8();
}
self.emit(SyntaxKind::LineComment, start)
}
fn lex_block_comment(&mut self, start: usize) -> Token<'a> {
self.pos += 2;
let mut depth = 1usize;
while depth > 0 {
let Some(c) = self.peek() else { break };
if c == '/' && self.peek_nth(1) == Some('*') {
self.pos += 2;
depth += 1;
} else if c == '*' && self.peek_nth(1) == Some('/') {
self.pos += 2;
depth -= 1;
} else {
self.pos += c.len_utf8();
}
}
self.emit(SyntaxKind::BlockComment, start)
}
fn lex_string(&mut self, start: usize) -> Token<'a> {
self.pos += 1;
while let Some(c) = self.peek() {
match c {
'\\' => {
self.pos += 1;
if let Some(esc) = self.peek() {
self.pos += esc.len_utf8();
}
}
'"' => {
self.pos += 1;
break;
}
_ => self.pos += c.len_utf8(),
}
}
self.emit(SyntaxKind::String, start)
}
fn lex_raw_string(&mut self, start: usize) -> Token<'a> {
self.pos += 1;
let mut hashes = 0usize;
while self.peek() == Some('#') {
self.pos += 1;
hashes += 1;
}
if self.peek() != Some('"') {
return self.emit(SyntaxKind::RawString, start);
}
self.pos += 1; while let Some(c) = self.peek() {
if c == '"' {
let after_quote = self.pos + 1;
let mut matched = 0usize;
let mut probe = after_quote;
while matched < hashes && self.src[probe..].starts_with('#') {
probe += 1;
matched += 1;
}
if matched == hashes {
self.pos = probe;
break;
}
self.pos += 1;
} else {
self.pos += c.len_utf8();
}
}
self.emit(SyntaxKind::RawString, start)
}
fn lex_char(&mut self, start: usize) -> Token<'a> {
self.pos += 1;
while let Some(c) = self.peek() {
match c {
'\\' => {
self.pos += 1;
if let Some(esc) = self.peek() {
self.pos += esc.len_utf8();
}
}
'\'' => {
self.pos += 1;
break;
}
_ => self.pos += c.len_utf8(),
}
}
self.emit(SyntaxKind::Char, start)
}
fn lex_number(&mut self, start: usize, signed: bool) -> Token<'a> {
if signed {
self.pos += 1; }
if self.peek() == Some('0') {
if let Some(radix) = self.peek_nth(1) {
let base = match radix {
'x' | 'X' => Some(16u32),
'b' | 'B' => Some(2),
'o' | 'O' => Some(8),
_ => None,
};
if let Some(base) = base {
self.pos += 2; self.consume_digits(base);
self.consume_type_suffix();
return self.emit(SyntaxKind::Integer, start);
}
}
}
self.consume_digits(10);
let mut is_float = false;
if self.peek() == Some('.') && self.peek_nth(1) != Some('.') {
is_float = true;
self.pos += 1;
self.consume_digits(10);
}
if matches!(self.peek(), Some('e') | Some('E')) {
let next = self.peek_nth(1);
let exp = matches!(next, Some('0'..='9'))
|| (matches!(next, Some('+') | Some('-'))
&& matches!(self.peek_nth(2), Some('0'..='9')));
if exp {
is_float = true;
self.pos += 1; if matches!(self.peek(), Some('+') | Some('-')) {
self.pos += 1;
}
self.consume_digits(10);
}
}
self.consume_type_suffix();
self.emit(
if is_float {
SyntaxKind::Float
} else {
SyntaxKind::Integer
},
start,
)
}
fn consume_digits(&mut self, base: u32) {
while let Some(c) = self.peek() {
if c == '_' || c.is_digit(base) {
self.pos += 1; } else {
break;
}
}
}
fn consume_type_suffix(&mut self) {
if let Some(c) = self.peek() {
if c.is_ascii_alphabetic() {
while let Some(c) = self.peek() {
if is_ident_continue(c) {
self.pos += c.len_utf8();
} else {
break;
}
}
}
}
}
fn lex_ident_or_keyword(&mut self, start: usize) -> Token<'a> {
while let Some(c) = self.peek() {
if is_ident_continue(c) {
self.pos += c.len_utf8();
} else {
break;
}
}
let text = &self.src[start..self.pos];
let kind = match text {
"true" => SyntaxKind::TrueKw,
"false" => SyntaxKind::FalseKw,
"enable" => SyntaxKind::EnableKw,
_ => SyntaxKind::Ident,
};
Token { kind, text }
}
}
#[inline]
fn is_whitespace(c: char) -> bool {
c.is_whitespace()
}
#[inline]
fn is_ident_start(c: char) -> bool {
c == '_' || c.is_alphabetic()
}
#[inline]
fn is_ident_continue(c: char) -> bool {
c == '_' || c.is_alphanumeric()
}
#[cfg(test)]
mod tests {
use super::*;
fn concat(tokens: &[Token<'_>]) -> String {
tokens.iter().map(|t| t.text).collect()
}
fn kinds(tokens: &[Token<'_>]) -> Vec<SyntaxKind> {
tokens.iter().map(|t| t.kind).collect()
}
#[test]
fn validate_utf8_accepts_valid() {
assert_eq!(validate_utf8(b"hello").unwrap(), "hello");
}
#[test]
fn validate_utf8_rejects_invalid_without_panic() {
let bad = [0xFF, 0xFE, 0x00];
let err = validate_utf8(&bad).unwrap_err();
assert_eq!(err.offset, Some(0));
}
#[test]
fn covers_every_byte() {
let inputs = [
"",
" ",
"// only a comment",
"/* nested /* block */ comment */",
"Foo(x: 1, y: 2.5)",
"[1, 2, 3,]",
"{ \"k\": 'c', 4: true }",
"r#\"raw \"quote\" string\"#",
"Some(())",
"#![enable(implicit_some)]\n42",
"0xFF_u8 0b1010 0o17 1_000.5e-3f64 -1 +2.0",
];
for input in inputs {
let toks = tokenize(input);
assert_eq!(concat(&toks), input, "round-trip for {input:?}");
}
}
#[test]
fn leading_bom_is_trivia() {
let src = "\u{FEFF}1";
let toks = tokenize(src);
assert_eq!(toks[0].kind, SyntaxKind::Bom);
assert_eq!(toks[0].text, "\u{FEFF}");
assert_eq!(concat(&toks), src);
}
#[test]
fn raw_string_with_hashes() {
let src = "r##\"has \"# inside\"##";
let toks = tokenize(src);
assert_eq!(kinds(&toks), vec![SyntaxKind::RawString]);
assert_eq!(toks[0].text, src);
}
#[test]
fn numbers_classified() {
assert_eq!(tokenize("42")[0].kind, SyntaxKind::Integer);
assert_eq!(tokenize("0xFF")[0].kind, SyntaxKind::Integer);
assert_eq!(tokenize("3.14")[0].kind, SyntaxKind::Float);
assert_eq!(tokenize("1e10")[0].kind, SyntaxKind::Float);
assert_eq!(tokenize("1_000i64")[0].kind, SyntaxKind::Integer);
}
#[test]
fn keywords_and_idents() {
assert_eq!(tokenize("true")[0].kind, SyntaxKind::TrueKw);
assert_eq!(tokenize("false")[0].kind, SyntaxKind::FalseKw);
assert_eq!(tokenize("enable")[0].kind, SyntaxKind::EnableKw);
assert_eq!(tokenize("Foo")[0].kind, SyntaxKind::Ident);
}
#[test]
fn unterminated_string_runs_to_eof_without_panic() {
let src = "\"no end";
let toks = tokenize(src);
assert_eq!(toks[0].kind, SyntaxKind::String);
assert_eq!(concat(&toks), src);
}
#[test]
fn crlf_preserved() {
let src = "1\r\n2";
let toks = tokenize(src);
assert_eq!(concat(&toks), src);
}
}