1use ocas_atom::{Atom, AtomArena};
7use thiserror::Error;
8
9use crate::lexer::{LexError, Token, lex};
10
11#[derive(Debug, Error, PartialEq)]
13pub enum ParseError {
14 #[error("lex error")]
16 Lex(#[from] LexError),
17 #[error("unexpected end of input")]
19 UnexpectedEof,
20 #[error("unexpected token")]
22 UnexpectedToken,
23}
24
25pub 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 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 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 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 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 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 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}