Implement auto return types, removing => returns (#850)

This implements #826, I think covering everything important there.

Regarding ReturnTypeContext, I broke that out because it started feeling like a significant number of args to be passing around, and I think this makes the association inside type checking clearer.

Co-authored-by: Geoff Romer <gromer@google.com>
This commit is contained in:
Jon Meow
2021-09-27 11:17:24 -07:00
committed by GitHub
co-authored by Geoff Romer
parent 04ab30f231
commit 5aa958345b
11 changed files with 164 additions and 46 deletions
@@ -22,9 +22,17 @@
using llvm::cast;
using llvm::dyn_cast;
using llvm::isa;
namespace Carbon {
TypeChecker::ReturnTypeContext::ReturnTypeContext(
Nonnull<const Value*> orig_return_type, bool is_omitted)
: is_auto_(isa<AutoType>(orig_return_type)),
deduced_return_type_(is_auto_ ? std::nullopt
: std::optional(orig_return_type)),
is_omitted_(is_omitted) {}
void PrintTypeEnv(TypeEnv types, llvm::raw_ostream& out) {
llvm::ListSeparator sep;
for (const auto& [name, type] : types) {
@@ -622,18 +630,17 @@ auto TypeChecker::TypeCheckPattern(
auto TypeChecker::TypeCheckCase(Nonnull<const Value*> expected,
Nonnull<Pattern*> pat, Nonnull<Statement*> body,
TypeEnv types, Env values,
Nonnull<const Value*>& ret_type,
bool is_omitted_ret_type)
Nonnull<ReturnTypeContext*> return_type_context)
-> std::pair<Nonnull<Pattern*>, Nonnull<Statement*>> {
auto pat_res = TypeCheckPattern(pat, types, values, expected);
auto res =
TypeCheckStmt(body, pat_res.types, values, ret_type, is_omitted_ret_type);
auto res = TypeCheckStmt(body, pat_res.types, values, return_type_context);
return std::make_pair(pat, res.stmt);
}
auto TypeChecker::TypeCheckStmt(Nonnull<Statement*> s, TypeEnv types,
Env values, Nonnull<const Value*>& ret_type,
bool is_omitted_ret_type) -> TCStatement {
Env values,
Nonnull<ReturnTypeContext*> return_type_context)
-> TCStatement {
switch (s->Tag()) {
case Statement::Kind::Match: {
auto& match = cast<Match>(*s);
@@ -644,7 +651,7 @@ auto TypeChecker::TypeCheckStmt(Nonnull<Statement*> s, TypeEnv types,
for (auto& clause : match.Clauses()) {
new_clauses.push_back(TypeCheckCase(res_type, clause.first,
clause.second, types, values,
ret_type, is_omitted_ret_type));
return_type_context));
}
auto new_s = arena->New<Match>(s->SourceLoc(), res.exp, new_clauses);
return TCStatement(new_s, types);
@@ -654,8 +661,8 @@ auto TypeChecker::TypeCheckStmt(Nonnull<Statement*> s, TypeEnv types,
auto cnd_res = TypeCheckExp(while_stmt.Cond(), types, values);
ExpectType(s->SourceLoc(), "condition of `while`", arena->New<BoolType>(),
cnd_res.type);
auto body_res = TypeCheckStmt(while_stmt.Body(), types, values, ret_type,
is_omitted_ret_type);
auto body_res =
TypeCheckStmt(while_stmt.Body(), types, values, return_type_context);
auto new_s =
arena->New<While>(s->SourceLoc(), cnd_res.exp, body_res.stmt);
return TCStatement(new_s, types);
@@ -666,8 +673,8 @@ auto TypeChecker::TypeCheckStmt(Nonnull<Statement*> s, TypeEnv types,
case Statement::Kind::Block: {
auto& block = cast<Block>(*s);
if (block.Stmt()) {
auto stmt_res = TypeCheckStmt(*block.Stmt(), types, values, ret_type,
is_omitted_ret_type);
auto stmt_res =
TypeCheckStmt(*block.Stmt(), types, values, return_type_context);
return TCStatement(arena->New<Block>(s->SourceLoc(), stmt_res.stmt),
types);
} else {
@@ -685,13 +692,13 @@ auto TypeChecker::TypeCheckStmt(Nonnull<Statement*> s, TypeEnv types,
}
case Statement::Kind::Sequence: {
auto& seq = cast<Sequence>(*s);
auto stmt_res = TypeCheckStmt(seq.Stmt(), types, values, ret_type,
is_omitted_ret_type);
auto stmt_res =
TypeCheckStmt(seq.Stmt(), types, values, return_type_context);
auto checked_types = stmt_res.types;
std::optional<Nonnull<Statement*>> next_stmt;
if (seq.Next()) {
auto next_res = TypeCheckStmt(*seq.Next(), checked_types, values,
ret_type, is_omitted_ret_type);
return_type_context);
next_stmt = next_res.stmt;
checked_types = next_res.types;
}
@@ -720,12 +727,12 @@ auto TypeChecker::TypeCheckStmt(Nonnull<Statement*> s, TypeEnv types,
auto cnd_res = TypeCheckExp(if_stmt.Cond(), types, values);
ExpectType(s->SourceLoc(), "condition of `if`", arena->New<BoolType>(),
cnd_res.type);
auto then_res = TypeCheckStmt(if_stmt.ThenStmt(), types, values, ret_type,
is_omitted_ret_type);
auto then_res =
TypeCheckStmt(if_stmt.ThenStmt(), types, values, return_type_context);
std::optional<Nonnull<Statement*>> else_stmt;
if (if_stmt.ElseStmt()) {
auto else_res = TypeCheckStmt(*if_stmt.ElseStmt(), types, values,
ret_type, is_omitted_ret_type);
return_type_context);
else_stmt = else_res.stmt;
}
auto new_s =
@@ -735,17 +742,24 @@ auto TypeChecker::TypeCheckStmt(Nonnull<Statement*> s, TypeEnv types,
case Statement::Kind::Return: {
auto& ret = cast<Return>(*s);
auto res = TypeCheckExp(ret.Exp(), types, values);
if (ret_type->Tag() == Value::Kind::AutoType) {
// The following infers the return type from the first 'return'
// statement. This will get more difficult with subtyping, when we
// should infer the least-upper bound of all the 'return' statements.
ret_type = res.type;
if (return_type_context->is_auto()) {
if (return_type_context->deduced_return_type()) {
// Only one return is allowed when the return type is `auto`.
FATAL_COMPILATION_ERROR(s->SourceLoc())
<< "Only one return is allowed in a function with an `auto` "
"return type.";
} else {
// Infer the auto return from the first `return` statement.
return_type_context->set_deduced_return_type(res.type);
}
} else {
ExpectType(s->SourceLoc(), "return", ret_type, res.type);
ExpectType(s->SourceLoc(), "return",
*return_type_context->deduced_return_type(), res.type);
}
if (ret.IsOmittedExp() != is_omitted_ret_type) {
if (ret.IsOmittedExp() != return_type_context->is_omitted()) {
FATAL_COMPILATION_ERROR(s->SourceLoc())
<< *s << " should" << (is_omitted_ret_type ? " not" : "")
<< *s << " should"
<< (return_type_context->is_omitted() ? " not" : "")
<< " provide a return value, to match the function's signature.";
}
return TCStatement(
@@ -754,8 +768,8 @@ auto TypeChecker::TypeCheckStmt(Nonnull<Statement*> s, TypeEnv types,
}
case Statement::Kind::Continuation: {
auto& cont = cast<Continuation>(*s);
TCStatement body_result = TypeCheckStmt(cont.Body(), types, values,
ret_type, is_omitted_ret_type);
TCStatement body_result =
TypeCheckStmt(cont.Body(), types, values, return_type_context);
auto new_continuation = arena->New<Continuation>(
s->SourceLoc(), cont.ContinuationVariable(), body_result.stmt);
types.Set(cont.ContinuationVariable(), arena->New<ContinuationType>());
@@ -875,9 +889,13 @@ auto TypeChecker::TypeCheckFunDef(FunctionDefinition* f, TypeEnv types,
}
std::optional<Nonnull<Statement*>> body_stmt;
if (f->body()) {
auto res = TypeCheckStmt(*f->body(), param_res.types, values, return_type,
f->is_omitted_return_type());
ReturnTypeContext return_type_context(return_type,
f->is_omitted_return_type());
auto res = TypeCheckStmt(*f->body(), param_res.types, values,
&return_type_context);
body_stmt = res.stmt;
// Save the return type in case it changed.
return_type = *return_type_context.deduced_return_type();
}
auto body = CheckOrEnsureReturn(body_stmt, f->is_omitted_return_type(),
f->source_loc());
@@ -38,6 +38,37 @@ class TypeChecker {
auto TopLevel(std::vector<Nonnull<Declaration*>>* fs) -> TypeCheckContext;
private:
// Context about the return type, which may be updated during type checking.
class ReturnTypeContext {
public:
// If orig_return_type is auto, deduced_return_type_ will be nullopt;
// otherwise, it's orig_return_type. is_auto_ is set accordingly.
ReturnTypeContext(Nonnull<const Value*> orig_return_type, bool is_omitted);
auto is_auto() const -> bool { return is_auto_; }
auto deduced_return_type() const -> std::optional<Nonnull<const Value*>> {
return deduced_return_type_;
}
void set_deduced_return_type(Nonnull<const Value*> type) {
deduced_return_type_ = type;
}
auto is_omitted() const -> bool { return is_omitted_; }
private:
// Indicates an `auto` return type, as in `fn Foo() -> auto { return 0; }`.
const bool is_auto_;
// The actual return type. May be nullopt for an `auto` return type that has
// yet to be determined.
std::optional<Nonnull<const Value*>> deduced_return_type_;
// Indicates the return type was omitted and is implicitly the empty tuple,
// as in `fn Foo() {}`.
const bool is_omitted_;
};
struct TCExpression {
TCExpression(Nonnull<Expression*> e, Nonnull<const Value*> t, TypeEnv types)
: exp(e), type(t), types(types) {}
@@ -90,7 +121,7 @@ class TypeChecker {
// type is "auto", then the return type is inferred from the first return
// statement.
auto TypeCheckStmt(Nonnull<Statement*> s, TypeEnv types, Env values,
Nonnull<const Value*>& ret_type, bool is_omitted_ret_type)
Nonnull<ReturnTypeContext*> return_type_context)
-> TCStatement;
auto TypeCheckFunDef(FunctionDefinition* f, TypeEnv types, Env values)
@@ -98,7 +129,7 @@ class TypeChecker {
auto TypeCheckCase(Nonnull<const Value*> expected, Nonnull<Pattern*> pat,
Nonnull<Statement*> body, TypeEnv types, Env values,
Nonnull<const Value*>& ret_type, bool is_omitted_ret_type)
Nonnull<ReturnTypeContext*> return_type_context)
-> std::pair<Nonnull<Pattern*>, Nonnull<Statement*>>;
auto TypeOfFunDef(TypeEnv types, Env values, FunctionDefinition* fun_def)