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> {
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 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 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}