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 -> unary (^ factor)?
137    fn factor(&mut self) -> Result<Atom<'a>, ParseError> {
138        let base = self.unary()?;
139        if self.current() == Some(Token::Caret) {
140            self.advance();
141            let exp = self.factor()?;
142            Ok(self.ctx.pow(base, exp))
143        } else {
144            Ok(base)
145        }
146    }
147
148    // unary -> - unary | primary
149    fn unary(&mut self) -> Result<Atom<'a>, ParseError> {
150        if self.current() == Some(Token::Minus) {
151            self.advance();
152            let operand = self.unary()?;
153            Ok(self.ctx.mul(&[self.ctx.num(-1), operand]))
154        } else {
155            self.primary()
156        }
157    }
158
159    // primary -> number | ident | ident ( arg_list ) | ( expr )
160    fn primary(&mut self) -> Result<Atom<'a>, ParseError> {
161        match self.current() {
162            Some(Token::Integer(n)) => {
163                self.advance();
164                Ok(self.ctx.num(n))
165            }
166            Some(Token::Ident(name)) => {
167                let name = name.to_owned();
168                self.advance();
169                if self.current() == Some(Token::LParen) {
170                    self.advance();
171                    let args = if self.current() == Some(Token::RParen) {
172                        Vec::new()
173                    } else {
174                        self.arg_list()?
175                    };
176                    self.expect(Token::RParen)?;
177                    Ok(self.ctx.fun(&name, &args))
178                } else {
179                    Ok(self.ctx.var(&name))
180                }
181            }
182            Some(Token::LParen) => {
183                self.advance();
184                let inner = self.expr()?;
185                self.expect(Token::RParen)?;
186                Ok(inner)
187            }
188            Some(_) => Err(ParseError::UnexpectedToken),
189            None => Err(ParseError::UnexpectedEof),
190        }
191    }
192
193    // arg_list -> expr (',' expr)*
194    fn arg_list(&mut self) -> Result<Vec<Atom<'a>>, ParseError> {
195        let mut args = vec![self.expr()?];
196        while self.current() == Some(Token::Comma) {
197            self.advance();
198            args.push(self.expr()?);
199        }
200        Ok(args)
201    }
202}
203
204#[cfg(test)]
205mod tests {
206    use super::*;
207    use ocas_core::arena::Arena;
208
209    #[test]
210    fn parse_number() {
211        let arena = Arena::new();
212        let ctx = AtomArena::new(&arena);
213        let atom = parse(&ctx, "42").unwrap();
214        assert_eq!(atom.to_string(), "42");
215    }
216
217    #[test]
218    fn parse_variable() {
219        let arena = Arena::new();
220        let ctx = AtomArena::new(&arena);
221        let atom = parse(&ctx, "x").unwrap();
222        assert_eq!(atom.to_string(), "x");
223    }
224
225    #[test]
226    fn parse_addition() {
227        let arena = Arena::new();
228        let ctx = AtomArena::new(&arena);
229        let atom = parse(&ctx, "x + y").unwrap();
230        assert_eq!(atom.to_string(), "x + y");
231    }
232
233    #[test]
234    fn parse_multiplication() {
235        let arena = Arena::new();
236        let ctx = AtomArena::new(&arena);
237        let atom = parse(&ctx, "x * y").unwrap();
238        assert_eq!(atom.to_string(), "x*y");
239    }
240
241    #[test]
242    fn parse_function_call() {
243        let arena = Arena::new();
244        let ctx = AtomArena::new(&arena);
245        let atom = parse(&ctx, "sin(x)").unwrap();
246        assert_eq!(atom.to_string(), "sin(x)");
247    }
248
249    #[test]
250    fn parse_function_call_multiple_args() {
251        let arena = Arena::new();
252        let ctx = AtomArena::new(&arena);
253        let atom = parse(&ctx, "f(x, y, 2)").unwrap();
254        assert_eq!(atom.to_string(), "f(x, y, 2)");
255    }
256
257    #[test]
258    fn parse_function_call_in_expression() {
259        let arena = Arena::new();
260        let ctx = AtomArena::new(&arena);
261        let atom = parse(&ctx, "sin(x) + cos(x)").unwrap();
262        assert_eq!(atom.to_string(), "(sin(x)) + (cos(x))");
263    }
264
265    #[test]
266    fn parse_operator_precedence() {
267        let arena = Arena::new();
268        let ctx = AtomArena::new(&arena);
269        let atom = parse(&ctx, "x + 2 * y").unwrap();
270        assert_eq!(atom.to_string(), "x + (2*y)");
271    }
272
273    #[test]
274    fn parse_right_associative_power() {
275        let arena = Arena::new();
276        let ctx = AtomArena::new(&arena);
277        let atom = parse(&ctx, "2 ^ 3 ^ 2").unwrap();
278        assert_eq!(atom.to_string(), "2^(3^2)");
279    }
280
281    #[test]
282    fn parse_parentheses() {
283        let arena = Arena::new();
284        let ctx = AtomArena::new(&arena);
285        let atom = parse(&ctx, "(x + y) * z").unwrap();
286        assert_eq!(atom.to_string(), "(x + y)*z");
287    }
288
289    #[test]
290    fn parse_unary_minus() {
291        let arena = Arena::new();
292        let ctx = AtomArena::new(&arena);
293        let atom = parse(&ctx, "-x").unwrap();
294        assert_eq!(atom.to_string(), "-1*x");
295    }
296
297    #[test]
298    fn parse_subtraction_normalized() {
299        let arena = Arena::new();
300        let ctx = AtomArena::new(&arena);
301        let atom = parse(&ctx, "x - y").unwrap();
302        assert_eq!(atom.to_string(), "x + (-1*y)");
303    }
304
305    #[test]
306    fn parse_division_normalized() {
307        let arena = Arena::new();
308        let ctx = AtomArena::new(&arena);
309        let atom = parse(&ctx, "x / y").unwrap();
310        assert_eq!(atom.to_string(), "x*(y^-1)");
311    }
312
313    #[test]
314    fn parse_polynomial_like() {
315        let arena = Arena::new();
316        let ctx = AtomArena::new(&arena);
317        let atom = parse(&ctx, "x^2 + 2*x + 1").unwrap();
318        assert_eq!(atom.to_string(), "((x^2) + (2*x)) + 1");
319    }
320
321    #[test]
322    fn parse_rejects_invalid_input() {
323        let arena = Arena::new();
324        let ctx = AtomArena::new(&arena);
325        assert!(parse(&ctx, "x +").is_err());
326    }
327}