From a99d882223f7b341071617d3bdc1ac92aa158e09 Mon Sep 17 00:00:00 2001 From: Geoff Romer Date: Mon, 11 Oct 2021 16:37:39 -0700 Subject: [PATCH] Perform type-checking in place (#867) --- .../ast/function_definition.h | 1 + .../interpreter/exec_program.cpp | 7 +- .../interpreter/interpreter.cpp | 13 +- .../interpreter/interpreter.h | 5 +- .../interpreter/type_checker.cpp | 437 +++++++----------- .../interpreter/type_checker.h | 45 +- 6 files changed, 189 insertions(+), 319 deletions(-) diff --git a/executable_semantics/ast/function_definition.h b/executable_semantics/ast/function_definition.h index cd57588c16ec..022b9bdef74c 100644 --- a/executable_semantics/ast/function_definition.h +++ b/executable_semantics/ast/function_definition.h @@ -49,6 +49,7 @@ class FunctionDefinition { auto param_pattern() const -> const TuplePattern& { return *param_pattern_; } auto param_pattern() -> TuplePattern& { return *param_pattern_; } auto return_type() const -> const Pattern& { return *return_type_; } + auto return_type() -> Pattern& { return *return_type_; } auto is_omitted_return_type() const -> bool { return is_omitted_return_type_; } diff --git a/executable_semantics/interpreter/exec_program.cpp b/executable_semantics/interpreter/exec_program.cpp index 86dcde83591f..76d87bbae587 100644 --- a/executable_semantics/interpreter/exec_program.cpp +++ b/executable_semantics/interpreter/exec_program.cpp @@ -49,14 +49,13 @@ void ExecProgram(Nonnull arena, AST ast) { TypeChecker::TypeCheckContext p = type_checker.TopLevel(&ast.declarations); TypeEnv top = p.types; Env ct_top = p.values; - std::vector> new_decls; for (const auto decl : ast.declarations) { - new_decls.push_back(type_checker.MakeTypeChecked(decl, top, ct_top)); + type_checker.TypeCheck(decl, top, ct_top); } if (tracing_output) { llvm::outs() << "\n"; llvm::outs() << "********** type checking complete **********\n"; - for (const auto decl : new_decls) { + for (const auto decl : ast.declarations) { llvm::outs() << *decl; } llvm::outs() << "********** starting execution **********\n"; @@ -66,7 +65,7 @@ void ExecProgram(Nonnull arena, AST ast) { Nonnull call_main = arena->New( source_loc, arena->New(source_loc, "main"), arena->New(source_loc)); - int result = Interpreter(arena).InterpProgram(new_decls, call_main); + int result = Interpreter(arena).InterpProgram(ast.declarations, call_main); llvm::outs() << "result: " << result << "\n"; } diff --git a/executable_semantics/interpreter/interpreter.cpp b/executable_semantics/interpreter/interpreter.cpp index c7887b583ac3..138ca5dabf7b 100644 --- a/executable_semantics/interpreter/interpreter.cpp +++ b/executable_semantics/interpreter/interpreter.cpp @@ -173,8 +173,7 @@ void Interpreter::InitEnv(const Declaration& d, Env* env) { } } -void Interpreter::InitGlobals( - const std::vector>& fs) { +void Interpreter::InitGlobals(llvm::ArrayRef> fs) { for (const auto d : fs) { InitEnv(*d, &globals); } @@ -1150,8 +1149,9 @@ class Interpreter::DoTransition { void Interpreter::Step() { Nonnull frame = stack.Top(); if (frame->todo.IsEmpty()) { - FATAL_RUNTIME_ERROR_NO_LINE() - << "fell off end of function " << frame->name << " without `return`"; + std::visit(DoTransition(this), + Transition{UnwindFunctionCall{TupleValue::Empty()}}); + return; } Nonnull act = frame->todo.Top(); @@ -1171,9 +1171,8 @@ void Interpreter::Step() { } // switch } -auto Interpreter::InterpProgram( - const std::vector>& fs, - Nonnull call_main) -> int { +auto Interpreter::InterpProgram(llvm::ArrayRef> fs, + Nonnull call_main) -> int { // Check that the interpreter is in a clean state. CHECK(globals.IsEmpty()); CHECK(stack.IsEmpty()); diff --git a/executable_semantics/interpreter/interpreter.h b/executable_semantics/interpreter/interpreter.h index ecb2d290089a..eb1b8fd1e1c3 100644 --- a/executable_semantics/interpreter/interpreter.h +++ b/executable_semantics/interpreter/interpreter.h @@ -17,6 +17,7 @@ #include "executable_semantics/interpreter/heap.h" #include "executable_semantics/interpreter/stack.h" #include "executable_semantics/interpreter/value.h" +#include "llvm/ADT/ArrayRef.h" namespace Carbon { @@ -28,7 +29,7 @@ class Interpreter { : arena(arena), globals(arena), heap(arena) {} // Interpret the whole program. - auto InterpProgram(const std::vector>& fs, + auto InterpProgram(llvm::ArrayRef> fs, Nonnull call_main) -> int; // Interpret an expression at compile-time. @@ -129,7 +130,7 @@ class Interpreter { // State transition for statements. auto StepStmt() -> Transition; - void InitGlobals(const std::vector>& fs); + void InitGlobals(llvm::ArrayRef> fs); auto CurrentEnv() -> Env; auto GetFromEnv(SourceLocation source_loc, const std::string& name) -> Address; diff --git a/executable_semantics/interpreter/type_checker.cpp b/executable_semantics/interpreter/type_checker.cpp index 78bcf7b1437f..736ace435df4 100644 --- a/executable_semantics/interpreter/type_checker.cpp +++ b/executable_semantics/interpreter/type_checker.cpp @@ -60,68 +60,52 @@ static void ExpectPointerType(SourceLocation source_loc, } } -auto TypeChecker::ReifyType(Nonnull t, SourceLocation source_loc) - -> Nonnull { - switch (t->kind()) { - case Value::Kind::IntType: - return arena->New(source_loc); - case Value::Kind::BoolType: - return arena->New(source_loc); - case Value::Kind::TypeType: - return arena->New(source_loc); - case Value::Kind::ContinuationType: - return arena->New(source_loc); - case Value::Kind::FunctionType: { - const auto& fn_type = cast(*t); - return arena->New( - source_loc, ReifyType(fn_type.Param(), source_loc), - ReifyType(fn_type.Ret(), source_loc), - /*is_omitted_return_type=*/false); - } - case Value::Kind::TupleValue: { - std::vector args; - for (const TupleElement& field : cast(*t).Elements()) { - args.push_back( - FieldInitializer(field.name, ReifyType(field.value, source_loc))); - } - return arena->New(source_loc, args); - } - case Value::Kind::StructType: { - std::vector args; - for (const auto& [name, type] : cast(*t).fields()) { - args.push_back(FieldInitializer(name, ReifyType(type, source_loc))); - } - return arena->New(source_loc, args); - } - case Value::Kind::NominalClassType: - return arena->New( - source_loc, cast(*t).Name()); - case Value::Kind::ChoiceType: - return arena->New(source_loc, - cast(*t).Name()); - case Value::Kind::PointerType: - return arena->New( - source_loc, Operator::Ptr, - std::vector>( - {ReifyType(cast(*t).Type(), source_loc)})); - case Value::Kind::VariableType: - return arena->New(source_loc, - cast(*t).Name()); - case Value::Kind::StringType: - return arena->New(source_loc); - case Value::Kind::AlternativeConstructorValue: - case Value::Kind::AlternativeValue: - case Value::Kind::AutoType: - case Value::Kind::BindingPlaceholderValue: - case Value::Kind::BoolValue: - case Value::Kind::ContinuationValue: - case Value::Kind::FunctionValue: +// Returns whether *value represents a concrete type, as opposed to a +// type pattern or a non-type value. +static auto IsConcreteType(Nonnull value) -> bool { + switch (value->kind()) { case Value::Kind::IntValue: + case Value::Kind::FunctionValue: case Value::Kind::PointerValue: - case Value::Kind::StringValue: + case Value::Kind::BoolValue: case Value::Kind::StructValue: case Value::Kind::NominalClassValue: - FATAL() << "expected a type, not " << *t; + case Value::Kind::AlternativeValue: + case Value::Kind::BindingPlaceholderValue: + case Value::Kind::AlternativeConstructorValue: + case Value::Kind::ContinuationValue: + case Value::Kind::StringValue: + return false; + case Value::Kind::IntType: + case Value::Kind::BoolType: + case Value::Kind::TypeType: + case Value::Kind::FunctionType: + case Value::Kind::PointerType: + case Value::Kind::StructType: + case Value::Kind::NominalClassType: + case Value::Kind::ChoiceType: + case Value::Kind::ContinuationType: + case Value::Kind::VariableType: + case Value::Kind::StringType: + return true; + case Value::Kind::AutoType: + // `auto` isn't a concrete type, it's a pattern that matches types. + return false; + case Value::Kind::TupleValue: + for (const TupleElement& field : cast(*value).Elements()) { + if (!IsConcreteType(field.value)) { + return false; + } + } + return true; + } +} + +void TypeChecker::ExpectIsConcreteType(SourceLocation source_loc, + Nonnull value) { + if (!IsConcreteType(value)) { + FATAL_COMPILATION_ERROR(source_loc) + << "Expected a type, but got " << *value; } } @@ -321,7 +305,7 @@ auto TypeChecker::Substitute(TypeEnv dict, Nonnull type) } auto TypeChecker::TypeCheckExp(Nonnull e, TypeEnv types, - Env values) -> TCExpression { + Env values) -> TCResult { if (tracing_output) { llvm::outs() << "checking expression " << *e << "\ntypes: "; PrintTypeEnv(types, llvm::outs()); @@ -346,10 +330,7 @@ auto TypeChecker::TypeCheckExp(Nonnull e, TypeEnv types, FATAL_COMPILATION_ERROR(e->source_loc()) << "field " << f << " is not in the tuple " << *t; } - auto new_e = arena->New( - e->source_loc(), res.exp, - arena->New(e->source_loc(), i)); - return TCExpression(new_e, *field_t, res.types); + return TCResult(*field_t, res.types); } default: FATAL_COMPILATION_ERROR(e->source_loc()) << "expected a tuple"; @@ -362,12 +343,11 @@ auto TypeChecker::TypeCheckExp(Nonnull e, TypeEnv types, for (auto& arg : cast(*e).fields()) { auto arg_res = TypeCheckExp(arg.expression(), new_types, values); new_types = arg_res.types; - new_args.push_back(FieldInitializer(arg.name(), arg_res.exp)); + new_args.push_back(FieldInitializer(arg.name(), arg.expression())); arg_types.push_back({.name = arg.name(), .value = arg_res.type}); } - auto tuple_e = arena->New(e->source_loc(), new_args); auto tuple_t = arena->New(std::move(arg_types)); - return TCExpression(tuple_e, tuple_t, new_types); + return TCResult(tuple_t, new_types); } case Expression::Kind::StructLiteral: { std::vector new_args; @@ -376,12 +356,11 @@ auto TypeChecker::TypeCheckExp(Nonnull e, TypeEnv types, for (auto& arg : cast(*e).fields()) { auto arg_res = TypeCheckExp(arg.expression(), new_types, values); new_types = arg_res.types; - new_args.push_back(FieldInitializer(arg.name(), arg_res.exp)); + new_args.push_back(FieldInitializer(arg.name(), arg.expression())); arg_types.push_back({arg.name(), arg_res.type}); } - auto new_e = arena->New(e->source_loc(), new_args); auto type = arena->New(std::move(arg_types)); - return TCExpression(new_e, type, new_types); + return TCResult(type, new_types); } case Expression::Kind::StructTypeLiteral: { auto& struct_type = cast(*e); @@ -390,11 +369,10 @@ auto TypeChecker::TypeCheckExp(Nonnull e, TypeEnv types, for (auto& arg : struct_type.fields()) { auto arg_res = TypeCheckExp(arg.expression(), new_types, values); new_types = arg_res.types; - Nonnull type = interpreter.InterpExp(values, arg_res.exp); - new_args.push_back( - FieldInitializer(arg.name(), ReifyType(type, e->source_loc()))); + ExpectIsConcreteType(arg.expression()->source_loc(), + interpreter.InterpExp(values, arg.expression())); + new_args.push_back(FieldInitializer(arg.name(), arg.expression())); } - auto new_e = arena->New(e->source_loc(), new_args); Nonnull type; if (struct_type.fields().empty()) { // `{}` is the type of `{}`, just as `()` is the type of `()`. @@ -405,7 +383,7 @@ auto TypeChecker::TypeCheckExp(Nonnull e, TypeEnv types, } else { type = arena->New(); } - return TCExpression(new_e, type, new_types); + return TCResult(type, new_types); } case Expression::Kind::FieldAccessExpression: { auto& access = cast(*e); @@ -416,9 +394,7 @@ auto TypeChecker::TypeCheckExp(Nonnull e, TypeEnv types, const auto& struct_type = cast(*t); for (const auto& [field_name, field_type] : struct_type.fields()) { if (access.Field() == field_name) { - Nonnull new_e = arena->New( - access.source_loc(), res.exp, access.Field()); - return TCExpression(new_e, field_type, res.types); + return TCResult(field_type, res.types); } } FATAL_COMPILATION_ERROR(access.source_loc()) @@ -430,17 +406,13 @@ auto TypeChecker::TypeCheckExp(Nonnull e, TypeEnv types, // Search for a field for (auto& field : t_class.Fields()) { if (access.Field() == field.first) { - Nonnull new_e = arena->New( - e->source_loc(), res.exp, access.Field()); - return TCExpression(new_e, field.second, res.types); + return TCResult(field.second, res.types); } } // Search for a method for (auto& method : t_class.Methods()) { if (access.Field() == method.first) { - Nonnull new_e = arena->New( - e->source_loc(), res.exp, access.Field()); - return TCExpression(new_e, method.second, res.types); + return TCResult(method.second, res.types); } } FATAL_COMPILATION_ERROR(e->source_loc()) @@ -451,9 +423,7 @@ auto TypeChecker::TypeCheckExp(Nonnull e, TypeEnv types, const auto& tup = cast(*t); for (const TupleElement& field : tup.Elements()) { if (access.Field() == field.name) { - auto new_e = arena->New( - e->source_loc(), res.exp, access.Field()); - return TCExpression(new_e, field.value, res.types); + return TCResult(field.value, res.types); } } FATAL_COMPILATION_ERROR(e->source_loc()) @@ -464,11 +434,9 @@ auto TypeChecker::TypeCheckExp(Nonnull e, TypeEnv types, const auto& choice = cast(*t); for (const auto& vt : choice.Alternatives()) { if (access.Field() == vt.first) { - Nonnull new_e = arena->New( - e->source_loc(), res.exp, access.Field()); auto fun_ty = arena->New( std::vector(), vt.second, t); - return TCExpression(new_e, fun_ty, res.types); + return TCResult(fun_ty, res.types); } } FATAL_COMPILATION_ERROR(e->source_loc()) @@ -485,16 +453,16 @@ auto TypeChecker::TypeCheckExp(Nonnull e, TypeEnv types, const auto& ident = cast(*e); std::optional> type = types.Get(ident.Name()); if (type) { - return TCExpression(e, *type, types); + return TCResult(*type, types); } else { FATAL_COMPILATION_ERROR(e->source_loc()) << "could not find `" << ident.Name() << "`"; } } case Expression::Kind::IntLiteral: - return TCExpression(e, arena->New(), types); + return TCResult(arena->New(), types); case Expression::Kind::BoolLiteral: - return TCExpression(e, arena->New(), types); + return TCResult(arena->New(), types); case Expression::Kind::PrimitiveOperatorExpression: { const auto& op = cast(*e); std::vector> es; @@ -503,54 +471,51 @@ auto TypeChecker::TypeCheckExp(Nonnull e, TypeEnv types, for (Nonnull argument : op.Arguments()) { auto res = TypeCheckExp(argument, types, values); new_types = res.types; - es.push_back(res.exp); + es.push_back(argument); ts.push_back(res.type); } - auto new_e = - arena->New(e->source_loc(), op.Op(), es); switch (op.Op()) { case Operator::Neg: ExpectType(e->source_loc(), "negation", arena->New(), ts[0]); - return TCExpression(new_e, arena->New(), new_types); + return TCResult(arena->New(), new_types); case Operator::Add: ExpectType(e->source_loc(), "addition(1)", arena->New(), ts[0]); ExpectType(e->source_loc(), "addition(2)", arena->New(), ts[1]); - return TCExpression(new_e, arena->New(), new_types); + return TCResult(arena->New(), new_types); case Operator::Sub: ExpectType(e->source_loc(), "subtraction(1)", arena->New(), ts[0]); ExpectType(e->source_loc(), "subtraction(2)", arena->New(), ts[1]); - return TCExpression(new_e, arena->New(), new_types); + return TCResult(arena->New(), new_types); case Operator::Mul: ExpectType(e->source_loc(), "multiplication(1)", arena->New(), ts[0]); ExpectType(e->source_loc(), "multiplication(2)", arena->New(), ts[1]); - return TCExpression(new_e, arena->New(), new_types); + return TCResult(arena->New(), new_types); case Operator::And: ExpectType(e->source_loc(), "&&(1)", arena->New(), ts[0]); ExpectType(e->source_loc(), "&&(2)", arena->New(), ts[1]); - return TCExpression(new_e, arena->New(), new_types); + return TCResult(arena->New(), new_types); case Operator::Or: ExpectType(e->source_loc(), "||(1)", arena->New(), ts[0]); ExpectType(e->source_loc(), "||(2)", arena->New(), ts[1]); - return TCExpression(new_e, arena->New(), new_types); + return TCResult(arena->New(), new_types); case Operator::Not: ExpectType(e->source_loc(), "!", arena->New(), ts[0]); - return TCExpression(new_e, arena->New(), new_types); + return TCResult(arena->New(), new_types); case Operator::Eq: ExpectType(e->source_loc(), "==", ts[0], ts[1]); - return TCExpression(new_e, arena->New(), new_types); + return TCResult(arena->New(), new_types); case Operator::Deref: ExpectPointerType(e->source_loc(), "*", ts[0]); - return TCExpression(new_e, cast(*ts[0]).Type(), - new_types); + return TCResult(cast(*ts[0]).Type(), new_types); case Operator::Ptr: ExpectType(e->source_loc(), "*", arena->New(), ts[0]); - return TCExpression(new_e, arena->New(), new_types); + return TCResult(arena->New(), new_types); } break; } @@ -580,9 +545,7 @@ auto TypeChecker::TypeCheckExp(Nonnull e, TypeEnv types, } else { ExpectType(e->source_loc(), "call", parameter_type, arg_res.type); } - auto new_e = arena->New(e->source_loc(), fun_res.exp, - arg_res.exp); - return TCExpression(new_e, return_type, arg_res.types); + return TCResult(return_type, arg_res.types); } default: { FATAL_COMPILATION_ERROR(e->source_loc()) @@ -593,34 +556,32 @@ auto TypeChecker::TypeCheckExp(Nonnull e, TypeEnv types, break; } case Expression::Kind::FunctionTypeLiteral: { - const auto& fn = cast(*e); - auto pt = interpreter.InterpExp(values, fn.Parameter()); - auto rt = interpreter.InterpExp(values, fn.ReturnType()); - auto new_e = arena->New( - e->source_loc(), ReifyType(pt, e->source_loc()), - ReifyType(rt, e->source_loc()), - /*is_omitted_return_type=*/false); - return TCExpression(new_e, arena->New(), types); + auto& fn = cast(*e); + ExpectIsConcreteType(fn.Parameter()->source_loc(), + interpreter.InterpExp(values, fn.Parameter())); + ExpectIsConcreteType(fn.ReturnType()->source_loc(), + interpreter.InterpExp(values, fn.ReturnType())); + return TCResult(arena->New(), types); } case Expression::Kind::StringLiteral: - return TCExpression(e, arena->New(), types); + return TCResult(arena->New(), types); case Expression::Kind::IntrinsicExpression: switch (cast(*e).Intrinsic()) { case IntrinsicExpression::IntrinsicKind::Print: - return TCExpression(e, TupleValue::Empty(), types); + return TCResult(TupleValue::Empty(), types); } case Expression::Kind::IntTypeLiteral: case Expression::Kind::BoolTypeLiteral: case Expression::Kind::StringTypeLiteral: case Expression::Kind::TypeTypeLiteral: case Expression::Kind::ContinuationTypeLiteral: - return TCExpression(e, arena->New(), types); + return TCResult(arena->New(), types); } } auto TypeChecker::TypeCheckPattern( Nonnull p, TypeEnv types, Env values, - std::optional> expected) -> TCPattern { + std::optional> expected) -> TCResult { if (tracing_output) { llvm::outs() << "checking pattern " << *p; if (expected) { @@ -634,14 +595,13 @@ auto TypeChecker::TypeCheckPattern( } switch (p->kind()) { case Pattern::Kind::AutoPattern: { - return {.pattern = p, .type = arena->New(), .types = types}; + return TCResult(arena->New(), types); } case Pattern::Kind::BindingPattern: { auto& binding = cast(*p); - TCPattern binding_type_result = - TypeCheckPattern(binding.Type(), types, values, std::nullopt); + TypeCheckPattern(binding.Type(), types, values, std::nullopt); Nonnull type = - interpreter.InterpPattern(values, binding_type_result.pattern); + interpreter.InterpPattern(values, binding.Type()); if (expected) { std::optional values = interpreter.PatternMatch( type, *expected, binding.Type()->source_loc()); @@ -654,13 +614,11 @@ auto TypeChecker::TypeCheckPattern( << "Name bindings within type patterns are unsupported"; type = *expected; } - auto new_p = arena->New( - binding.source_loc(), binding.Name(), - arena->New(ReifyType(type, binding.source_loc()))); + ExpectIsConcreteType(binding.source_loc(), type); if (binding.Name().has_value()) { types.Set(*binding.Name(), type); } - return {.pattern = new_p, .type = type, .types = types}; + return TCResult(type, types); } case Pattern::Kind::TuplePattern: { auto& tuple = cast(*p); @@ -691,13 +649,11 @@ auto TypeChecker::TypeCheckPattern( auto field_result = TypeCheckPattern(field.pattern, new_types, values, expected_field_type); new_types = field_result.types; - new_fields.push_back( - TuplePattern::Field(field.name, field_result.pattern)); + new_fields.push_back(TuplePattern::Field(field.name, field.pattern)); field_types.push_back({.name = field.name, .value = field_result.type}); } - auto new_tuple = arena->New(tuple.source_loc(), new_fields); auto tuple_t = arena->New(std::move(field_types)); - return {.pattern = new_tuple, .type = tuple_t, .types = new_types}; + return TCResult(tuple_t, new_types); } case Pattern::Kind::AlternativePattern: { auto& alternative = cast(*p); @@ -719,25 +675,14 @@ auto TypeChecker::TypeCheckPattern( << "'" << alternative.AlternativeName() << "' is not an alternative of " << *choice_type; } - TCPattern arg_results = TypeCheckPattern(alternative.Arguments(), types, - values, *parameter_types); - // TODO: Think about a cleaner way to cast between Ptr types. - // (multiple TODOs) - auto arguments = - Nonnull(cast(arg_results.pattern)); - return {.pattern = arena->New( - alternative.source_loc(), - ReifyType(choice_type, alternative.source_loc()), - alternative.AlternativeName(), arguments), - .type = choice_type, - .types = arg_results.types}; + TCResult arg_results = TypeCheckPattern(alternative.Arguments(), types, + values, *parameter_types); + return TCResult(choice_type, arg_results.types); } case Pattern::Kind::ExpressionPattern: { - TCExpression result = + TCResult result = TypeCheckExp(cast(*p).Expression(), types, values); - return {.pattern = arena->New(result.exp), - .type = result.type, - .types = result.types}; + return TCResult(result.type, result.types); } } } @@ -748,14 +693,14 @@ auto TypeChecker::TypeCheckCase(Nonnull expected, Nonnull return_type_context) -> Match::Clause { auto pat_res = TypeCheckPattern(pat, types, values, expected); - auto res = TypeCheckStmt(body, pat_res.types, values, return_type_context); - return Match::Clause(pat, res.stmt); + TypeCheckStmt(body, pat_res.types, values, return_type_context); + return Match::Clause(pat, body); } auto TypeChecker::TypeCheckStmt(Nonnull s, TypeEnv types, Env values, Nonnull return_type_context) - -> TCStatement { + -> TCResult { switch (s->kind()) { case Statement::Kind::Match: { auto& match = cast(*s); @@ -767,32 +712,26 @@ auto TypeChecker::TypeCheckStmt(Nonnull s, TypeEnv types, &clause.statement(), types, values, return_type_context)); } - auto new_s = arena->New(s->source_loc(), res.exp, new_clauses); - return TCStatement(new_s, types); + return TCResult(TupleValue::Empty(), types); } case Statement::Kind::While: { auto& while_stmt = cast(*s); auto cnd_res = TypeCheckExp(while_stmt.Cond(), types, values); ExpectType(s->source_loc(), "condition of `while`", arena->New(), cnd_res.type); - auto body_res = - TypeCheckStmt(while_stmt.Body(), types, values, return_type_context); - auto new_s = - arena->New(s->source_loc(), cnd_res.exp, body_res.stmt); - return TCStatement(new_s, types); + TypeCheckStmt(while_stmt.Body(), types, values, return_type_context); + return TCResult(TupleValue::Empty(), types); } case Statement::Kind::Break: case Statement::Kind::Continue: - return TCStatement(s, types); + return TCResult(TupleValue::Empty(), types); case Statement::Kind::Block: { auto& block = cast(*s); if (block.Stmt()) { - auto stmt_res = - TypeCheckStmt(*block.Stmt(), types, values, return_type_context); - return TCStatement(arena->New(s->source_loc(), stmt_res.stmt), - types); + TypeCheckStmt(*block.Stmt(), types, values, return_type_context); + return TCResult(TupleValue::Empty(), types); } else { - return TCStatement(s, types); + return TCResult(TupleValue::Empty(), types); } } case Statement::Kind::VariableDefinition: { @@ -800,25 +739,19 @@ auto TypeChecker::TypeCheckStmt(Nonnull s, TypeEnv types, auto res = TypeCheckExp(var.Init(), types, values); Nonnull rhs_ty = res.type; auto lhs_res = TypeCheckPattern(var.Pat(), types, values, rhs_ty); - auto new_s = - arena->New(s->source_loc(), var.Pat(), res.exp); - return TCStatement(new_s, lhs_res.types); + return TCResult(TupleValue::Empty(), lhs_res.types); } case Statement::Kind::Sequence: { auto& seq = cast(*s); auto stmt_res = TypeCheckStmt(seq.Stmt(), types, values, return_type_context); auto checked_types = stmt_res.types; - std::optional> next_stmt; if (seq.Next()) { auto next_res = TypeCheckStmt(*seq.Next(), checked_types, values, return_type_context); - next_stmt = next_res.stmt; checked_types = next_res.types; } - return TCStatement( - arena->New(s->source_loc(), stmt_res.stmt, next_stmt), - checked_types); + return TCResult(TupleValue::Empty(), checked_types); } case Statement::Kind::Assign: { auto& assign = cast(*s); @@ -827,32 +760,22 @@ auto TypeChecker::TypeCheckStmt(Nonnull s, TypeEnv types, auto lhs_res = TypeCheckExp(assign.Lhs(), types, values); auto lhs_t = lhs_res.type; ExpectType(s->source_loc(), "assign", lhs_t, rhs_t); - auto new_s = - arena->New(s->source_loc(), lhs_res.exp, rhs_res.exp); - return TCStatement(new_s, lhs_res.types); + return TCResult(TupleValue::Empty(), lhs_res.types); } case Statement::Kind::ExpressionStatement: { - auto res = - TypeCheckExp(cast(*s).Exp(), types, values); - auto new_s = arena->New(s->source_loc(), res.exp); - return TCStatement(new_s, types); + TypeCheckExp(cast(*s).Exp(), types, values); + return TCResult(TupleValue::Empty(), types); } case Statement::Kind::If: { auto& if_stmt = cast(*s); auto cnd_res = TypeCheckExp(if_stmt.Cond(), types, values); ExpectType(s->source_loc(), "condition of `if`", arena->New(), cnd_res.type); - auto then_res = - TypeCheckStmt(if_stmt.ThenStmt(), types, values, return_type_context); - std::optional> else_stmt; + TypeCheckStmt(if_stmt.ThenStmt(), types, values, return_type_context); if (if_stmt.ElseStmt()) { - auto else_res = TypeCheckStmt(*if_stmt.ElseStmt(), types, values, - return_type_context); - else_stmt = else_res.stmt; + TypeCheckStmt(*if_stmt.ElseStmt(), types, values, return_type_context); } - auto new_s = arena->New(s->source_loc(), cnd_res.exp, then_res.stmt, - else_stmt); - return TCStatement(new_s, types); + return TCResult(TupleValue::Empty(), types); } case Statement::Kind::Return: { auto& ret = cast(*s); @@ -877,45 +800,34 @@ auto TypeChecker::TypeCheckStmt(Nonnull s, TypeEnv types, << (return_type_context->is_omitted() ? " not" : "") << " provide a return value, to match the function's signature."; } - return TCStatement( - arena->New(s->source_loc(), res.exp, ret.IsOmittedExp()), - types); + return TCResult(TupleValue::Empty(), types); } case Statement::Kind::Continuation: { auto& cont = cast(*s); - TCStatement body_result = - TypeCheckStmt(cont.Body(), types, values, return_type_context); - auto new_continuation = arena->New( - s->source_loc(), cont.ContinuationVariable(), body_result.stmt); + TypeCheckStmt(cont.Body(), types, values, return_type_context); types.Set(cont.ContinuationVariable(), arena->New()); - return TCStatement(new_continuation, types); + return TCResult(TupleValue::Empty(), types); } case Statement::Kind::Run: { - TCExpression argument_result = + TCResult argument_result = TypeCheckExp(cast(*s).Argument(), types, values); ExpectType(s->source_loc(), "argument of `run`", arena->New(), argument_result.type); - auto new_run = arena->New(s->source_loc(), argument_result.exp); - return TCStatement(new_run, types); + return TCResult(TupleValue::Empty(), types); } case Statement::Kind::Await: { // nothing to do here - return TCStatement(s, types); + return TCResult(TupleValue::Empty(), types); } } // switch } -auto TypeChecker::CheckOrEnsureReturn( - std::optional> opt_stmt, bool omitted_ret_type, - SourceLocation source_loc) -> Nonnull { +void TypeChecker::ExpectReturnOnAllPaths( + std::optional> opt_stmt, SourceLocation source_loc) { if (!opt_stmt) { - if (omitted_ret_type) { - return arena->New(arena, source_loc); - } else { - FATAL_COMPILATION_ERROR(source_loc) - << "control-flow reaches end of function that provides a `->` return " - "type without reaching a return statement"; - } + FATAL_COMPILATION_ERROR(source_loc) + << "control-flow reaches end of function that provides a `->` return " + "type without reaching a return statement"; } Nonnull stmt = *opt_stmt; switch (stmt->kind()) { @@ -923,59 +835,43 @@ auto TypeChecker::CheckOrEnsureReturn( auto& match = cast(*stmt); std::vector new_clauses; for (auto& clause : match.clauses()) { - auto s = CheckOrEnsureReturn(&clause.statement(), omitted_ret_type, - stmt->source_loc()); - new_clauses.push_back(Match::Clause(&clause.pattern(), s)); + ExpectReturnOnAllPaths(&clause.statement(), stmt->source_loc()); } - return arena->New(stmt->source_loc(), &match.expression(), - new_clauses); + return; } case Statement::Kind::Block: - return arena->New( - stmt->source_loc(), - CheckOrEnsureReturn(cast(*stmt).Stmt(), omitted_ret_type, - stmt->source_loc())); + ExpectReturnOnAllPaths(cast(*stmt).Stmt(), stmt->source_loc()); + return; case Statement::Kind::If: { auto& if_stmt = cast(*stmt); - return arena->New( - stmt->source_loc(), if_stmt.Cond(), - CheckOrEnsureReturn(if_stmt.ThenStmt(), omitted_ret_type, - stmt->source_loc()), - CheckOrEnsureReturn(if_stmt.ElseStmt(), omitted_ret_type, - stmt->source_loc())); + ExpectReturnOnAllPaths(if_stmt.ThenStmt(), stmt->source_loc()); + ExpectReturnOnAllPaths(if_stmt.ElseStmt(), stmt->source_loc()); + return; } case Statement::Kind::Return: - return stmt; + return; case Statement::Kind::Sequence: { auto& seq = cast(*stmt); if (seq.Next()) { - return arena->New( - stmt->source_loc(), seq.Stmt(), - CheckOrEnsureReturn(seq.Next(), omitted_ret_type, - stmt->source_loc())); + ExpectReturnOnAllPaths(seq.Next(), stmt->source_loc()); } else { - return CheckOrEnsureReturn(seq.Stmt(), omitted_ret_type, - stmt->source_loc()); + ExpectReturnOnAllPaths(seq.Stmt(), stmt->source_loc()); } + return; } case Statement::Kind::Continuation: case Statement::Kind::Run: case Statement::Kind::Await: - return stmt; + return; case Statement::Kind::Assign: case Statement::Kind::ExpressionStatement: case Statement::Kind::While: case Statement::Kind::Break: case Statement::Kind::Continue: case Statement::Kind::VariableDefinition: - if (omitted_ret_type) { - return arena->New(stmt->source_loc(), stmt, - arena->New(arena, source_loc)); - } else { - FATAL_COMPILATION_ERROR(stmt->source_loc()) - << "control-flow reaches end of function that provides a `->` " - "return type without reaching a return statement"; - } + FATAL_COMPILATION_ERROR(stmt->source_loc()) + << "control-flow reaches end of function that provides a `->` " + "return type without reaching a return statement"; } } @@ -984,7 +880,7 @@ auto TypeChecker::CheckOrEnsureReturn( // TODO: Add checking to function definitions to ensure that // all deduced type parameters will be deduced. auto TypeChecker::TypeCheckFunDef(FunctionDefinition* f, TypeEnv types, - Env values) -> Nonnull { + Env values) -> TCResult { // Bring the deduced parameters into scope for (const auto& deduced : f->deduced_parameters()) { // auto t = interpreter.InterpExp(values, deduced.type); @@ -1006,20 +902,20 @@ auto TypeChecker::TypeCheckFunDef(FunctionDefinition* f, TypeEnv types, if (f->body()) { 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; + TypeCheckStmt(*f->body(), param_res.types, values, &return_type_context); + body_stmt = *f->body(); // Save the return type in case it changed. if (return_type_context.deduced_return_type().has_value()) { return_type = *return_type_context.deduced_return_type(); } } - auto body = CheckOrEnsureReturn(body_stmt, f->is_omitted_return_type(), - f->source_loc()); - return arena->New( - f->source_loc(), f->name(), f->deduced_parameters(), &f->param_pattern(), - arena->New(ReifyType(return_type, f->source_loc())), - /*is_omitted_return_type=*/false, body); + if (!f->is_omitted_return_type()) { + ExpectReturnOnAllPaths(body_stmt, f->source_loc()); + } + ExpectIsConcreteType(f->return_type().source_loc(), return_type); + return TCResult(arena->New(f->deduced_parameters(), + param_res.type, return_type), + types); } auto TypeChecker::TypeOfFunDef(TypeEnv types, Env values, @@ -1038,8 +934,7 @@ auto TypeChecker::TypeOfFunDef(TypeEnv types, Env values, // Evaluate the return type expression auto ret = interpreter.InterpPattern(values, &fun_def->return_type()); if (ret->kind() == Value::Kind::AutoType) { - auto f = TypeCheckFunDef(fun_def, types, values); - ret = interpreter.InterpPattern(values, &f->return_type()); + return TypeCheckFunDef(fun_def, types, values).type; } return arena->New(fun_def->deduced_parameters(), param_res.type, ret); @@ -1092,39 +987,27 @@ static auto GetName(const Declaration& d) -> const std::string& { } } -auto TypeChecker::MakeTypeChecked(Nonnull d, const TypeEnv& types, - const Env& values) -> Nonnull { +void TypeChecker::TypeCheck(Nonnull d, const TypeEnv& types, + const Env& values) { switch (d->kind()) { case Declaration::Kind::FunctionDeclaration: - return arena->New(TypeCheckFunDef( - &cast(*d).definition(), types, values)); - - case Declaration::Kind::ClassDeclaration: { - const ClassDefinition& class_def = - cast(*d).definition(); - std::vector> fields; - for (Nonnull m : class_def.members()) { - switch (m->kind()) { - case Member::Kind::FieldMember: - // TODO: Interpret the type expression and store the result. - fields.push_back(m); - break; - } - } - return arena->New(class_def.source_loc(), - class_def.name(), std::move(fields)); - } + TypeCheckFunDef(&cast(*d).definition(), types, + values); + return; + case Declaration::Kind::ClassDeclaration: + // TODO + return; case Declaration::Kind::ChoiceDeclaration: // TODO - return d; + return; case Declaration::Kind::VariableDeclaration: { auto& var = cast(*d); // Signals a type error if the initializing expression does not have // the declared type of the variable, otherwise returns this // declaration with annotated types. - TCExpression type_checked_initializer = + TCResult type_checked_initializer = TypeCheckExp(&var.initializer(), types, values); const auto* binding_type = dyn_cast(var.binding().Type()); @@ -1137,7 +1020,7 @@ auto TypeChecker::MakeTypeChecked(Nonnull d, const TypeEnv& types, interpreter.InterpExp(values, binding_type->Expression()); ExpectType(var.source_loc(), "initializer of variable", declared_type, type_checked_initializer.type); - return d; + return; } } } diff --git a/executable_semantics/interpreter/type_checker.h b/executable_semantics/interpreter/type_checker.h index 7eb80a2e4c19..19787b65b695 100644 --- a/executable_semantics/interpreter/type_checker.h +++ b/executable_semantics/interpreter/type_checker.h @@ -32,8 +32,8 @@ class TypeChecker { Env values; }; - auto MakeTypeChecked(Nonnull d, const TypeEnv& types, - const Env& values) -> Nonnull; + void TypeCheck(Nonnull d, const TypeEnv& types, + const Env& values); auto TopLevel(std::vector>* fs) -> TypeCheckContext; @@ -69,28 +69,13 @@ class TypeChecker { const bool is_omitted_; }; - struct TCExpression { - TCExpression(Nonnull e, Nonnull t, TypeEnv types) - : exp(e), type(t), types(types) {} + struct TCResult { + TCResult(Nonnull t, TypeEnv types) : type(t), types(types) {} - Nonnull exp; Nonnull type; TypeEnv types; }; - struct TCPattern { - Nonnull pattern; - Nonnull type; - TypeEnv types; - }; - - struct TCStatement { - TCStatement(Nonnull s, TypeEnv types) : stmt(s), types(types) {} - - Nonnull stmt; - TypeEnv types; - }; - // TypeCheckExp performs semantic analysis on an expression. It returns a new // version of the expression, its type, and an updated environment which are // bundled into a TCResult object. The purpose of the updated environment is @@ -103,7 +88,7 @@ class TypeChecker { // values maps variable names to their compile-time values. It is not // directly used in this function but is passed to InterExp. auto TypeCheckExp(Nonnull e, TypeEnv types, Env values) - -> TCExpression; + -> TCResult; // Equivalent to TypeCheckExp, but operates on Patterns instead of // Expressions. `expected` is the type that this pattern is expected to have, @@ -111,7 +96,7 @@ class TypeChecker { // nullopt. auto TypeCheckPattern(Nonnull p, TypeEnv types, Env values, std::optional> expected) - -> TCPattern; + -> TCResult; // TypeCheckStmt performs semantic analysis on a statement. It returns a new // version of the statement and a new type environment. @@ -122,10 +107,10 @@ class TypeChecker { // statement. auto TypeCheckStmt(Nonnull s, TypeEnv types, Env values, Nonnull return_type_context) - -> TCStatement; + -> TCResult; auto TypeCheckFunDef(FunctionDefinition* f, TypeEnv types, Env values) - -> Nonnull; + -> TCResult; auto TypeCheckCase(Nonnull expected, Nonnull pat, Nonnull body, TypeEnv types, Env values, @@ -139,13 +124,15 @@ class TypeChecker { void TopLevel(Nonnull d, TypeCheckContext* tops); - auto CheckOrEnsureReturn(std::optional> opt_stmt, - bool omitted_ret_type, SourceLocation source_loc) - -> Nonnull; + // Verifies that opt_stmt holds a statement, and it is structurally impossible + // for control flow to leave that statement except via a `return`. + void ExpectReturnOnAllPaths(std::optional> opt_stmt, + SourceLocation source_loc); - // Reify type to type expression. - auto ReifyType(Nonnull t, SourceLocation source_loc) - -> Nonnull; + // Verifies that *value represents a concrete type, as opposed to a + // type pattern or a non-type value. + void ExpectIsConcreteType(SourceLocation source_loc, + Nonnull value); auto Substitute(TypeEnv dict, Nonnull type) -> Nonnull;