Skip to main content

math_parser_rs/dsl/
parser.rs

1use super::{
2    ast_types::{AstNode, OperatorKind, ParserError},
3    optimizations::{general::constant_folding, trig::rewrite_cos_as_sin},
4    types::{LexResult, Token, TokenType},
5};
6
7pub struct Parser {
8    tokens: Vec<Token>,
9    pos: usize,
10}
11
12impl Parser {
13    pub fn new(tokens: LexResult) -> Parser {
14        Parser {
15            pos: 0,
16            tokens: tokens
17                .rhs
18                .into_iter()
19                .filter(|tok| tok.token_type != TokenType::Unimplemented)
20                .collect(),
21        }
22    }
23
24    fn optimize_ast(&self, mut ast: AstNode) -> AstNode {
25        ast = rewrite_cos_as_sin(&ast);
26        ast = constant_folding(&ast);
27        ast
28    }
29
30    pub fn parse(&mut self) -> Result<AstNode, ParserError> {
31        let ast = self.parse_expression()?;
32        Ok(self.optimize_ast(ast))
33    }
34
35    /// Handlest lowest-precedance operations like `+` and `-`
36    fn parse_expression(&mut self) -> Result<AstNode, ParserError> {
37        let mut node = self.parse_term()?;
38
39        while self.current_token_is(TokenType::Operator, Some("+"))
40            || self.current_token_is(TokenType::Operator, Some("-"))
41        {
42            let op = if self.current_token_is(TokenType::Operator, Some("+")) {
43                self.advance();
44                OperatorKind::Plus
45            } else {
46                self.advance();
47                OperatorKind::Minus
48            };
49
50            let right = self.parse_term()?;
51
52            node = AstNode::BinaryOp {
53                op,
54                left: Box::new(node),
55                right: Box::new(right),
56            };
57        }
58        Ok(node)
59    }
60    /// Handlest higer-precedance operations like `*` and `/`
61    fn parse_term(&mut self) -> Result<AstNode, ParserError> {
62        let mut node = self.parse_factor()?;
63
64        while self.current_token_is(TokenType::Operator, Some("*"))
65            || self.current_token_is(TokenType::Operator, Some("/"))
66        {
67            let op = if self.current_token_is(TokenType::Operator, Some("*")) {
68                self.advance();
69                OperatorKind::Multiply
70            } else {
71                self.advance();
72                OperatorKind::Divide
73            };
74
75            let right = self.parse_factor()?;
76
77            node = AstNode::BinaryOp {
78                op,
79                left: Box::new(node),
80                right: Box::new(right),
81            };
82        }
83        Ok(node)
84    }
85    /// Handles number, parantheses
86    fn parse_factor(&mut self) -> Result<AstNode, ParserError> {
87        let current = self.current();
88        match current {
89            Some(tok) if self.current_token_is(TokenType::Number, None) => {
90                let number = tok.text.parse::<f64>();
91                match number {
92                    Ok(number) => {
93                        self.advance();
94                        Ok(AstNode::Number(number))
95                    }
96                    Err(_) => Err(ParserError {
97                        message: "Could not convert number to float".to_string(),
98                    }),
99                }
100            }
101            Some(tok) if self.current_token_is(TokenType::Identifier, None) => {
102                let var_name = tok.text.clone();
103                if let Some(next) = self.peek() {
104                    if next.token_type == TokenType::ParaOpen {
105                        self.advance(); // consume identifier
106                        self.advance(); // consume '('
107                        let mut args = Vec::new();
108                        if !self.current_token_is(TokenType::ParaClose, None) {
109                            loop {
110                                let expr = self.parse_expression()?;
111                                args.push(expr);
112
113                                if self.current_token_is(TokenType::Comma, None) {
114                                    self.advance(); // consume ','
115                                    if self.current_token_is(TokenType::ParaClose, None) {
116                                        return Err(ParserError {
117                                                       message: "Unexpected ',' before ')' in function argument list".to_string(),
118                                                   });
119                                    }
120                                } else if self.current_token_is(TokenType::ParaClose, None) {
121                                    break;
122                                } else {
123                                    return Err(ParserError {
124                                        message: "Expected ',' or ')' in function argument list"
125                                            .to_string(),
126                                    });
127                                }
128                            }
129                        }
130                        self.eat(TokenType::ParaClose, None);
131                        Ok(AstNode::FunctionCall {
132                            name: var_name,
133                            args,
134                        })
135                    } else {
136                        self.advance();
137                        Ok(AstNode::Variable(var_name))
138                    }
139                } else {
140                    self.advance();
141                    Ok(AstNode::Variable(var_name))
142                }
143            }
144            Some(_) if self.current_token_is(TokenType::ParaOpen, None) => {
145                self.advance();
146                let expr = self.parse_expression()?;
147                if self.eat(TokenType::ParaClose, None) {
148                    Ok(expr)
149                } else {
150                    Err(ParserError {
151                        message: "Expected ')'".to_string(),
152                    })
153                }
154            }
155            Some(tok) => Err(ParserError {
156                message: format!(
157                    "Unexpected token: {:?} '{}', expected number or '('",
158                    tok.token_type, tok.text
159                ),
160            }),
161            None => Err(ParserError {
162                message: "Unexpected end of input".to_string(),
163            }),
164        }
165    }
166
167    /// Pops and returns the current token, advancing the position.
168    fn advance(&mut self) -> Option<&Token> {
169        if self.pos < self.tokens.len() {
170            let token = &self.tokens[self.pos];
171            self.pos += 1;
172            Some(token)
173        } else {
174            None
175        }
176    }
177
178    /// Returns a reference to the next token without consuming it.
179    fn peek(&self) -> Option<&Token> {
180        if self.pos + 1 < self.tokens.len() {
181            let token = &self.tokens[self.pos + 1];
182            Some(token)
183        } else {
184            None
185        }
186    }
187
188    /// Returns a reference to the current token.
189    fn current(&self) -> Option<&Token> {
190        if self.pos < self.tokens.len() {
191            let token = &self.tokens[self.pos];
192            Some(token)
193        } else {
194            None
195        }
196    }
197
198    fn current_token_is(&self, ttype: TokenType, text: Option<&str>) -> bool {
199        let current_token = self.current();
200        match current_token {
201            Some(token) => {
202                if token.token_type != ttype {
203                    return false;
204                }
205                match text {
206                    Some(text) => text == token.text,
207                    None => true,
208                }
209            }
210            None => false,
211        }
212    }
213
214    fn eat(&mut self, ttype: TokenType, text: Option<&str>) -> bool {
215        let is_match = self.current_token_is(ttype, text);
216        if is_match {
217            self.advance();
218        }
219        is_match
220    }
221}