From 36ed79dc25dae36f283d36587beb4eb906d5529f Mon Sep 17 00:00:00 2001 From: Jon Meow <46229924+jonmeow@users.noreply.github.com> Date: Mon, 30 Aug 2021 14:41:14 -0700 Subject: [PATCH] Convert Statement to use Ptr (#788) Note this makes a few cases where the Statement was optional explicit (Block, If, Sequence). I do add a few CHECKs around where statements were optional and assumed but unverified. I switch TypeCheckStmt to not take an optional Statement because I think it makes the call sites clearer in behavior. It's also a smaller change than the converse, because taking an optional Statement means the returned statement would also need to be optional. Arguably a wrapper for optional statements could be added, but this still seems cleaner to me, and there aren't that many cases of an optional statement. Co-authored-by: Geoff Romer --- .../ast/function_definition.cpp | 2 +- .../ast/function_definition.h | 5 +- executable_semantics/ast/statement.cpp | 6 +- executable_semantics/ast/statement.h | 51 ++++---- executable_semantics/interpreter/action.h | 6 +- .../interpreter/interpreter.cpp | 27 ++-- .../interpreter/typecheck.cpp | 118 ++++++++++-------- executable_semantics/interpreter/typecheck.h | 6 +- executable_semantics/interpreter/value.cpp | 10 +- executable_semantics/interpreter/value.h | 7 +- executable_semantics/syntax/parser.ypp | 56 ++++----- .../syntax/syntax_helpers.cpp | 10 +- 12 files changed, 166 insertions(+), 138 deletions(-) diff --git a/executable_semantics/ast/function_definition.cpp b/executable_semantics/ast/function_definition.cpp index ec69180bb193..b7759f62cd25 100644 --- a/executable_semantics/ast/function_definition.cpp +++ b/executable_semantics/ast/function_definition.cpp @@ -27,7 +27,7 @@ void FunctionDefinition::PrintDepth(int depth, llvm::raw_ostream& out) const { } if (body) { out << " {\n"; - body->PrintDepth(depth, out); + (*body)->PrintDepth(depth, out); out << "\n}\n"; } else { out << ";\n"; diff --git a/executable_semantics/ast/function_definition.h b/executable_semantics/ast/function_definition.h index 71a0f0c088fa..dbb4de86e7f9 100644 --- a/executable_semantics/ast/function_definition.h +++ b/executable_semantics/ast/function_definition.h @@ -26,7 +26,8 @@ struct FunctionDefinition { std::vector deduced_params, Ptr param_pattern, Ptr return_type, - bool is_omitted_return_type, const Statement* body) + bool is_omitted_return_type, + std::optional> body) : source_location(source_location), name(std::move(name)), deduced_parameters(deduced_params), @@ -45,7 +46,7 @@ struct FunctionDefinition { Ptr param_pattern; Ptr return_type; bool is_omitted_return_type; - const Statement* body; + std::optional> body; }; } // namespace Carbon diff --git a/executable_semantics/ast/statement.cpp b/executable_semantics/ast/statement.cpp index 4990101304a1..5eb2dd513fc0 100644 --- a/executable_semantics/ast/statement.cpp +++ b/executable_semantics/ast/statement.cpp @@ -65,7 +65,7 @@ void Statement::PrintDepth(int depth, llvm::raw_ostream& out) const { if_stmt.ThenStmt()->PrintDepth(depth - 1, out); if (if_stmt.ElseStmt()) { out << "\nelse\n"; - if_stmt.ElseStmt()->PrintDepth(depth - 1, out); + (*if_stmt.ElseStmt())->PrintDepth(depth - 1, out); } break; } @@ -87,7 +87,7 @@ void Statement::PrintDepth(int depth, llvm::raw_ostream& out) const { out << " "; } if (seq.Next()) { - seq.Next()->PrintDepth(depth - 1, out); + (*seq.Next())->PrintDepth(depth - 1, out); } break; } @@ -98,7 +98,7 @@ void Statement::PrintDepth(int depth, llvm::raw_ostream& out) const { out << "\n"; } if (block.Stmt()) { - block.Stmt()->PrintDepth(depth, out); + (*block.Stmt())->PrintDepth(depth, out); if (depth < 0 || depth > 1) { out << "\n"; } diff --git a/executable_semantics/ast/statement.h b/executable_semantics/ast/statement.h index 56b171cc4ce3..255ab8db0365 100644 --- a/executable_semantics/ast/statement.h +++ b/executable_semantics/ast/statement.h @@ -109,8 +109,9 @@ class VariableDefinition : public Statement { class If : public Statement { public: - If(SourceLocation loc, Ptr cond, const Statement* then_stmt, - const Statement* else_stmt) + If(SourceLocation loc, Ptr cond, + Ptr then_stmt, + std::optional> else_stmt) : Statement(Kind::If, loc), cond(cond), then_stmt(then_stmt), @@ -121,13 +122,15 @@ class If : public Statement { } auto Cond() const -> Ptr { return cond; } - auto ThenStmt() const -> const Statement* { return then_stmt; } - auto ElseStmt() const -> const Statement* { return else_stmt; } + auto ThenStmt() const -> Ptr { return then_stmt; } + auto ElseStmt() const -> std::optional> { + return else_stmt; + } private: Ptr cond; - const Statement* then_stmt; - const Statement* else_stmt; + Ptr then_stmt; + std::optional> else_stmt; }; class Return : public Statement { @@ -153,39 +156,41 @@ class Return : public Statement { class Sequence : public Statement { public: - Sequence(SourceLocation loc, const Statement* stmt, const Statement* next) + Sequence(SourceLocation loc, Ptr stmt, + std::optional> next) : Statement(Kind::Sequence, loc), stmt(stmt), next(next) {} static auto classof(const Statement* stmt) -> bool { return stmt->Tag() == Kind::Sequence; } - auto Stmt() const -> const Statement* { return stmt; } - auto Next() const -> const Statement* { return next; } + auto Stmt() const -> Ptr { return stmt; } + auto Next() const -> std::optional> { return next; } private: - const Statement* stmt; - const Statement* next; + Ptr stmt; + std::optional> next; }; class Block : public Statement { public: - Block(SourceLocation loc, const Statement* stmt) + Block(SourceLocation loc, std::optional> stmt) : Statement(Kind::Block, loc), stmt(stmt) {} static auto classof(const Statement* stmt) -> bool { return stmt->Tag() == Kind::Block; } - auto Stmt() const -> const Statement* { return stmt; } + auto Stmt() const -> std::optional> { return stmt; } private: - const Statement* stmt; + std::optional> stmt; }; class While : public Statement { public: - While(SourceLocation loc, Ptr cond, const Statement* body) + While(SourceLocation loc, Ptr cond, + Ptr body) : Statement(Kind::While, loc), cond(cond), body(body) {} static auto classof(const Statement* stmt) -> bool { @@ -193,11 +198,11 @@ class While : public Statement { } auto Cond() const -> Ptr { return cond; } - auto Body() const -> const Statement* { return body; } + auto Body() const -> Ptr { return body; } private: Ptr cond; - const Statement* body; + Ptr body; }; class Break : public Statement { @@ -221,7 +226,7 @@ class Continue : public Statement { class Match : public Statement { public: Match(SourceLocation loc, Ptr exp, - std::list, const Statement*>>* clauses) + std::list, Ptr>>* clauses) : Statement(Kind::Match, loc), exp(exp), clauses(clauses) {} static auto classof(const Statement* stmt) -> bool { @@ -230,13 +235,13 @@ class Match : public Statement { auto Exp() const -> Ptr { return exp; } auto Clauses() const - -> const std::list, const Statement*>>* { + -> const std::list, Ptr>>* { return clauses; } private: Ptr exp; - std::list, const Statement*>>* clauses; + std::list, Ptr>>* clauses; }; // A continuation statement. @@ -247,7 +252,7 @@ class Match : public Statement { class Continuation : public Statement { public: Continuation(SourceLocation loc, std::string continuation_variable, - const Statement* body) + Ptr body) : Statement(Kind::Continuation, loc), continuation_variable(std::move(continuation_variable)), body(body) {} @@ -259,11 +264,11 @@ class Continuation : public Statement { auto ContinuationVariable() const -> const std::string& { return continuation_variable; } - auto Body() const -> const Statement* { return body; } + auto Body() const -> Ptr { return body; } private: std::string continuation_variable; - const Statement* body; + Ptr body; }; // A run statement. diff --git a/executable_semantics/interpreter/action.h b/executable_semantics/interpreter/action.h index 79950ef02361..5d341f78fbd5 100644 --- a/executable_semantics/interpreter/action.h +++ b/executable_semantics/interpreter/action.h @@ -117,17 +117,17 @@ class PatternAction : public Action { class StatementAction : public Action { public: - explicit StatementAction(const Statement* stmt) + explicit StatementAction(Ptr stmt) : Action(Kind::StatementAction), stmt(stmt) {} static auto classof(const Action* action) -> bool { return action->Tag() == Kind::StatementAction; } - auto Stmt() const -> const Statement* { return stmt; } + auto Stmt() const -> Ptr { return stmt; } private: - const Statement* stmt; + Ptr stmt; }; } // namespace Carbon diff --git a/executable_semantics/interpreter/interpreter.cpp b/executable_semantics/interpreter/interpreter.cpp index d19c7cacd393..b72b458192b3 100644 --- a/executable_semantics/interpreter/interpreter.cpp +++ b/executable_semantics/interpreter/interpreter.cpp @@ -812,8 +812,7 @@ auto IsBlockAct(Ptr act) -> bool { Transition StepStmt() { Ptr frame = state->stack.Top(); Ptr act = frame->todo.Top(); - const Statement* stmt = cast(*act).Stmt(); - CHECK(stmt != nullptr) << "null statement!"; + Ptr stmt = cast(*act).Stmt(); if (tracing_output) { llvm::outs() << "--- step stmt "; stmt->PrintDepth(1, llvm::outs()); @@ -861,9 +860,8 @@ Transition StepStmt() { vars.push_back(name); } frame->scopes.Push(global_arena->New(values, vars)); - const Statement* body_block = - global_arena->RawNew(stmt->SourceLoc(), c->second); - auto body_act = global_arena->New(body_block); + auto body_act = global_arena->New( + global_arena->New(stmt->SourceLoc(), c->second)); body_act->IncrementPos(); frame->todo.Pop(1); frame->todo.Push(body_act); @@ -925,9 +923,9 @@ Transition StepStmt() { case Statement::Kind::Block: { if (act->Pos() == 0) { const Block& block = cast(*stmt); - if (block.Stmt() != nullptr) { + if (block.Stmt()) { frame->scopes.Push(global_arena->New(CurrentEnv(state))); - return Spawn{global_arena->New(block.Stmt())}; + return Spawn{global_arena->New(*block.Stmt())}; } else { return Done{}; } @@ -1007,7 +1005,7 @@ Transition StepStmt() { // S, H} // -> { { else_stmt :: C, E, F } :: S, H} return Delegate{ - global_arena->New(cast(*stmt).ElseStmt())}; + global_arena->New(*cast(*stmt).ElseStmt())}; } else { return Done{}; } @@ -1030,9 +1028,9 @@ Transition StepStmt() { if (act->Pos() == 0) { return Spawn{global_arena->New(seq.Stmt())}; } else { - if (seq.Next() != nullptr) { - return Delegate{ - global_arena->New(cast(*stmt).Next())}; + if (seq.Next()) { + return Delegate{global_arena->New( + *cast(*stmt).Next())}; } else { return Done{}; } @@ -1046,7 +1044,7 @@ Transition StepStmt() { Stack>(global_arena->New(CurrentEnv(state))); Stack> todo; todo.Push(global_arena->New( - global_arena->RawNew(stmt->SourceLoc()))); + global_arena->New(stmt->SourceLoc()))); todo.Push( global_arena->New(cast(*stmt).Body())); auto continuation_frame = @@ -1074,7 +1072,7 @@ Transition StepStmt() { // Push an expression statement action to ignore the result // value from the continuation. auto ignore_result = global_arena->New( - global_arena->RawNew( + global_arena->New( stmt->SourceLoc(), global_arena->New(stmt->SourceLoc()))); frame->todo.Push(ignore_result); @@ -1172,8 +1170,9 @@ struct DoTransition { params.push_back(name); } auto scopes = Stack>(global_arena->New(values, params)); + CHECK(call.function->Body()) << "Calling a function that's missing a body"; auto todo = Stack>( - global_arena->New(call.function->Body())); + global_arena->New(*call.function->Body())); auto frame = global_arena->New(call.function->Name(), scopes, todo); state->stack.Push(frame); } diff --git a/executable_semantics/interpreter/typecheck.cpp b/executable_semantics/interpreter/typecheck.cpp index 06d787de50bc..1b9d00a44977 100644 --- a/executable_semantics/interpreter/typecheck.cpp +++ b/executable_semantics/interpreter/typecheck.cpp @@ -657,9 +657,9 @@ auto TypeCheckPattern(Ptr p, TypeEnv types, Env values, } static auto TypecheckCase(const Value* expected, Ptr pat, - const Statement* body, TypeEnv types, Env values, + Ptr body, TypeEnv types, Env values, const Value*& ret_type, bool is_omitted_ret_type) - -> std::pair, const Statement*> { + -> std::pair, Ptr> { auto pat_res = TypeCheckPattern(pat, types, values, expected); auto res = TypeCheckStmt(body, pat_res.types, values, ret_type, is_omitted_ret_type); @@ -673,26 +673,23 @@ static auto TypecheckCase(const Value* expected, Ptr pat, // It is the declared return type of the enclosing function definition. // If the return type is "auto", then the return type is inferred from // the first return statement. -auto TypeCheckStmt(const Statement* s, TypeEnv types, Env values, +auto TypeCheckStmt(Ptr s, TypeEnv types, Env values, const Value*& ret_type, bool is_omitted_ret_type) -> TCStatement { - if (!s) { - return TCStatement(s, types); - } switch (s->Tag()) { case Statement::Kind::Match: { const auto& match = cast(*s); auto res = TypeCheckExp(match.Exp(), types, values); auto res_type = res.type; auto new_clauses = global_arena->RawNew< - std::list, const Statement*>>>(); + std::list, Ptr>>>(); for (auto& clause : *match.Clauses()) { new_clauses->push_back(TypecheckCase(res_type, clause.first, clause.second, types, values, ret_type, is_omitted_ret_type)); } - const Statement* new_s = - global_arena->RawNew(s->SourceLoc(), res.exp, new_clauses); + auto new_s = + global_arena->New(s->SourceLoc(), res.exp, new_clauses); return TCStatement(new_s, types); } case Statement::Kind::While: { @@ -702,39 +699,48 @@ auto TypeCheckStmt(const Statement* s, TypeEnv types, Env values, global_arena->RawNew(), cnd_res.type); auto body_res = TypeCheckStmt(while_stmt.Body(), types, values, ret_type, is_omitted_ret_type); - auto new_s = global_arena->RawNew(s->SourceLoc(), cnd_res.exp, - body_res.stmt); + auto new_s = + global_arena->New(s->SourceLoc(), cnd_res.exp, body_res.stmt); return TCStatement(new_s, types); } case Statement::Kind::Break: case Statement::Kind::Continue: return TCStatement(s, types); case Statement::Kind::Block: { - auto stmt_res = TypeCheckStmt(cast(*s).Stmt(), types, values, - ret_type, is_omitted_ret_type); - return TCStatement( - global_arena->RawNew(s->SourceLoc(), stmt_res.stmt), types); + const auto& block = cast(*s); + if (block.Stmt()) { + auto stmt_res = TypeCheckStmt(*block.Stmt(), types, values, ret_type, + is_omitted_ret_type); + return TCStatement( + global_arena->New(s->SourceLoc(), stmt_res.stmt), types); + } else { + return TCStatement(s, types); + } } case Statement::Kind::VariableDefinition: { const auto& var = cast(*s); auto res = TypeCheckExp(var.Init(), types, values); const Value* rhs_ty = res.type; auto lhs_res = TypeCheckPattern(var.Pat(), types, values, rhs_ty); - const Statement* new_s = global_arena->RawNew( - s->SourceLoc(), var.Pat(), res.exp); + auto new_s = global_arena->New(s->SourceLoc(), + var.Pat(), res.exp); return TCStatement(new_s, lhs_res.types); } case Statement::Kind::Sequence: { const auto& seq = cast(*s); auto stmt_res = TypeCheckStmt(seq.Stmt(), types, values, ret_type, is_omitted_ret_type); - auto types2 = stmt_res.types; - auto next_res = TypeCheckStmt(seq.Next(), types2, values, ret_type, - is_omitted_ret_type); - auto types3 = next_res.types; - return TCStatement(global_arena->RawNew( - s->SourceLoc(), stmt_res.stmt, next_res.stmt), - types3); + auto checked_types = stmt_res.types; + std::optional> next_stmt; + if (seq.Next()) { + auto next_res = TypeCheckStmt(*seq.Next(), checked_types, values, + ret_type, is_omitted_ret_type); + next_stmt = next_res.stmt; + checked_types = next_res.types; + } + return TCStatement( + global_arena->New(s->SourceLoc(), stmt_res.stmt, next_stmt), + checked_types); } case Statement::Kind::Assign: { const auto& assign = cast(*s); @@ -743,15 +749,15 @@ auto TypeCheckStmt(const Statement* s, TypeEnv types, Env values, auto lhs_res = TypeCheckExp(assign.Lhs(), types, values); auto lhs_t = lhs_res.type; ExpectType(s->SourceLoc(), "assign", lhs_t, rhs_t); - auto new_s = global_arena->RawNew(s->SourceLoc(), lhs_res.exp, - rhs_res.exp); + auto new_s = + global_arena->New(s->SourceLoc(), lhs_res.exp, rhs_res.exp); return TCStatement(new_s, lhs_res.types); } case Statement::Kind::ExpressionStatement: { auto res = TypeCheckExp(cast(*s).Exp(), types, values); auto new_s = - global_arena->RawNew(s->SourceLoc(), res.exp); + global_arena->New(s->SourceLoc(), res.exp); return TCStatement(new_s, types); } case Statement::Kind::If: { @@ -761,10 +767,14 @@ auto TypeCheckStmt(const Statement* s, TypeEnv types, Env values, global_arena->RawNew(), cnd_res.type); auto then_res = TypeCheckStmt(if_stmt.ThenStmt(), types, values, ret_type, is_omitted_ret_type); - auto else_res = TypeCheckStmt(if_stmt.ElseStmt(), types, values, ret_type, - is_omitted_ret_type); - auto new_s = global_arena->RawNew(s->SourceLoc(), cnd_res.exp, - then_res.stmt, else_res.stmt); + std::optional> else_stmt; + if (if_stmt.ElseStmt()) { + auto else_res = TypeCheckStmt(*if_stmt.ElseStmt(), types, values, + ret_type, is_omitted_ret_type); + else_stmt = else_res.stmt; + } + auto new_s = global_arena->New(s->SourceLoc(), cnd_res.exp, + then_res.stmt, else_stmt); return TCStatement(new_s, types); } case Statement::Kind::Return: { @@ -783,15 +793,15 @@ auto TypeCheckStmt(const Statement* s, TypeEnv types, Env values, << *s << " should" << (is_omitted_ret_type ? " not" : "") << " provide a return value, to match the function's signature."; } - return TCStatement(global_arena->RawNew(s->SourceLoc(), res.exp, - ret.IsOmittedExp()), + return TCStatement(global_arena->New(s->SourceLoc(), res.exp, + ret.IsOmittedExp()), types); } case Statement::Kind::Continuation: { const auto& cont = cast(*s); TCStatement body_result = TypeCheckStmt(cont.Body(), types, values, ret_type, is_omitted_ret_type); - const Statement* new_continuation = global_arena->RawNew( + auto new_continuation = global_arena->New( s->SourceLoc(), cont.ContinuationVariable(), body_result.stmt); types.Set(cont.ContinuationVariable(), global_arena->RawNew()); @@ -803,8 +813,8 @@ auto TypeCheckStmt(const Statement* s, TypeEnv types, Env values, ExpectType(s->SourceLoc(), "argument of `run`", global_arena->RawNew(), argument_result.type); - const Statement* new_run = - global_arena->RawNew(s->SourceLoc(), argument_result.exp); + auto new_run = + global_arena->New(s->SourceLoc(), argument_result.exp); return TCStatement(new_run, types); } case Statement::Kind::Await: { @@ -814,38 +824,40 @@ auto TypeCheckStmt(const Statement* s, TypeEnv types, Env values, } // switch } -static auto CheckOrEnsureReturn(const Statement* stmt, bool omitted_ret_type, - SourceLocation loc) -> const Statement* { - if (!stmt) { +static auto CheckOrEnsureReturn(std::optional> opt_stmt, + bool omitted_ret_type, SourceLocation loc) + -> Ptr { + if (!opt_stmt) { if (omitted_ret_type) { - return global_arena->RawNew(loc); + return global_arena->New(loc); } else { FATAL_COMPILATION_ERROR(loc) << "control-flow reaches end of function that provides a `->` return " "type without reaching a return statement"; } } + Ptr stmt = *opt_stmt; switch (stmt->Tag()) { case Statement::Kind::Match: { const auto& match = cast(*stmt); auto new_clauses = global_arena->RawNew< - std::list, const Statement*>>>(); + std::list, Ptr>>>(); for (const auto& clause : *match.Clauses()) { auto s = CheckOrEnsureReturn(clause.second, omitted_ret_type, stmt->SourceLoc()); new_clauses->push_back(std::make_pair(clause.first, s)); } - return global_arena->RawNew(stmt->SourceLoc(), match.Exp(), - new_clauses); + return global_arena->New(stmt->SourceLoc(), match.Exp(), + new_clauses); } case Statement::Kind::Block: - return global_arena->RawNew( + return global_arena->New( stmt->SourceLoc(), CheckOrEnsureReturn(cast(*stmt).Stmt(), omitted_ret_type, stmt->SourceLoc())); case Statement::Kind::If: { const auto& if_stmt = cast(*stmt); - return global_arena->RawNew( + return global_arena->New( stmt->SourceLoc(), if_stmt.Cond(), CheckOrEnsureReturn(if_stmt.ThenStmt(), omitted_ret_type, stmt->SourceLoc()), @@ -857,7 +869,7 @@ static auto CheckOrEnsureReturn(const Statement* stmt, bool omitted_ret_type, case Statement::Kind::Sequence: { const auto& seq = cast(*stmt); if (seq.Next()) { - return global_arena->RawNew( + return global_arena->New( stmt->SourceLoc(), seq.Stmt(), CheckOrEnsureReturn(seq.Next(), omitted_ret_type, stmt->SourceLoc())); @@ -877,8 +889,8 @@ static auto CheckOrEnsureReturn(const Statement* stmt, bool omitted_ret_type, case Statement::Kind::Continue: case Statement::Kind::VariableDefinition: if (omitted_ret_type) { - return global_arena->RawNew( - stmt->SourceLoc(), stmt, global_arena->RawNew(loc)); + return global_arena->New(stmt->SourceLoc(), stmt, + global_arena->New(loc)); } else { FATAL_COMPILATION_ERROR(stmt->SourceLoc()) << "control-flow reaches end of function that provides a `->` " @@ -909,9 +921,13 @@ static auto TypeCheckFunDef(const FunctionDefinition* f, TypeEnv types, global_arena->RawNew(), return_type); // TODO: Check that main doesn't have any parameters. } - auto res = TypeCheckStmt(f->body, param_res.types, values, return_type, - f->is_omitted_return_type); - auto body = CheckOrEnsureReturn(res.stmt, f->is_omitted_return_type, + std::optional> body_stmt; + if (f->body) { + auto res = TypeCheckStmt(*f->body, param_res.types, values, return_type, + f->is_omitted_return_type); + body_stmt = res.stmt; + } + auto body = CheckOrEnsureReturn(body_stmt, f->is_omitted_return_type, f->source_location); return global_arena->New( f->source_location, f->name, f->deduced_parameters, f->param_pattern, diff --git a/executable_semantics/interpreter/typecheck.h b/executable_semantics/interpreter/typecheck.h index 3b01d0919620..eda106268a03 100644 --- a/executable_semantics/interpreter/typecheck.h +++ b/executable_semantics/interpreter/typecheck.h @@ -34,9 +34,9 @@ struct TCPattern { }; struct TCStatement { - TCStatement(const Statement* s, TypeEnv types) : stmt(s), types(types) {} + TCStatement(Ptr s, TypeEnv types) : stmt(s), types(types) {} - const Statement* stmt; + Ptr stmt; TypeEnv types; }; @@ -52,7 +52,7 @@ auto TypeCheckExp(Ptr e, TypeEnv types, Env values) auto TypeCheckPattern(Ptr p, TypeEnv types, Env values, const Value* expected) -> TCPattern; -auto TypeCheckStmt(const Statement* s, TypeEnv types, Env values, +auto TypeCheckStmt(Ptr s, TypeEnv types, Env values, const Value*& ret_type, bool is_omitted_ret_type) -> TCStatement; diff --git a/executable_semantics/interpreter/value.cpp b/executable_semantics/interpreter/value.cpp index afee2b72b646..a0944107f62d 100644 --- a/executable_semantics/interpreter/value.cpp +++ b/executable_semantics/interpreter/value.cpp @@ -399,8 +399,14 @@ auto ValueEqual(const Value* v1, const Value* v2, SourceLocation loc) -> bool { return cast(*v1).Val() == cast(*v2).Val(); case Value::Kind::PointerValue: return cast(*v1).Val() == cast(*v2).Val(); - case Value::Kind::FunctionValue: - return cast(*v1).Body() == cast(*v2).Body(); + case Value::Kind::FunctionValue: { + std::optional> body1 = + cast(*v1).Body(); + std::optional> body2 = + cast(*v2).Body(); + return body1.has_value() == body2.has_value() && + (!body1.has_value() || *body1 == *body2); + } case Value::Kind::TupleValue: return FieldsValueEqual(cast(*v1).Elements(), cast(*v2).Elements(), loc); diff --git a/executable_semantics/interpreter/value.h b/executable_semantics/interpreter/value.h index ea1d1ec0f4ff..e1198b799334 100644 --- a/executable_semantics/interpreter/value.h +++ b/executable_semantics/interpreter/value.h @@ -121,7 +121,8 @@ class IntValue : public Value { // A function value. class FunctionValue : public Value { public: - FunctionValue(std::string name, const Value* param, const Statement* body) + FunctionValue(std::string name, const Value* param, + std::optional> body) : Value(Kind::FunctionValue), name(std::move(name)), param(param), @@ -133,12 +134,12 @@ class FunctionValue : public Value { auto Name() const -> const std::string& { return name; } auto Param() const -> const Value* { return param; } - auto Body() const -> const Statement* { return body; } + auto Body() const -> std::optional> { return body; } private: std::string name; const Value* param; - const Statement* body; + std::optional> body; }; // A pointer value. diff --git a/executable_semantics/syntax/parser.ypp b/executable_semantics/syntax/parser.ypp index a80d09c4ef40..d4f2a83fa392 100644 --- a/executable_semantics/syntax/parser.ypp +++ b/executable_semantics/syntax/parser.ypp @@ -101,12 +101,12 @@ void Carbon::Parser::error(const location_type&, const std::string& message) { %type >> function_declaration %type >> function_definition %type >> declaration_list -%type statement -%type if_statement -%type optional_else +%type >> statement +%type >> if_statement +%type >> optional_else %type , bool>>> return_expression -%type block -%type statement_list +%type >> block +%type >> statement_list %type >> expression %type > generic_binding %type > deduced_params @@ -131,8 +131,8 @@ void Carbon::Parser::error(const location_type&, const std::string& message) { %type > paren_pattern_contents %type >>> alternative %type >>> alternative_list -%type , const Statement*>*> clause -%type , const Statement*>>*> clause_list +%type , Ptr>*> clause +%type , Ptr>>*> clause_list %token END_OF_FILE 0 %token AND %token OR @@ -408,61 +408,61 @@ maybe_empty_tuple_pattern: ; clause: CASE pattern DBLARROW statement - { $$ = global_arena->RawNew, const Statement*>>($2, $4); } + { $$ = global_arena->RawNew, Ptr>>($2, $4); } | DEFAULT DBLARROW statement { auto vp = global_arena->New( context.SourceLoc(), std::nullopt, global_arena->New(context.SourceLoc())); - $$ = global_arena->RawNew, const Statement*>>(vp, $3); + $$ = global_arena->RawNew, Ptr>>(vp, $3); } ; clause_list: // Empty { $$ = global_arena->RawNew, const Statement*>>>(); + std::pair, Ptr>>>(); } | clause clause_list { $$ = $2; $$->push_front(*$1); } ; statement: expression "=" expression ";" - { $$ = global_arena->RawNew(context.SourceLoc(), $1, $3); } + { $$ = global_arena->New(context.SourceLoc(), $1, $3); } | VAR pattern "=" expression ";" - { $$ = global_arena->RawNew(context.SourceLoc(), $2, $4); } + { $$ = global_arena->New(context.SourceLoc(), $2, $4); } | expression ";" - { $$ = global_arena->RawNew(context.SourceLoc(), $1); } + { $$ = global_arena->New(context.SourceLoc(), $1); } | if_statement { $$ = $1; } | WHILE "(" expression ")" block - { $$ = global_arena->RawNew(context.SourceLoc(), $3, $5); } + { $$ = global_arena->New(context.SourceLoc(), $3, $5); } | BREAK ";" - { $$ = global_arena->RawNew(context.SourceLoc()); } + { $$ = global_arena->New(context.SourceLoc()); } | CONTINUE ";" - { $$ = global_arena->RawNew(context.SourceLoc()); } + { $$ = global_arena->New(context.SourceLoc()); } | RETURN return_expression ";" { auto [return_exp, is_omitted_exp] = $2.Release(); - $$ = global_arena->RawNew(context.SourceLoc(), return_exp, is_omitted_exp); + $$ = global_arena->New(context.SourceLoc(), return_exp, is_omitted_exp); } | block { $$ = $1; } | MATCH "(" expression ")" "{" clause_list "}" - { $$ = global_arena->RawNew(context.SourceLoc(), $3, $6); } + { $$ = global_arena->New(context.SourceLoc(), $3, $6); } | CONTINUATION identifier statement - { $$ = global_arena->RawNew(context.SourceLoc(), $2, $3); } + { $$ = global_arena->New(context.SourceLoc(), $2, $3); } | RUN expression ";" - { $$ = global_arena->RawNew(context.SourceLoc(), $2); } + { $$ = global_arena->New(context.SourceLoc(), $2); } | AWAIT ";" - { $$ = global_arena->RawNew(context.SourceLoc()); } + { $$ = global_arena->New(context.SourceLoc()); } ; if_statement: IF "(" expression ")" block optional_else - { $$ = global_arena->RawNew(context.SourceLoc(), $3, $5, $6); } + { $$ = global_arena->New(context.SourceLoc(), $3, $5, $6); } ; optional_else: // Empty - { $$ = 0; } + { $$ = std::nullopt; } | ELSE if_statement { $$ = $2; } | ELSE block @@ -476,13 +476,13 @@ return_expression: ; statement_list: // Empty - { $$ = 0; } + { $$ = std::nullopt; } | statement statement_list - { $$ = global_arena->RawNew(context.SourceLoc(), $1, $2); } + { $$ = global_arena->New(context.SourceLoc(), $1, $2); } ; block: "{" statement_list "}" - { $$ = global_arena->RawNew(context.SourceLoc(), $2); } + { $$ = global_arena->New(context.SourceLoc(), $2); } ; return_type: // Empty @@ -532,7 +532,7 @@ function_definition: $$ = global_arena->New( context.SourceLoc(), $2, $3, $4, global_arena->New(context.SourceLoc()), true, - global_arena->RawNew(context.SourceLoc(), $6, true)); + global_arena->New(context.SourceLoc(), $6, true)); } ; function_declaration: @@ -542,7 +542,7 @@ function_declaration: $$ = global_arena->New( context.SourceLoc(), $2, $3, $4, global_arena->New(return_exp), - is_omitted_exp, nullptr); + is_omitted_exp, std::nullopt); } ; variable_declaration: identifier ":" pattern diff --git a/executable_semantics/syntax/syntax_helpers.cpp b/executable_semantics/syntax/syntax_helpers.cpp index 12a4cc07b83c..ba601de3f36a 100644 --- a/executable_semantics/syntax/syntax_helpers.cpp +++ b/executable_semantics/syntax/syntax_helpers.cpp @@ -22,11 +22,11 @@ static void AddIntrinsics(std::list>* fs) { loc, "format_str", global_arena->New( global_arena->New(loc))))}; - auto* print_return = global_arena->RawNew( - loc, - global_arena->New( - IntrinsicExpression::IntrinsicKind::Print), - false); + auto print_return = + global_arena->New(loc, + global_arena->New( + IntrinsicExpression::IntrinsicKind::Print), + false); auto print = global_arena->New( global_arena->New( loc, "Print", std::vector(),