Skip to main content

ocas_parse/
parser.rs

1//! Recursive-descent parser for oCAS expressions.
2//!
3//! Parses the token stream produced by [`crate::lexer`] into an
4//! [`ocas_atom::Atom`] expression tree allocated in the provided arena.
5
6use ocas_atom::{Atom, AtomArena};
7use thiserror::Error;
8
9use crate::lexer::{LexError, Token, lex};
10
11/// Errors that can occur while parsing an expression.
12#[derive(Debug, Error, PartialEq)]
13pub enum ParseError {
14    /// The input could not be lexed.
15    #[error("lex error")]
16    Lex(#[from] LexError),
17    /// Unexpected end of input.
18    #[error("unexpected end of input")]
19    UnexpectedEof,
20    /// An unexpected token was encountered.
21    #[error("unexpected token")]
22    UnexpectedToken,
23}
24
25/// Parse an expression string into an [`Atom`].
26///
27/// # Errors
28///
29/// Returns a [`ParseError`] if the input cannot be lexed or parsed.
30///
31/// # Example
32///
33/// ```
34/// use ocas_atom::AtomArena;
35/// use ocas_core::arena::Arena;
36/// use ocas_parse::parse;
37///
38/// let arena = Arena::new();
39/// let ctx = AtomArena::new(&arena);
40/// let expr = parse(&ctx, "x^2 + 2*x + 1").unwrap();
41/// assert_eq!(expr.to_string(), "((x^2) + (2*x)) + 1");
42/// ```
43pub fn parse<'a>(ctx: &'a AtomArena<'a>, input: &'a str) -> Result<Atom<'a>, ParseError> {
44    let tokens = lex(input)?;
45    let mut parser = Parser::new(ctx, &tokens);
46    parser.parse()
47}
48
49struct Parser<'a, 'tokens> {
50    ctx: &'a AtomArena<'a>,
51    tokens: &'tokens [Token<'tokens>],
52    pos: usize,
53}
54
55impl<'a, 'tokens> Parser<'a, 'tokens> {
56    fn new(ctx: &'a AtomArena<'a>, tokens: &'tokens [Token<'tokens>]) -> Self {
57        Self {
58            ctx,
59            tokens,
60            pos: 0,
61        }
62    }
63
64    fn parse(&mut self) -> Result<Atom<'a>, ParseError> {
65        let expr = self.expr()?;
66        self.expect(Token::Eof)?;
67        Ok(expr)
68    }
69
70    fn current(&self) -> Option<Token<'_>> {
71        self.tokens.get(self.pos).copied()
72    }
73
74    fn advance(&mut self) -> Token<'_> {
75        let token = self.tokens[self.pos];
76        self.pos += 1;
77        token
78    }
79
80    fn expect(&mut self, expected: Token) -> Result<(), ParseError> {
81        match self.current() {
82            Some(token) if token == expected => {
83                self.advance();
84                Ok(())
85            }
86            Some(_) => Err(ParseError::UnexpectedToken),
87            None => Err(ParseError::UnexpectedEof),
88        }
89    }
90
91    // expr -> term ((+|-) term)*
92    fn expr(&mut self) -> Result<Atom<'a>, ParseError> {
93        let mut left = self.term()?;
94        while let Some(token) = self.current() {
95            match token {
96                Token::Plus => {
97                    self.advance();
98                    let right = self.term()?;
99                    left = self.ctx.add(&[left, right]);
100                }
101                Token::Minus => {
102                    self.advance();
103                    let right = self.term()?;
104                    let neg_right = self.ctx.mul(&[self.ctx.num(-1), right]);
105                    left = self.ctx.add(&[left, neg_right]);
106                }
107                _ => break,
108            }
109        }
110        Ok(left)
111    }
112
113    // term -> factor ((*|/) factor)*
114    fn term(&mut self) -> Result<Atom<'a>, ParseError> {
115        let mut left = self.factor()?;
116        while let Some(token) = self.current() {
117            match token {
118                Token::Star => {
119                    self.advance();
120                    let right = self.factor()?;
121                    left = self.ctx.mul(&[left, right]);
122                }
123                Token::Slash => {
124                    self.advance();
125                    let right = self.factor()?;
126                    let neg_one = self.ctx.num(-1);
127                    let inv_right = self.ctx.pow(right, neg_one);
128                    left = self.ctx.mul(&[left, inv_right]);
129                }
130                _ => break,
131            }
132        }
133        Ok(left)
134    }
135
136    // factor -> primary (^ factor)? with a leading minus binding looser
137    // than exponentiation: -x^2 = -(x^2).
138    fn factor(&mut self) -> Result<Atom<'a>, ParseError> {
139        if self.current() == Some(Token::Minus) {
140            self.advance();
141            let operand = self.factor()?;
142            return Ok(self.ctx.mul(&[self.ctx.num(-1), operand]));
143        }
144        let base = self.primary()?;
145        if self.current() == Some(Token::Caret) {
146            self.advance();
147            let exp = self.factor()?;
148            Ok(self.ctx.pow(base, exp))
149        } else {
150            Ok(base)
151        }
152    }
153
154    // primary -> number | ident | ident ( arg_list ) | ( expr )
155    fn primary(&mut self) -> Result<Atom<'a>, ParseError> {
156        match self.current() {
157            Some(Token::Integer(n)) => {
158                self.advance();
159                Ok(self.ctx.num(n))
160            }
161            Some(Token::Ident(name)) => {
162                let name = name.to_owned();
163                self.advance();
164                if self.current() == Some(Token::LParen) {
165                    self.advance();
166                    let args = if self.current() == Some(Token::RParen) {
167                        Vec::new()
168                    } else {
169                        self.arg_list()?
170                    };
171                    self.expect(Token::RParen)?;
172                    Ok(self.ctx.fun(&name, &args))
173                } else {
174                    Ok(self.ctx.var(&name))
175                }
176            }
177            Some(Token::LParen) => {
178                self.advance();
179                let inner = self.expr()?;
180                self.expect(Token::RParen)?;
181                Ok(inner)
182            }
183            Some(_) => Err(ParseError::UnexpectedToken),
184            None => Err(ParseError::UnexpectedEof),
185        }
186    }
187
188    // arg_list -> expr (',' expr)*
189    fn arg_list(&mut self) -> Result<Vec<Atom<'a>>, ParseError> {
190        let mut args = vec![self.expr()?];
191        while self.current() == Some(Token::Comma) {
192            self.advance();
193            args.push(self.expr()?);
194        }
195        Ok(args)
196    }
197}
198
199#[cfg(test)]
200mod tests {
201    use super::*;
202    use ocas_core::arena::Arena;
203
204    #[test]
205    fn parse_number() {
206        let arena = Arena::new();
207        let ctx = AtomArena::new(&arena);
208        let atom = parse(&ctx, "42").unwrap();
209        assert_eq!(atom.to_string(), "42");
210    }
211
212    #[test]
213    fn parse_variable() {
214        let arena = Arena::new();
215        let ctx = AtomArena::new(&arena);
216        let atom = parse(&ctx, "x").unwrap();
217        assert_eq!(atom.to_string(), "x");
218    }
219
220    #[test]
221    fn parse_addition() {
222        let arena = Arena::new();
223        let ctx = AtomArena::new(&arena);
224        let atom = parse(&ctx, "x + y").unwrap();
225        assert_eq!(atom.to_string(), "x + y");
226    }
227
228    #[test]
229    fn parse_multiplication() {
230        let arena = Arena::new();
231        let ctx = AtomArena::new(&arena);
232        let atom = parse(&ctx, "x * y").unwrap();
233        assert_eq!(atom.to_string(), "x*y");
234    }
235
236    #[test]
237    fn parse_function_call() {
238        let arena = Arena::new();
239        let ctx = AtomArena::new(&arena);
240        let atom = parse(&ctx, "sin(x)").unwrap();
241        assert_eq!(atom.to_string(), "sin(x)");
242    }
243
244    #[test]
245    fn parse_function_call_multiple_args() {
246        let arena = Arena::new();
247        let ctx = AtomArena::new(&arena);
248        let atom = parse(&ctx, "f(x, y, 2)").unwrap();
249        assert_eq!(atom.to_string(), "f(x, y, 2)");
250    }
251
252    #[test]
253    fn parse_function_call_in_expression() {
254        let arena = Arena::new();
255        let ctx = AtomArena::new(&arena);
256        let atom = parse(&ctx, "sin(x) + cos(x)").unwrap();
257        assert_eq!(atom.to_string(), "(sin(x)) + (cos(x))");
258    }
259
260    #[test]
261    fn parse_operator_precedence() {
262        let arena = Arena::new();
263        let ctx = AtomArena::new(&arena);
264        let atom = parse(&ctx, "x + 2 * y").unwrap();
265        assert_eq!(atom.to_string(), "x + (2*y)");
266    }
267
268    #[test]
269    fn parse_right_associative_power() {
270        let arena = Arena::new();
271        let ctx = AtomArena::new(&arena);
272        let atom = parse(&ctx, "2 ^ 3 ^ 2").unwrap();
273        assert_eq!(atom.to_string(), "2^(3^2)");
274    }
275
276    #[test]
277    fn parse_parentheses() {
278        let arena = Arena::new();
279        let ctx = AtomArena::new(&arena);
280        let atom = parse(&ctx, "(x + y) * z").unwrap();
281        assert_eq!(atom.to_string(), "(x + y)*z");
282    }
283
284    #[test]
285    fn parse_unary_minus() {
286        let arena = Arena::new();
287        let ctx = AtomArena::new(&arena);
288        let atom = parse(&ctx, "-x").unwrap();
289        assert_eq!(atom.to_string(), "-1*x");
290    }
291
292    #[test]
293    fn parse_subtraction_normalized() {
294        let arena = Arena::new();
295        let ctx = AtomArena::new(&arena);
296        let atom = parse(&ctx, "x - y").unwrap();
297        assert_eq!(atom.to_string(), "x + (-1*y)");
298    }
299
300    #[test]
301    fn parse_division_normalized() {
302        let arena = Arena::new();
303        let ctx = AtomArena::new(&arena);
304        let atom = parse(&ctx, "x / y").unwrap();
305        assert_eq!(atom.to_string(), "x*(y^-1)");
306    }
307
308    #[test]
309    fn parse_polynomial_like() {
310        let arena = Arena::new();
311        let ctx = AtomArena::new(&arena);
312        let atom = parse(&ctx, "x^2 + 2*x + 1").unwrap();
313        assert_eq!(atom.to_string(), "((x^2) + (2*x)) + 1");
314    }
315
316    #[test]
317    fn parse_rejects_invalid_input() {
318        let arena = Arena::new();
319        let ctx = AtomArena::new(&arena);
320        assert!(parse(&ctx, "x +").is_err());
321    }
322}