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 <gromer@google.com>
This commit is contained in:
Jon Meow
2021-08-30 14:41:14 -07:00
committed by GitHub
co-authored by Geoff Romer
parent ed2d171703
commit 36ed79dc25
12 changed files with 166 additions and 138 deletions
+3 -3
View File
@@ -117,17 +117,17 @@ class PatternAction : public Action {
class StatementAction : public Action {
public:
explicit StatementAction(const Statement* stmt)
explicit StatementAction(Ptr<const Statement> 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<const Statement> { return stmt; }
private:
const Statement* stmt;
Ptr<const Statement> stmt;
};
} // namespace Carbon
@@ -812,8 +812,7 @@ auto IsBlockAct(Ptr<Action> act) -> bool {
Transition StepStmt() {
Ptr<Frame> frame = state->stack.Top();
Ptr<Action> act = frame->todo.Top();
const Statement* stmt = cast<StatementAction>(*act).Stmt();
CHECK(stmt != nullptr) << "null statement!";
Ptr<const Statement> stmt = cast<StatementAction>(*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<Scope>(values, vars));
const Statement* body_block =
global_arena->RawNew<Block>(stmt->SourceLoc(), c->second);
auto body_act = global_arena->New<StatementAction>(body_block);
auto body_act = global_arena->New<StatementAction>(
global_arena->New<Block>(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<Block>(*stmt);
if (block.Stmt() != nullptr) {
if (block.Stmt()) {
frame->scopes.Push(global_arena->New<Scope>(CurrentEnv(state)));
return Spawn{global_arena->New<StatementAction>(block.Stmt())};
return Spawn{global_arena->New<StatementAction>(*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<StatementAction>(cast<If>(*stmt).ElseStmt())};
global_arena->New<StatementAction>(*cast<If>(*stmt).ElseStmt())};
} else {
return Done{};
}
@@ -1030,9 +1028,9 @@ Transition StepStmt() {
if (act->Pos() == 0) {
return Spawn{global_arena->New<StatementAction>(seq.Stmt())};
} else {
if (seq.Next() != nullptr) {
return Delegate{
global_arena->New<StatementAction>(cast<Sequence>(*stmt).Next())};
if (seq.Next()) {
return Delegate{global_arena->New<StatementAction>(
*cast<Sequence>(*stmt).Next())};
} else {
return Done{};
}
@@ -1046,7 +1044,7 @@ Transition StepStmt() {
Stack<Ptr<Scope>>(global_arena->New<Scope>(CurrentEnv(state)));
Stack<Ptr<Action>> todo;
todo.Push(global_arena->New<StatementAction>(
global_arena->RawNew<Return>(stmt->SourceLoc())));
global_arena->New<Return>(stmt->SourceLoc())));
todo.Push(
global_arena->New<StatementAction>(cast<Continuation>(*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<StatementAction>(
global_arena->RawNew<ExpressionStatement>(
global_arena->New<ExpressionStatement>(
stmt->SourceLoc(),
global_arena->New<TupleLiteral>(stmt->SourceLoc())));
frame->todo.Push(ignore_result);
@@ -1172,8 +1170,9 @@ struct DoTransition {
params.push_back(name);
}
auto scopes = Stack<Ptr<Scope>>(global_arena->New<Scope>(values, params));
CHECK(call.function->Body()) << "Calling a function that's missing a body";
auto todo = Stack<Ptr<Action>>(
global_arena->New<StatementAction>(call.function->Body()));
global_arena->New<StatementAction>(*call.function->Body()));
auto frame = global_arena->New<Frame>(call.function->Name(), scopes, todo);
state->stack.Push(frame);
}
+67 -51
View File
@@ -657,9 +657,9 @@ auto TypeCheckPattern(Ptr<const Pattern> p, TypeEnv types, Env values,
}
static auto TypecheckCase(const Value* expected, Ptr<const Pattern> pat,
const Statement* body, TypeEnv types, Env values,
Ptr<const Statement> body, TypeEnv types, Env values,
const Value*& ret_type, bool is_omitted_ret_type)
-> std::pair<Ptr<const Pattern>, const Statement*> {
-> std::pair<Ptr<const Pattern>, Ptr<const Statement>> {
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<const Pattern> 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<const Statement> 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<Match>(*s);
auto res = TypeCheckExp(match.Exp(), types, values);
auto res_type = res.type;
auto new_clauses = global_arena->RawNew<
std::list<std::pair<Ptr<const Pattern>, const Statement*>>>();
std::list<std::pair<Ptr<const Pattern>, Ptr<const Statement>>>>();
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<Match>(s->SourceLoc(), res.exp, new_clauses);
auto new_s =
global_arena->New<Match>(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<BoolType>(), cnd_res.type);
auto body_res = TypeCheckStmt(while_stmt.Body(), types, values, ret_type,
is_omitted_ret_type);
auto new_s = global_arena->RawNew<While>(s->SourceLoc(), cnd_res.exp,
body_res.stmt);
auto new_s =
global_arena->New<While>(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<Block>(*s).Stmt(), types, values,
ret_type, is_omitted_ret_type);
return TCStatement(
global_arena->RawNew<Block>(s->SourceLoc(), stmt_res.stmt), types);
const auto& block = cast<Block>(*s);
if (block.Stmt()) {
auto stmt_res = TypeCheckStmt(*block.Stmt(), types, values, ret_type,
is_omitted_ret_type);
return TCStatement(
global_arena->New<Block>(s->SourceLoc(), stmt_res.stmt), types);
} else {
return TCStatement(s, types);
}
}
case Statement::Kind::VariableDefinition: {
const auto& var = cast<VariableDefinition>(*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<VariableDefinition>(
s->SourceLoc(), var.Pat(), res.exp);
auto new_s = global_arena->New<VariableDefinition>(s->SourceLoc(),
var.Pat(), res.exp);
return TCStatement(new_s, lhs_res.types);
}
case Statement::Kind::Sequence: {
const auto& seq = cast<Sequence>(*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<Sequence>(
s->SourceLoc(), stmt_res.stmt, next_res.stmt),
types3);
auto checked_types = stmt_res.types;
std::optional<Ptr<const Statement>> 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<Sequence>(s->SourceLoc(), stmt_res.stmt, next_stmt),
checked_types);
}
case Statement::Kind::Assign: {
const auto& assign = cast<Assign>(*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<Assign>(s->SourceLoc(), lhs_res.exp,
rhs_res.exp);
auto new_s =
global_arena->New<Assign>(s->SourceLoc(), lhs_res.exp, rhs_res.exp);
return TCStatement(new_s, lhs_res.types);
}
case Statement::Kind::ExpressionStatement: {
auto res =
TypeCheckExp(cast<ExpressionStatement>(*s).Exp(), types, values);
auto new_s =
global_arena->RawNew<ExpressionStatement>(s->SourceLoc(), res.exp);
global_arena->New<ExpressionStatement>(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<BoolType>(), 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<If>(s->SourceLoc(), cnd_res.exp,
then_res.stmt, else_res.stmt);
std::optional<Ptr<const Statement>> 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<If>(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<Return>(s->SourceLoc(), res.exp,
ret.IsOmittedExp()),
return TCStatement(global_arena->New<Return>(s->SourceLoc(), res.exp,
ret.IsOmittedExp()),
types);
}
case Statement::Kind::Continuation: {
const auto& cont = cast<Continuation>(*s);
TCStatement body_result = TypeCheckStmt(cont.Body(), types, values,
ret_type, is_omitted_ret_type);
const Statement* new_continuation = global_arena->RawNew<Continuation>(
auto new_continuation = global_arena->New<Continuation>(
s->SourceLoc(), cont.ContinuationVariable(), body_result.stmt);
types.Set(cont.ContinuationVariable(),
global_arena->RawNew<ContinuationType>());
@@ -803,8 +813,8 @@ auto TypeCheckStmt(const Statement* s, TypeEnv types, Env values,
ExpectType(s->SourceLoc(), "argument of `run`",
global_arena->RawNew<ContinuationType>(),
argument_result.type);
const Statement* new_run =
global_arena->RawNew<Run>(s->SourceLoc(), argument_result.exp);
auto new_run =
global_arena->New<Run>(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<Ptr<const Statement>> opt_stmt,
bool omitted_ret_type, SourceLocation loc)
-> Ptr<const Statement> {
if (!opt_stmt) {
if (omitted_ret_type) {
return global_arena->RawNew<Return>(loc);
return global_arena->New<Return>(loc);
} else {
FATAL_COMPILATION_ERROR(loc)
<< "control-flow reaches end of function that provides a `->` return "
"type without reaching a return statement";
}
}
Ptr<const Statement> stmt = *opt_stmt;
switch (stmt->Tag()) {
case Statement::Kind::Match: {
const auto& match = cast<Match>(*stmt);
auto new_clauses = global_arena->RawNew<
std::list<std::pair<Ptr<const Pattern>, const Statement*>>>();
std::list<std::pair<Ptr<const Pattern>, Ptr<const Statement>>>>();
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<Match>(stmt->SourceLoc(), match.Exp(),
new_clauses);
return global_arena->New<Match>(stmt->SourceLoc(), match.Exp(),
new_clauses);
}
case Statement::Kind::Block:
return global_arena->RawNew<Block>(
return global_arena->New<Block>(
stmt->SourceLoc(),
CheckOrEnsureReturn(cast<Block>(*stmt).Stmt(), omitted_ret_type,
stmt->SourceLoc()));
case Statement::Kind::If: {
const auto& if_stmt = cast<If>(*stmt);
return global_arena->RawNew<If>(
return global_arena->New<If>(
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<Sequence>(*stmt);
if (seq.Next()) {
return global_arena->RawNew<Sequence>(
return global_arena->New<Sequence>(
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<Sequence>(
stmt->SourceLoc(), stmt, global_arena->RawNew<Return>(loc));
return global_arena->New<Sequence>(stmt->SourceLoc(), stmt,
global_arena->New<Return>(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<IntType>(), 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<Ptr<const Statement>> 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<FunctionDefinition>(
f->source_location, f->name, f->deduced_parameters, f->param_pattern,
+3 -3
View File
@@ -34,9 +34,9 @@ struct TCPattern {
};
struct TCStatement {
TCStatement(const Statement* s, TypeEnv types) : stmt(s), types(types) {}
TCStatement(Ptr<const Statement> s, TypeEnv types) : stmt(s), types(types) {}
const Statement* stmt;
Ptr<const Statement> stmt;
TypeEnv types;
};
@@ -52,7 +52,7 @@ auto TypeCheckExp(Ptr<const Expression> e, TypeEnv types, Env values)
auto TypeCheckPattern(Ptr<const Pattern> p, TypeEnv types, Env values,
const Value* expected) -> TCPattern;
auto TypeCheckStmt(const Statement* s, TypeEnv types, Env values,
auto TypeCheckStmt(Ptr<const Statement> s, TypeEnv types, Env values,
const Value*& ret_type, bool is_omitted_ret_type)
-> TCStatement;
+8 -2
View File
@@ -399,8 +399,14 @@ auto ValueEqual(const Value* v1, const Value* v2, SourceLocation loc) -> bool {
return cast<BoolValue>(*v1).Val() == cast<BoolValue>(*v2).Val();
case Value::Kind::PointerValue:
return cast<PointerValue>(*v1).Val() == cast<PointerValue>(*v2).Val();
case Value::Kind::FunctionValue:
return cast<FunctionValue>(*v1).Body() == cast<FunctionValue>(*v2).Body();
case Value::Kind::FunctionValue: {
std::optional<Ptr<const Statement>> body1 =
cast<FunctionValue>(*v1).Body();
std::optional<Ptr<const Statement>> body2 =
cast<FunctionValue>(*v2).Body();
return body1.has_value() == body2.has_value() &&
(!body1.has_value() || *body1 == *body2);
}
case Value::Kind::TupleValue:
return FieldsValueEqual(cast<TupleValue>(*v1).Elements(),
cast<TupleValue>(*v2).Elements(), loc);
+4 -3
View File
@@ -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<Ptr<const Statement>> 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<Ptr<const Statement>> { return body; }
private:
std::string name;
const Value* param;
const Statement* body;
std::optional<Ptr<const Statement>> body;
};
// A pointer value.