diff --git a/src/parser.rs b/src/parser.rs index 8335d7f..7eda6f6 100644 --- a/src/parser.rs +++ b/src/parser.rs @@ -6,7 +6,7 @@ use crate::{ lexer::{Lexer, token::Token}, parser::{ ast::{Line, PrintItem, Program, Statement}, - expression::Expression, + expression::{BinaryOperator, Expression}, }, }; pub mod ast; @@ -110,12 +110,42 @@ impl<'a> Parser<'a> { } fn parse_expression(&mut self) -> Expression { + self.parse_binary_expression(0) + } + + fn parse_binary_expression(&mut self, minimum_precedence: u8) -> Expression { + let mut left = self.parse_primary_expression(); + + while let Some(operator) = BinaryOperator::from_token(&self.lexer.peek()) { + let precedence = operator.precedence(); + if precedence < minimum_precedence { + break; + } + + self.lexer.next_token(); + let right = self.parse_binary_expression(precedence + 1); + left = Expression::binary(left, operator, right); + } + + left + } + + fn parse_primary_expression(&mut self) -> Expression { match self.lexer.next_token() { Token::Integer(value) => Expression::integer(value), Token::Float(value) => Expression::float(value), Token::String(value) => Expression::string(value), Token::Char(value) => Expression::char(value), Token::Identifier(value) => Expression::identifier(value), + Token::LParen => { + let expression = self.parse_expression(); + assert_eq!( + self.lexer.next_token(), + Token::RParen, + "Expected right paren" + ); + Expression::grouping(expression) + } token => panic!("Expected expression, got {token:?}"), } } @@ -126,7 +156,7 @@ mod tests { use crate::parser::{ Parser, ast::{Line, PrintItem, Program, Statement}, - expression::Expression, + expression::{BinaryOperator, Expression}, }; fn parse_statement_list(input: &'static str) -> Vec { @@ -198,6 +228,62 @@ mod tests { ); } + #[test] + fn parse_add_expression_statement_list() { + assert_eq!( + parse_statement_list("1 + 2;"), + vec![Statement::Expression(Expression::binary( + Expression::integer(1), + BinaryOperator::Plus, + Expression::integer(2), + ))] + ); + } + + #[test] + fn parse_binary_expression_precedence() { + assert_eq!( + parse_statement_list("1 + 2 * 3;"), + vec![Statement::Expression(Expression::binary( + Expression::integer(1), + BinaryOperator::Plus, + Expression::binary( + Expression::integer(2), + BinaryOperator::Multiply, + Expression::integer(3), + ), + ))] + ); + } + + #[test] + fn parse_grouped_binary_expression() { + assert_eq!( + parse_statement_list("(1 + 2) * 3;"), + vec![Statement::Expression(Expression::binary( + Expression::grouping(Expression::binary( + Expression::integer(1), + BinaryOperator::Plus, + Expression::integer(2), + )), + BinaryOperator::Multiply, + Expression::integer(3), + ))] + ); + } + + #[test] + fn parse_comparison_expression() { + assert_eq!( + parse_statement_list("1 <= 2;"), + vec![Statement::Expression(Expression::binary( + Expression::integer(1), + BinaryOperator::LessEqual, + Expression::integer(2), + ))] + ); + } + #[test] fn parse_return_statement_list() { assert_eq!(parse_statement_list("return;"), vec![Statement::Return]); @@ -404,4 +490,67 @@ mod tests { } ); } + + #[test] + fn parse_program_from_basic_math_fixture() { + let mut parser = Parser::from_input(include_str!("../tests/input/basic_math.bsc")); + + assert_eq!( + parser.parse_program(), + Program { + lines: vec![Line { + number: Some(1), + statements: vec![Statement::Expression(Expression::binary( + Expression::integer(1), + BinaryOperator::Plus, + Expression::integer(2), + ))] + }] + } + ); + } + + #[test] + fn parse_program_from_complicated_expression_fixture() { + let mut parser = + Parser::from_input(include_str!("../tests/input/complicated_expression.bsc")); + + assert_eq!( + parser.parse_program(), + Program { + lines: vec![Line { + number: None, + statements: vec![Statement::Expression(Expression::binary( + Expression::grouping(Expression::binary( + Expression::binary( + Expression::grouping(Expression::binary( + Expression::integer(1), + BinaryOperator::Plus, + Expression::integer(2), + )), + BinaryOperator::Multiply, + Expression::grouping(Expression::binary( + Expression::integer(3), + BinaryOperator::Plus, + Expression::integer(4), + )), + ), + BinaryOperator::Minus, + Expression::binary( + Expression::integer(5), + BinaryOperator::Divide, + Expression::grouping(Expression::binary( + Expression::integer(6), + BinaryOperator::Minus, + Expression::integer(1), + )), + ), + )), + BinaryOperator::GreaterEqual, + Expression::integer(20), + ))] + }] + } + ); + } } diff --git a/tests/input/basic_math.bsc b/tests/input/basic_math.bsc index 6912615..0e92774 100644 --- a/tests/input/basic_math.bsc +++ b/tests/input/basic_math.bsc @@ -1 +1 @@ -1; \ No newline at end of file +1 + 2; diff --git a/tests/input/complicated_expression.bsc b/tests/input/complicated_expression.bsc new file mode 100644 index 0000000..43acdab --- /dev/null +++ b/tests/input/complicated_expression.bsc @@ -0,0 +1 @@ +((1 + 2) * (3 + 4) - 5 / (6 - 1)) >= 20;