1use logos::Logos;
7
8#[derive(Debug, Clone, Copy, PartialEq, Eq, Logos)]
10#[logos(skip r"[ \t\n\r]+")]
11#[logos(error = LexError)]
12pub enum Token<'a> {
13 #[regex(r"-?[0-9]+", |lex| lex.slice().parse::<i64>())]
15 Integer(i64),
16
17 #[regex(r"[a-zA-Z_][a-zA-Z0-9_]*", |lex| lex.slice())]
19 Ident(&'a str),
20
21 #[token("+")]
23 Plus,
24
25 #[token("-")]
27 Minus,
28
29 #[token("*")]
31 Star,
32
33 #[token("/")]
35 Slash,
36
37 #[token("^")]
39 Caret,
40
41 #[token("(")]
43 LParen,
44
45 #[token(")")]
47 RParen,
48
49 #[token(",")]
51 Comma,
52
53 Eof,
55}
56
57pub fn lex(input: &str) -> Result<Vec<Token<'_>>, LexError> {
63 Token::lexer(input)
64 .collect::<Result<Vec<_>, _>>()
65 .map(|mut tokens| {
66 tokens.push(Token::Eof);
67 tokens
68 })
69}
70
71#[derive(Debug, Default, Clone, Copy, PartialEq, Eq)]
73pub struct LexError;
74
75impl std::fmt::Display for LexError {
76 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
77 write!(f, "invalid token")
78 }
79}
80
81impl std::error::Error for LexError {}
82
83impl From<std::num::ParseIntError> for LexError {
84 fn from(_: std::num::ParseIntError) -> Self {
85 LexError
86 }
87}
88
89#[cfg(test)]
90mod tests {
91 use super::*;
92
93 #[test]
94 fn lex_integer() {
95 let tokens = lex("42").unwrap();
96 assert_eq!(tokens, vec![Token::Integer(42), Token::Eof]);
97 }
98
99 #[test]
100 fn lex_negative_integer() {
101 let tokens = lex("-7").unwrap();
102 assert_eq!(tokens, vec![Token::Integer(-7), Token::Eof]);
103 }
104
105 #[test]
106 fn lex_identifier() {
107 let tokens = lex("x").unwrap();
108 assert_eq!(tokens, vec![Token::Ident("x"), Token::Eof]);
109 }
110
111 #[test]
112 fn lex_operators() {
113 let tokens = lex("+ - * / ^").unwrap();
114 assert_eq!(
115 tokens,
116 vec![
117 Token::Plus,
118 Token::Minus,
119 Token::Star,
120 Token::Slash,
121 Token::Caret,
122 Token::Eof
123 ]
124 );
125 }
126
127 #[test]
128 fn lex_punctuation() {
129 let tokens = lex("(, )").unwrap();
130 assert_eq!(
131 tokens,
132 vec![Token::LParen, Token::Comma, Token::RParen, Token::Eof]
133 );
134 }
135
136 #[test]
137 fn lex_expression() {
138 let tokens = lex("x + 2*y^3").unwrap();
139 assert_eq!(
140 tokens,
141 vec![
142 Token::Ident("x"),
143 Token::Plus,
144 Token::Integer(2),
145 Token::Star,
146 Token::Ident("y"),
147 Token::Caret,
148 Token::Integer(3),
149 Token::Eof
150 ]
151 );
152 }
153
154 #[test]
155 fn lex_skips_whitespace() {
156 let tokens = lex(" x \n\t+ 1 ").unwrap();
157 assert_eq!(
158 tokens,
159 vec![
160 Token::Ident("x"),
161 Token::Plus,
162 Token::Integer(1),
163 Token::Eof
164 ]
165 );
166 }
167
168 #[test]
169 fn lex_rejects_invalid_input() {
170 assert!(lex("@").is_err());
171 }
172}