diff --git a/executable_semantics/interpreter/action.h b/executable_semantics/interpreter/action.h index 5d341f78fbd5..d8a849dd4bfb 100644 --- a/executable_semantics/interpreter/action.h +++ b/executable_semantics/interpreter/action.h @@ -40,7 +40,7 @@ class Action { // Results from a subexpression. auto Results() const -> const std::vector& { return results; } - void IncrementPos() { ++pos; } + void SetPos(int pos) { this->pos = pos; } void AddResult(const Value* result) { results.push_back(result); } diff --git a/executable_semantics/interpreter/interpreter.cpp b/executable_semantics/interpreter/interpreter.cpp index 749aec54ea6b..88f5982fe3db 100644 --- a/executable_semantics/interpreter/interpreter.cpp +++ b/executable_semantics/interpreter/interpreter.cpp @@ -722,11 +722,12 @@ static auto IsWhileAct(Ptr act) -> bool { } } -static auto IsBlockAct(Ptr act) -> bool { +static auto HasLocalScope(Ptr act) -> bool { switch (act->Tag()) { case Action::Kind::StatementAction: switch (cast(*act).Stmt()->Tag()) { case Statement::Kind::Block: + case Statement::Kind::Match: return true; default: return false; @@ -746,12 +747,13 @@ auto Interpreter::StepStmt() -> Transition { llvm::outs() << " --->\n"; } switch (stmt->Tag()) { - case Statement::Kind::Match: + case Statement::Kind::Match: { + const auto& match_stmt = cast(*stmt); if (act->Pos() == 0) { // { { (match (e) ...) :: C, E, F} :: S, H} // -> { { e :: (match ([]) ...) :: C, E, F} :: S, H} - return Spawn{ - global_arena->New(cast(*stmt).Exp())}; + frame->scopes.Push(global_arena->New(CurrentEnv())); + return Spawn{global_arena->New(match_stmt.Exp())}; } else { // Regarding act->Pos(): // * odd: start interpreting the pattern of a clause @@ -763,11 +765,12 @@ auto Interpreter::StepStmt() -> Transition { // * 2: the pattern for clause 1 // * ... auto clause_num = (act->Pos() - 1) / 2; - if (clause_num >= - static_cast(cast(*stmt).Clauses()->size())) { + if (clause_num >= static_cast(match_stmt.Clauses()->size())) { + DeallocateScope(frame->scopes.Top()); + frame->scopes.Pop(); return Done{}; } - auto c = cast(*stmt).Clauses()->begin(); + auto c = match_stmt.Clauses()->begin(); std::advance(c, clause_num); if (act->Pos() % 2 == 1) { @@ -780,31 +783,20 @@ auto Interpreter::StepStmt() -> Transition { auto pat = act->Results()[clause_num + 1]; std::optional matches = PatternMatch(pat, v, stmt->SourceLoc()); if (matches) { // we have a match, start the body - Env values = CurrentEnv(); - std::list vars; + // Ensure we don't process any more clauses. + act->SetPos(2 * match_stmt.Clauses()->size() + 1); + for (const auto& [name, value] : *matches) { - values.Set(name, value); - vars.push_back(name); + frame->scopes.Top()->values.Set(name, value); + frame->scopes.Top()->locals.push_back(name); } - frame->scopes.Push(global_arena->New(values, vars)); - 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); - frame->todo.Push(global_arena->New(c->second)); - return ManualTransition{}; + return Spawn{global_arena->New(c->second)}; } else { - // this case did not match, moving on - int next_clause_num = act->Pos() / 2; - if (next_clause_num == - static_cast(cast(*stmt).Clauses()->size())) { - return Done{}; - } return RunAgain{}; } } } + } case Statement::Kind::While: if (act->Pos() % 2 == 0) { // { { (while (e) s) :: C, E, F} :: S, H} @@ -1050,7 +1042,8 @@ class Interpreter::DoTransition { void operator()(const Spawn& spawn) { Ptr frame = interpreter->stack.Top(); - frame->todo.Top()->IncrementPos(); + Ptr action = frame->todo.Top(); + action->SetPos(action->Pos() + 1); frame->todo.Push(spawn.child); } @@ -1061,13 +1054,14 @@ class Interpreter::DoTransition { } void operator()(const RunAgain&) { - interpreter->stack.Top()->todo.Top()->IncrementPos(); + Ptr action = interpreter->stack.Top()->todo.Top(); + action->SetPos(action->Pos() + 1); } void operator()(const UnwindTo& unwind_to) { Ptr frame = interpreter->stack.Top(); while (frame->todo.Top() != unwind_to.new_top) { - if (IsBlockAct(frame->todo.Top())) { + if (HasLocalScope(frame->todo.Top())) { interpreter->DeallocateScope(frame->scopes.Top()); frame->scopes.Pop(); }