diff --git a/crates/parser/src/ast/statement.rs b/crates/parser/src/ast/statement.rs index ff89e07..361a434 100644 --- a/crates/parser/src/ast/statement.rs +++ b/crates/parser/src/ast/statement.rs @@ -8,27 +8,33 @@ pub struct Block( pub Option>, ); +#[derive(Debug, Clone, Serialize)] +pub enum StatementBody { + Statement(Box), + Expression(Expression), +} + #[derive(Debug, Clone, Serialize)] pub enum Statement { Block(Block), If { initial: StatementBranch, else_if: Vec, - else_branch: Option, + else_branch: Option, }, - Loop(Expression), + Loop(StatementBody), While(StatementBranch), CStyleFor { init: Expression, condition: Expression, update: Expression, - body: Expression, + body: StatementBody, }, For { mutable: bool, pattern: Pattern, iterator: Expression, - body: Expression, + body: StatementBody, }, Match(Expression, Vec<(Vec, Expression)>), @@ -54,7 +60,7 @@ pub struct VarDeclStmt { #[derive(Debug, Clone, Serialize)] pub struct StatementBranch { pub condition: Expression, - pub body: Box, + pub body: Box, } impl Statement { @@ -84,3 +90,12 @@ impl Statement { } } } + +impl StatementBody { + pub fn is_block(&self) -> bool { + match self { + Self::Expression(_) => false, + _ => true, + } + } +} diff --git a/crates/parser/src/grammar.pest b/crates/parser/src/grammar.pest index d255253..6651ae5 100644 --- a/crates/parser/src/grammar.pest +++ b/crates/parser/src/grammar.pest @@ -343,6 +343,10 @@ statement = _{ | (expr ~ semicolon) } +statement_wrapper = { statement } + +statement_body = { statement_wrapper | expr } + // ------------------------------------------------------ // BASIC STATEMENTS // ------------------------------------------------------ @@ -384,7 +388,7 @@ control_flow = { | block } -statement_branch = { "(" ~ expr ~ ")" ~ statement } +statement_branch = { "(" ~ expr ~ ")" ~ statement_body } else_if = { "else" ~ "if" ~ statement_branch @@ -395,7 +399,7 @@ else_if_list = { } if_stmt = { - "if" ~ statement_branch ~ else_if_list ~ ("else" ~ statement)? + "if" ~ statement_branch ~ else_if_list ~ ("else" ~ statement_body)? } // ------------------------------------------------------ @@ -407,15 +411,15 @@ while_stmt = { } c_for_stmt = { - "for" ~ "(" ~ statement ~ statement ~ expr ~ ")" ~ statement + "for" ~ "(" ~ statement ~ statement ~ expr ~ ")" ~ statement_body } for_stmt = { - "for" ~ "(" ~ mutable? ~ pattern ~ ":" ~ expr ~ ")" ~ statement + "for" ~ "(" ~ mutable? ~ pattern ~ ":" ~ expr ~ ")" ~ statement_body } loop_stmt = { - "loop" ~ statement + "loop" ~ statement_body } // ------------------------------------------------------ diff --git a/crates/parser/src/parser/common/statement.rs b/crates/parser/src/parser/common/statement.rs index 07c57f2..5fc632f 100644 --- a/crates/parser/src/parser/common/statement.rs +++ b/crates/parser/src/parser/common/statement.rs @@ -18,6 +18,24 @@ impl<'a> TryFrom> for Block { } } +impl<'a> TryFrom> for StatementBody { + type Error = AstError<'a, Self>; + + fn try_from(pair: pest::iterators::Pair<'a, Rule>) -> Result { + let mut inner = pair.clone().into_inner(); + + ast_ensure!(pair, Rule::statement_body => { + let i = inner.next().unwrap(); + + match i.as_rule() { + Rule::expr => ast_expr!(StatementBody::Expression(i.try_into())), + Rule::statement_wrapper => ast_expr!(StatementBody::Statement(i.try_into().map(Box::new).get_map(Box::new))), + _ => AstError::bug_unimplemented(i), + } + }) + } +} + impl<'a> TryFrom> for StatementBranch { type Error = AstError<'a, Self>; @@ -63,7 +81,7 @@ impl<'a> TryFrom> for Statement { ast_expr!(Statement::If { initial: inner.next().unwrap().try_into(), else_if: collect_recovered(inner.next().unwrap().into_inner()), - else_branch: inner.next().map(Expression::try_from).transpose(), + else_branch: inner.next().map(StatementBody::try_from).transpose(), }) }