Skip to main content

ocas_parse/
lexer.rs

1//! Lexer for oCAS expression syntax.
2//!
3//! The lexer recognises a minimal CAS language: integers, identifiers,
4//! arithmetic operators `+ - * / ^`, parentheses, and commas.
5
6use logos::Logos;
7
8/// A token in the oCAS expression language.
9#[derive(Debug, Clone, Copy, PartialEq, Eq, Logos)]
10#[logos(skip r"[ \t\n\r]+")]
11#[logos(error = LexError)]
12pub enum Token<'a> {
13    /// An integer literal, e.g. `42` or `-7`.
14    #[regex(r"-?[0-9]+", |lex| lex.slice().parse::<i64>())]
15    Integer(i64),
16
17    /// An identifier (variable or function name), e.g. `x` or `sin`.
18    #[regex(r"[a-zA-Z_][a-zA-Z0-9_]*", |lex| lex.slice())]
19    Ident(&'a str),
20
21    /// `+`
22    #[token("+")]
23    Plus,
24
25    /// `-`
26    #[token("-")]
27    Minus,
28
29    /// `*`
30    #[token("*")]
31    Star,
32
33    /// `/`
34    #[token("/")]
35    Slash,
36
37    /// `^`
38    #[token("^")]
39    Caret,
40
41    /// `(`
42    #[token("(")]
43    LParen,
44
45    /// `)`
46    #[token(")")]
47    RParen,
48
49    /// `,`
50    #[token(",")]
51    Comma,
52
53    /// End-of-file sentinel added by [`lex`].
54    Eof,
55}
56
57/// Lex an input string into a vector of tokens.
58///
59/// # Errors
60///
61/// Returns the first lexing error encountered.
62pub 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/// Lexing error produced when input does not match any known token.
72#[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}