diff --git a/src/lexer/mod.rs b/src/lexer/mod.rs index a086cf1..78c1021 100644 --- a/src/lexer/mod.rs +++ b/src/lexer/mod.rs @@ -15,6 +15,8 @@ pub enum Token { Percent, PlusEqual, MinusEqual, + PlusPlus, + MinusMinus, StarEqual, SlashEqual, PercentEqual, @@ -51,20 +53,30 @@ pub fn tokenize(source: &str) -> Result, String> { match ch { '+' => { chars.next(); - if chars.peek() == Some(&'=') { - chars.next(); - tokens.push(Token::PlusEqual); - } else { - tokens.push(Token::Plus); + match chars.peek() { + Some(&'=') => { + chars.next(); + tokens.push(Token::PlusEqual); + } + Some(&'+') => { + chars.next(); + tokens.push(Token::PlusPlus); + } + _ => tokens.push(Token::Plus), } } '-' => { chars.next(); - if chars.peek() == Some(&'=') { - chars.next(); - tokens.push(Token::MinusEqual); - } else { - tokens.push(Token::Minus); + match chars.peek() { + Some(&'=') => { + chars.next(); + tokens.push(Token::MinusEqual); + } + Some(&'-') => { + chars.next(); + tokens.push(Token::MinusMinus); + } + _ => tokens.push(Token::Minus), } } '*' => { @@ -414,6 +426,23 @@ mod tests { ); } + #[test] + fn tokenizes_increment_and_decrement_operators() { + let tokens = tokenize("++ -- + + - -").expect("tokenization should succeed"); + + assert_eq!( + tokens, + vec![ + Token::PlusPlus, + Token::MinusMinus, + Token::Plus, + Token::Plus, + Token::Minus, + Token::Minus, + ] + ); + } + #[test] fn tokenizes_percent_operator() { let tokens = tokenize("7 % 3").expect("tokenization should succeed"); diff --git a/src/parser/mod.rs b/src/parser/mod.rs index f398a89..81f5514 100644 --- a/src/parser/mod.rs +++ b/src/parser/mod.rs @@ -13,7 +13,10 @@ pub enum ParseError { found: Option, }, #[allow(dead_code)] - TrailingTokens { found: Token }, + TrailingTokens { + found: Token, + }, + InvalidIncrementTarget, } impl std::fmt::Display for ParseError { @@ -25,6 +28,9 @@ impl std::fmt::Display for ParseError { Self::TrailingTokens { found } => { write!(f, "unexpected trailing token after program: {found:?}") } + Self::InvalidIncrementTarget => { + write!(f, "`++`/`--` target must be a variable") + } } } } @@ -189,9 +195,7 @@ impl<'a> Parser<'a> { let post = if self.peek() == Some(&Token::RightParen) { None } else { - let target = self.parse_unary()?; - let assign = self.parse_assignment(target)?; - Some(Box::new(Statement::Assign(assign))) + Some(Box::new(self.parse_assignment_statement()?)) }; self.expect_right_paren()?; let body = self.parse_statement()?; @@ -203,10 +207,9 @@ impl<'a> Parser<'a> { })) } Some(_) => { - let target = self.parse_unary()?; - let assign = self.parse_assignment(target)?; + let statement = self.parse_assignment_statement()?; self.expect_semicolon()?; - Ok(Statement::Assign(assign)) + Ok(statement) } None => Err(ParseError::UnexpectedToken { expected: "statement (return, variable declaration, assignment, block, if, while, or for)", @@ -215,6 +218,44 @@ impl<'a> Parser<'a> { } } + /// Parses an assignment-like statement without its trailing semicolon: + /// a prefix/postfix increment/decrement or a (compound) assignment. + fn parse_assignment_statement(&mut self) -> Result { + if let Some(op) = self.peek_increment_op() { + self.next(); + let target = self.parse_unary()?; + return Ok(Statement::Assign(Self::increment_assignment(target, op)?)); + } + + let target = self.parse_unary()?; + if let Some(op) = self.peek_increment_op() { + self.next(); + return Ok(Statement::Assign(Self::increment_assignment(target, op)?)); + } + + Ok(Statement::Assign(self.parse_assignment(target)?)) + } + + fn peek_increment_op(&self) -> Option { + match self.peek() { + Some(Token::PlusPlus) => Some(BinaryOp::Add), + Some(Token::MinusMinus) => Some(BinaryOp::Subtract), + _ => None, + } + } + + /// Desugars `++x` / `x++` / `--x` / `x--` into `x += 1` / `x -= 1`. + fn increment_assignment(target: Expr, op: BinaryOp) -> Result { + if !matches!(target, Expr::Variable(_)) { + return Err(ParseError::InvalidIncrementTarget); + } + Ok(VarAssignStatement { + target, + op: Some(op), + expr: Expr::IntegerLiteral(IntegerLiteral { value: 1 }), + }) + } + /// Parses the `= expr` / `op= expr` tail of an assignment statement, /// given the already-parsed assignment target. fn parse_assignment(&mut self, target: Expr) -> Result { @@ -849,6 +890,68 @@ mod tests { ); } + #[test] + fn parses_increment_and_decrement_statements() { + let tokens = vec![ + Token::Int, + Token::Identifier("main".to_string()), + Token::LeftParen, + Token::RightParen, + Token::LeftBrace, + Token::Int, + Token::Identifier("x".to_string()), + Token::Equals, + Token::Integer(5), + Token::Semicolon, + Token::Identifier("x".to_string()), + Token::PlusPlus, + Token::Semicolon, + Token::MinusMinus, + Token::Identifier("x".to_string()), + Token::Semicolon, + Token::Return, + Token::Identifier("x".to_string()), + Token::Semicolon, + Token::RightBrace, + ]; + + let program = parse(&tokens).expect("parser should accept increment statements"); + + let expected_incdec = |op| { + Statement::Assign(VarAssignStatement { + target: Expr::Variable(VariableExpr { + name: "x".to_string(), + }), + op: Some(op), + expr: Expr::IntegerLiteral(IntegerLiteral { value: 1 }), + }) + }; + assert_eq!(program.functions[0].body[1], expected_incdec(BinaryOp::Add)); + assert_eq!( + program.functions[0].body[2], + expected_incdec(BinaryOp::Subtract) + ); + } + + #[test] + fn rejects_increment_of_non_variable() { + let tokens = vec![ + Token::Int, + Token::Identifier("main".to_string()), + Token::LeftParen, + Token::RightParen, + Token::LeftBrace, + Token::Star, + Token::Identifier("p".to_string()), + Token::PlusPlus, + Token::Semicolon, + Token::RightBrace, + ]; + + let error = parse(&tokens).expect_err("parser should reject `*p++`"); + assert_eq!(error, ParseError::InvalidIncrementTarget); + } + #[test] fn parses_comparisons_and_logical_operators() { let tokens = vec![ diff --git a/tests/compiler_e2e.rs b/tests/compiler_e2e.rs index 95bf7a4..2d243fb 100644 --- a/tests/compiler_e2e.rs +++ b/tests/compiler_e2e.rs @@ -33,6 +33,22 @@ fn evaluates_division_and_subtraction() { assert_program_exit_code("int main() { return 20 / 5 - 1; }\n", 3); } +#[test] +fn evaluates_increment_and_decrement_statements() { + assert_program_exit_code( + "int main() { int x = 5; x++; ++x; x--; --x; x++; return x; }\n", + 6, + ); +} + +#[test] +fn evaluates_increment_in_for_post() { + assert_program_exit_code( + "int main() { int total = 0; for (int i = 0; i < 5; i++) { total += i; } return total; }\n", + 10, + ); +} + #[test] fn evaluates_compound_assignments() { assert_program_exit_code(