diff --git a/executable_semantics/ast/statement.cpp b/executable_semantics/ast/statement.cpp index 7101d25eef22..79d92af3ad98 100644 --- a/executable_semantics/ast/statement.cpp +++ b/executable_semantics/ast/statement.cpp @@ -10,67 +10,63 @@ namespace Carbon { -const Expression* Statement::GetExpression() const { - CHECK(tag == StatementKind::ExpressionStatement); - return u.exp; +auto Statement::GetExpressionStatement() const -> const ExpressionStatement& { + return std::get(value); } -Assignment Statement::GetAssign() const { - CHECK(tag == StatementKind::Assign); - return u.assign; +auto Statement::GetAssign() const -> const Assign& { + return std::get(value); } -VariableDefinition Statement::GetVariableDefinition() const { - CHECK(tag == StatementKind::VariableDefinition); - return u.variable_definition; +auto Statement::GetVariableDefinition() const -> const VariableDefinition& { + return std::get(value); } -IfStatement Statement::GetIf() const { - CHECK(tag == StatementKind::If); - return u.if_stmt; +auto Statement::GetIf() const -> const If& { return std::get(value); } + +auto Statement::GetReturn() const -> const Return& { + return std::get(value); } -const Expression* Statement::GetReturn() const { - CHECK(tag == StatementKind::Return); - return u.return_stmt; +auto Statement::GetSequence() const -> const Sequence& { + return std::get(value); } -Sequence Statement::GetSequence() const { - CHECK(tag == StatementKind::Sequence); - return u.sequence; +auto Statement::GetBlock() const -> const Block& { + return std::get(value); } -Block Statement::GetBlock() const { - CHECK(tag == StatementKind::Block); - return u.block; +auto Statement::GetWhile() const -> const While& { + return std::get(value); } -While Statement::GetWhile() const { - CHECK(tag == StatementKind::While); - return u.while_stmt; +auto Statement::GetBreak() const -> const Break& { + return std::get(value); } -Match Statement::GetMatch() const { - CHECK(tag == StatementKind::Match); - return u.match_stmt; +auto Statement::GetContinue() const -> const Continue& { + return std::get(value); } -Continuation Statement::GetContinuation() const { - CHECK(tag == StatementKind::Continuation); - return u.continuation; +auto Statement::GetMatch() const -> const Match& { + return std::get(value); } -Run Statement::GetRun() const { - CHECK(tag == StatementKind::Run); - return u.run; +auto Statement::GetContinuation() const -> const Continuation& { + return std::get(value); } -auto Statement::MakeExpStmt(int line_num, const Expression* exp) +auto Statement::GetRun() const -> const Run& { return std::get(value); } + +auto Statement::GetAwait() const -> const Await& { + return std::get(value); +} + +auto Statement::MakeExpressionStatement(int line_num, const Expression* exp) -> const Statement* { auto* s = new Statement(); s->line_num = line_num; - s->tag = StatementKind::ExpressionStatement; - s->u.exp = exp; + s->value = ExpressionStatement({.exp = exp}); return s; } @@ -78,19 +74,16 @@ auto Statement::MakeAssign(int line_num, const Expression* lhs, const Expression* rhs) -> const Statement* { auto* s = new Statement(); s->line_num = line_num; - s->tag = StatementKind::Assign; - s->u.assign.lhs = lhs; - s->u.assign.rhs = rhs; + s->value = Assign({.lhs = lhs, .rhs = rhs}); return s; } -auto Statement::MakeVarDef(int line_num, const Expression* pat, - const Expression* init) -> const Statement* { +auto Statement::MakeVariableDefinition(int line_num, const Expression* pat, + const Expression* init) + -> const Statement* { auto* s = new Statement(); s->line_num = line_num; - s->tag = StatementKind::VariableDefinition; - s->u.variable_definition.pat = pat; - s->u.variable_definition.init = init; + s->value = VariableDefinition({.pat = pat, .init = init}); return s; } @@ -99,10 +92,7 @@ auto Statement::MakeIf(int line_num, const Expression* cond, -> const Statement* { auto* s = new Statement(); s->line_num = line_num; - s->tag = StatementKind::If; - s->u.if_stmt.cond = cond; - s->u.if_stmt.then_stmt = then_stmt; - s->u.if_stmt.else_stmt = else_stmt; + s->value = If({.cond = cond, .then_stmt = then_stmt, .else_stmt = else_stmt}); return s; } @@ -110,23 +100,21 @@ auto Statement::MakeWhile(int line_num, const Expression* cond, const Statement* body) -> const Statement* { auto* s = new Statement(); s->line_num = line_num; - s->tag = StatementKind::While; - s->u.while_stmt.cond = cond; - s->u.while_stmt.body = body; + s->value = While({.cond = cond, .body = body}); return s; } auto Statement::MakeBreak(int line_num) -> const Statement* { auto* s = new Statement(); s->line_num = line_num; - s->tag = StatementKind::Break; + s->value = Break(); return s; } auto Statement::MakeContinue(int line_num) -> const Statement* { auto* s = new Statement(); s->line_num = line_num; - s->tag = StatementKind::Continue; + s->value = Continue(); return s; } @@ -134,18 +122,15 @@ auto Statement::MakeReturn(int line_num, const Expression* e) -> const Statement* { auto* s = new Statement(); s->line_num = line_num; - s->tag = StatementKind::Return; - s->u.return_stmt = e; + s->value = Return({.exp = e}); return s; } -auto Statement::MakeSeq(int line_num, const Statement* s1, const Statement* s2) - -> const Statement* { +auto Statement::MakeSequence(int line_num, const Statement* s1, + const Statement* s2) -> const Statement* { auto* s = new Statement(); s->line_num = line_num; - s->tag = StatementKind::Sequence; - s->u.sequence.stmt = s1; - s->u.sequence.next = s2; + s->value = Sequence({.stmt = s1, .next = s2}); return s; } @@ -153,8 +138,7 @@ auto Statement::MakeBlock(int line_num, const Statement* stmt) -> const Statement* { auto* s = new Statement(); s->line_num = line_num; - s->tag = StatementKind::Block; - s->u.block.stmt = stmt; + s->value = Block({.stmt = stmt}); return s; } @@ -164,9 +148,7 @@ auto Statement::MakeMatch( -> const Statement* { auto* s = new Statement(); s->line_num = line_num; - s->tag = StatementKind::Match; - s->u.match_stmt.exp = exp; - s->u.match_stmt.clauses = clauses; + s->value = Match({.exp = exp, .clauses = clauses}); return s; } @@ -175,31 +157,29 @@ auto Statement::MakeMatch( auto Statement::MakeContinuation(int line_num, std::string continuation_variable, const Statement* body) -> const Statement* { - auto* continuation = new Statement(); - continuation->line_num = line_num; - continuation->tag = StatementKind::Continuation; - continuation->u.continuation.continuation_variable = - new std::string(continuation_variable); - continuation->u.continuation.body = body; - return continuation; + auto* s = new Statement(); + s->line_num = line_num; + s->value = + Continuation({.continuation_variable = std::move(continuation_variable), + .body = body}); + return s; } // Returns an AST node for a run statement give its line number and argument. auto Statement::MakeRun(int line_num, const Expression* argument) -> const Statement* { - auto* run = new Statement(); - run->line_num = line_num; - run->tag = StatementKind::Run; - run->u.run.argument = argument; - return run; + auto* s = new Statement(); + s->line_num = line_num; + s->value = Run({.argument = argument}); + return s; } // Returns an AST node for an await statement give its line number. auto Statement::MakeAwait(int line_num) -> const Statement* { - auto* await = new Statement(); - await->line_num = line_num; - await->tag = StatementKind::Await; - return await; + auto* s = new Statement(); + s->line_num = line_num; + s->value = Await(); + return s; } void PrintStatement(const Statement* s, int depth) { @@ -210,7 +190,7 @@ void PrintStatement(const Statement* s, int depth) { std::cout << " ... "; return; } - switch (s->tag) { + switch (s->tag()) { case StatementKind::Match: std::cout << "match ("; PrintExp(s->GetMatch().exp); @@ -249,7 +229,7 @@ void PrintStatement(const Statement* s, int depth) { std::cout << ";"; break; case StatementKind::ExpressionStatement: - PrintExp(s->GetExpression()); + PrintExp(s->GetExpressionStatement().exp); std::cout << ";"; break; case StatementKind::Assign: @@ -268,7 +248,7 @@ void PrintStatement(const Statement* s, int depth) { break; case StatementKind::Return: std::cout << "return "; - PrintExp(s->GetReturn()); + PrintExp(s->GetReturn().exp); std::cout << ";"; break; case StatementKind::Sequence: @@ -295,8 +275,8 @@ void PrintStatement(const Statement* s, int depth) { } break; case StatementKind::Continuation: - std::cout << "continuation " - << *s->GetContinuation().continuation_variable << " "; + std::cout << "continuation " << s->GetContinuation().continuation_variable + << " "; if (depth < 0 || depth > 1) { std::cout << std::endl; } diff --git a/executable_semantics/ast/statement.h b/executable_semantics/ast/statement.h index 8f4c5c880bc7..7409269beb8d 100644 --- a/executable_semantics/ast/statement.h +++ b/executable_semantics/ast/statement.h @@ -30,68 +30,96 @@ enum class StatementKind { struct Statement; -struct Assignment { +struct ExpressionStatement { + static constexpr StatementKind Kind = StatementKind::ExpressionStatement; + const Expression* exp; +}; + +struct Assign { + static constexpr StatementKind Kind = StatementKind::Assign; const Expression* lhs; const Expression* rhs; }; struct VariableDefinition { + static constexpr StatementKind Kind = StatementKind::VariableDefinition; const Expression* pat; const Expression* init; }; -struct IfStatement { +struct If { + static constexpr StatementKind Kind = StatementKind::If; const Expression* cond; const Statement* then_stmt; const Statement* else_stmt; }; +struct Return { + static constexpr StatementKind Kind = StatementKind::Return; + const Expression* exp; +}; + struct Sequence { + static constexpr StatementKind Kind = StatementKind::Sequence; const Statement* stmt; const Statement* next; }; struct Block { + static constexpr StatementKind Kind = StatementKind::Block; const Statement* stmt; }; struct While { + static constexpr StatementKind Kind = StatementKind::While; const Expression* cond; const Statement* body; }; +struct Break { + static constexpr StatementKind Kind = StatementKind::Break; +}; + +struct Continue { + static constexpr StatementKind Kind = StatementKind::Continue; +}; + struct Match { + static constexpr StatementKind Kind = StatementKind::Match; const Expression* exp; std::list>* clauses; }; struct Continuation { - std::string* continuation_variable; + static constexpr StatementKind Kind = StatementKind::Continuation; + std::string continuation_variable; const Statement* body; }; struct Run { + static constexpr StatementKind Kind = StatementKind::Run; const Expression* argument; }; -struct Statement { - // TODO: change Statement to a class and make all members private - int line_num; - StatementKind tag; +struct Await { + static constexpr StatementKind Kind = StatementKind::Await; +}; +struct Statement { // Constructors - static auto MakeExpStmt(int line_num, const Expression* exp) + static auto MakeExpressionStatement(int line_num, const Expression* exp) -> const Statement*; static auto MakeAssign(int line_num, const Expression* lhs, const Expression* rhs) -> const Statement*; - static auto MakeVarDef(int line_num, const Expression* pat, - const Expression* init) -> const Statement*; + static auto MakeVariableDefinition(int line_num, const Expression* pat, + const Expression* init) + -> const Statement*; static auto MakeIf(int line_num, const Expression* cond, const Statement* then_stmt, const Statement* else_stmt) -> const Statement*; static auto MakeReturn(int line_num, const Expression* e) -> const Statement*; - static auto MakeSeq(int line_num, const Statement* s1, const Statement* s2) - -> const Statement*; + static auto MakeSequence(int line_num, const Statement* s1, + const Statement* s2) -> const Statement*; static auto MakeBlock(int line_num, const Statement* s) -> const Statement*; static auto MakeWhile(int line_num, const Expression* cond, const Statement* body) -> const Statement*; @@ -119,33 +147,32 @@ struct Statement { // __await; static auto MakeAwait(int line_num) -> const Statement*; - // Access to the alternatives - const Expression* GetExpression() const; - Assignment GetAssign() const; - VariableDefinition GetVariableDefinition() const; - IfStatement GetIf() const; - const Expression* GetReturn() const; - Sequence GetSequence() const; - Block GetBlock() const; - While GetWhile() const; - Match GetMatch() const; - Continuation GetContinuation() const; - Run GetRun() const; + auto GetExpressionStatement() const -> const ExpressionStatement&; + auto GetAssign() const -> const Assign&; + auto GetVariableDefinition() const -> const VariableDefinition&; + auto GetIf() const -> const If&; + auto GetReturn() const -> const Return&; + auto GetSequence() const -> const Sequence&; + auto GetBlock() const -> const Block&; + auto GetWhile() const -> const While&; + auto GetBreak() const -> const Break&; + auto GetContinue() const -> const Continue&; + auto GetMatch() const -> const Match&; + auto GetContinuation() const -> const Continuation&; + auto GetRun() const -> const Run&; + auto GetAwait() const -> const Await&; + + inline auto tag() const -> StatementKind { + return std::visit([](const auto& t) { return t.Kind; }, value); + } + + int line_num; private: - union { - const Expression* exp; - Assignment assign; - VariableDefinition variable_definition; - IfStatement if_stmt; - const Expression* return_stmt; - Sequence sequence; - Block block; - While while_stmt; - Match match_stmt; - Continuation continuation; - Run run; - } u; + std::variant + value; }; void PrintStatement(const Statement*, int); diff --git a/executable_semantics/interpreter/interpreter.cpp b/executable_semantics/interpreter/interpreter.cpp index 73f2a4dc6b82..a3d1d0d3e3aa 100644 --- a/executable_semantics/interpreter/interpreter.cpp +++ b/executable_semantics/interpreter/interpreter.cpp @@ -967,7 +967,7 @@ void StepExp() { auto IsWhileAct(Action* act) -> bool { switch (act->tag()) { case ActionKind::StatementAction: - switch (act->GetStatementAction().stmt->tag) { + switch (act->GetStatementAction().stmt->tag()) { case StatementKind::While: return true; default: @@ -981,7 +981,7 @@ auto IsWhileAct(Action* act) -> bool { auto IsBlockAct(Action* act) -> bool { switch (act->tag()) { case ActionKind::StatementAction: - switch (act->GetStatementAction().stmt->tag) { + switch (act->GetStatementAction().stmt->tag()) { case StatementKind::Block: return true; default: @@ -1004,7 +1004,7 @@ void StepStmt() { PrintStatement(stmt, 1); std::cout << " --->" << std::endl; } - switch (stmt->tag) { + switch (stmt->tag()) { case StatementKind::Match: if (act->pos == 0) { // { { (match (e) ...) :: C, E, F} :: S, H} @@ -1164,7 +1164,8 @@ void StepStmt() { if (act->pos == 0) { // { {e :: C, E, F} :: S, H} // -> { {e :: C, E, F} :: S, H} - frame->todo.Push(Action::MakeExpressionAction(stmt->GetExpression())); + frame->todo.Push( + Action::MakeExpressionAction(stmt->GetExpressionStatement().exp)); act->pos++; } else { frame->todo.Pop(1); @@ -1216,7 +1217,7 @@ void StepStmt() { if (act->pos == 0) { // { {return e :: C, E, F} :: S, H} // -> { {e :: return [] :: C, E, F} :: S, H} - frame->todo.Push(Action::MakeExpressionAction(stmt->GetReturn())); + frame->todo.Push(Action::MakeExpressionAction(stmt->GetReturn().exp)); act->pos++; } else { // { {v :: return [] :: C, E, F} :: {C', E', F'} :: S, H} @@ -1256,7 +1257,7 @@ void StepStmt() { continuation_frame->continuation = continuation_address; // Bind the continuation object to the continuation variable frame->scopes.Top()->values.Set( - *stmt->GetContinuation().continuation_variable, continuation_address); + stmt->GetContinuation().continuation_variable, continuation_address); // Pop the continuation statement. frame->todo.Pop(); break; @@ -1270,9 +1271,10 @@ void StepStmt() { frame->todo.Pop(1); // Push an expression statement action to ignore the result // value from the continuation. - Action* ignore_result = Action::MakeStatementAction( - Statement::MakeExpStmt(stmt->line_num, Expression::MakeTupleLiteral( - stmt->line_num, {}))); + Action* ignore_result = + Action::MakeStatementAction(Statement::MakeExpressionStatement( + stmt->line_num, + Expression::MakeTupleLiteral(stmt->line_num, {}))); ignore_result->pos = 0; frame->todo.Push(ignore_result); // Push the continuation onto the current stack. diff --git a/executable_semantics/interpreter/typecheck.cpp b/executable_semantics/interpreter/typecheck.cpp index 32b2affff715..47c583f91474 100644 --- a/executable_semantics/interpreter/typecheck.cpp +++ b/executable_semantics/interpreter/typecheck.cpp @@ -456,7 +456,7 @@ auto TypeCheckStmt(const Statement* s, TypeEnv types, Env values, if (!s) { return TCStatement(s, types); } - switch (s->tag) { + switch (s->tag()) { case StatementKind::Match: { auto res = TypeCheckExp(s->GetMatch().exp, types, values, nullptr, TCContext::ValueContext); @@ -497,7 +497,7 @@ auto TypeCheckStmt(const Statement* s, TypeEnv types, Env values, const Value* rhs_ty = res.type; auto lhs_res = TypeCheckExp(s->GetVariableDefinition().pat, types, values, rhs_ty, TCContext::PatternContext); - const Statement* new_s = Statement::MakeVarDef( + const Statement* new_s = Statement::MakeVariableDefinition( s->line_num, s->GetVariableDefinition().pat, res.exp); return TCStatement(new_s, lhs_res.types); } @@ -509,7 +509,7 @@ auto TypeCheckStmt(const Statement* s, TypeEnv types, Env values, TypeCheckStmt(s->GetSequence().next, types2, values, ret_type); auto types3 = next_res.types; return TCStatement( - Statement::MakeSeq(s->line_num, stmt_res.stmt, next_res.stmt), + Statement::MakeSequence(s->line_num, stmt_res.stmt, next_res.stmt), types3); } case StatementKind::Assign: { @@ -524,9 +524,9 @@ auto TypeCheckStmt(const Statement* s, TypeEnv types, Env values, return TCStatement(new_s, lhs_res.types); } case StatementKind::ExpressionStatement: { - auto res = TypeCheckExp(s->GetExpression(), types, values, nullptr, - TCContext::ValueContext); - auto new_s = Statement::MakeExpStmt(s->line_num, res.exp); + auto res = TypeCheckExp(s->GetExpressionStatement().exp, types, values, + nullptr, TCContext::ValueContext); + auto new_s = Statement::MakeExpressionStatement(s->line_num, res.exp); return TCStatement(new_s, types); } case StatementKind::If: { @@ -543,7 +543,7 @@ auto TypeCheckStmt(const Statement* s, TypeEnv types, Env values, return TCStatement(new_s, types); } case StatementKind::Return: { - auto res = TypeCheckExp(s->GetReturn(), types, values, nullptr, + auto res = TypeCheckExp(s->GetReturn().exp, types, values, nullptr, TCContext::ValueContext); if (ret_type->tag() == ValKind::AutoType) { // The following infers the return type from the first 'return' @@ -559,9 +559,9 @@ auto TypeCheckStmt(const Statement* s, TypeEnv types, Env values, TCStatement body_result = TypeCheckStmt(s->GetContinuation().body, types, values, ret_type); const Statement* new_continuation = Statement::MakeContinuation( - s->line_num, *s->GetContinuation().continuation_variable, + s->line_num, s->GetContinuation().continuation_variable, body_result.stmt); - types.Set(*s->GetContinuation().continuation_variable, + types.Set(s->GetContinuation().continuation_variable, Value::MakeContinuationType()); return TCStatement(new_continuation, types); } @@ -595,7 +595,7 @@ auto CheckOrEnsureReturn(const Statement* stmt, bool void_return, int line_num) exit(-1); } } - switch (stmt->tag) { + switch (stmt->tag()) { case StatementKind::Match: { auto new_clauses = new std::list>(); @@ -622,7 +622,7 @@ auto CheckOrEnsureReturn(const Statement* stmt, bool void_return, int line_num) return stmt; case StatementKind::Sequence: if (stmt->GetSequence().next) { - return Statement::MakeSeq( + return Statement::MakeSequence( stmt->line_num, stmt->GetSequence().stmt, CheckOrEnsureReturn(stmt->GetSequence().next, void_return, stmt->line_num)); @@ -641,7 +641,7 @@ auto CheckOrEnsureReturn(const Statement* stmt, bool void_return, int line_num) case StatementKind::Continue: case StatementKind::VariableDefinition: if (void_return) { - return Statement::MakeSeq( + return Statement::MakeSequence( stmt->line_num, stmt, Statement::MakeReturn(stmt->line_num, Expression::MakeTupleLiteral( stmt->line_num, {}))); diff --git a/executable_semantics/syntax/parser.ypp b/executable_semantics/syntax/parser.ypp index 7a93c7d954d0..733351f8a6e2 100644 --- a/executable_semantics/syntax/parser.ypp +++ b/executable_semantics/syntax/parser.ypp @@ -327,9 +327,9 @@ statement: expression "=" expression ";" { $$ = Carbon::Statement::MakeAssign(yylineno, $1, $3); } | VAR pattern "=" expression ";" - { $$ = Carbon::Statement::MakeVarDef(yylineno, $2, $4); } + { $$ = Carbon::Statement::MakeVariableDefinition(yylineno, $2, $4); } | expression ";" - { $$ = Carbon::Statement::MakeExpStmt(yylineno, $1); } + { $$ = Carbon::Statement::MakeExpressionStatement(yylineno, $1); } | IF "(" expression ")" statement optional_else { $$ = Carbon::Statement::MakeIf(yylineno, $3, $5, $6); } | WHILE "(" expression ")" statement @@ -360,7 +360,7 @@ statement_list: // Empty { $$ = 0; } | statement statement_list - { $$ = Carbon::Statement::MakeSeq(yylineno, $1, $2); } + { $$ = Carbon::Statement::MakeSequence(yylineno, $1, $2); } ; return_type: // Empty