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
@@ -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";
@@ -26,7 +26,8 @@ struct FunctionDefinition {
std::vector<GenericBinding> deduced_params,
Ptr<const TuplePattern> param_pattern,
Ptr<const Pattern> return_type,
bool is_omitted_return_type, const Statement* body)
bool is_omitted_return_type,
std::optional<Ptr<const Statement>> body)
: source_location(source_location),
name(std::move(name)),
deduced_parameters(deduced_params),
@@ -45,7 +46,7 @@ struct FunctionDefinition {
Ptr<const TuplePattern> param_pattern;
Ptr<const Pattern> return_type;
bool is_omitted_return_type;
const Statement* body;
std::optional<Ptr<const Statement>> body;
};
} // namespace Carbon
+3 -3
View File
@@ -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";
}
+28 -23
View File
@@ -109,8 +109,9 @@ class VariableDefinition : public Statement {
class If : public Statement {
public:
If(SourceLocation loc, Ptr<const Expression> cond, const Statement* then_stmt,
const Statement* else_stmt)
If(SourceLocation loc, Ptr<const Expression> cond,
Ptr<const Statement> then_stmt,
std::optional<Ptr<const Statement>> else_stmt)
: Statement(Kind::If, loc),
cond(cond),
then_stmt(then_stmt),
@@ -121,13 +122,15 @@ class If : public Statement {
}
auto Cond() const -> Ptr<const Expression> { return cond; }
auto ThenStmt() const -> const Statement* { return then_stmt; }
auto ElseStmt() const -> const Statement* { return else_stmt; }
auto ThenStmt() const -> Ptr<const Statement> { return then_stmt; }
auto ElseStmt() const -> std::optional<Ptr<const Statement>> {
return else_stmt;
}
private:
Ptr<const Expression> cond;
const Statement* then_stmt;
const Statement* else_stmt;
Ptr<const Statement> then_stmt;
std::optional<Ptr<const Statement>> 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<const Statement> stmt,
std::optional<Ptr<const Statement>> 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<const Statement> { return stmt; }
auto Next() const -> std::optional<Ptr<const Statement>> { return next; }
private:
const Statement* stmt;
const Statement* next;
Ptr<const Statement> stmt;
std::optional<Ptr<const Statement>> next;
};
class Block : public Statement {
public:
Block(SourceLocation loc, const Statement* stmt)
Block(SourceLocation loc, std::optional<Ptr<const Statement>> 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<Ptr<const Statement>> { return stmt; }
private:
const Statement* stmt;
std::optional<Ptr<const Statement>> stmt;
};
class While : public Statement {
public:
While(SourceLocation loc, Ptr<const Expression> cond, const Statement* body)
While(SourceLocation loc, Ptr<const Expression> cond,
Ptr<const Statement> 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<const Expression> { return cond; }
auto Body() const -> const Statement* { return body; }
auto Body() const -> Ptr<const Statement> { return body; }
private:
Ptr<const Expression> cond;
const Statement* body;
Ptr<const Statement> body;
};
class Break : public Statement {
@@ -221,7 +226,7 @@ class Continue : public Statement {
class Match : public Statement {
public:
Match(SourceLocation loc, Ptr<const Expression> exp,
std::list<std::pair<Ptr<const Pattern>, const Statement*>>* clauses)
std::list<std::pair<Ptr<const Pattern>, Ptr<const Statement>>>* 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<const Expression> { return exp; }
auto Clauses() const
-> const std::list<std::pair<Ptr<const Pattern>, const Statement*>>* {
-> const std::list<std::pair<Ptr<const Pattern>, Ptr<const Statement>>>* {
return clauses;
}
private:
Ptr<const Expression> exp;
std::list<std::pair<Ptr<const Pattern>, const Statement*>>* clauses;
std::list<std::pair<Ptr<const Pattern>, Ptr<const Statement>>>* 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<const Statement> 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<const Statement> { return body; }
private:
std::string continuation_variable;
const Statement* body;
Ptr<const Statement> body;
};
// A run statement.