diff --git a/executable_semantics/ast/expression.h b/executable_semantics/ast/expression.h index c5ae774c9521..1391de742208 100644 --- a/executable_semantics/ast/expression.h +++ b/executable_semantics/ast/expression.h @@ -54,7 +54,7 @@ class Expression { auto source_loc() const -> SourceLocation { return source_loc_; } // The static type of this expression. Cannot be called before typechecking. - auto static_type() const -> Nonnull { return *static_type_; } + auto static_type() const -> const Value& { return **static_type_; } // Sets the static type of this expression. Can only be called once, during // typechecking. diff --git a/executable_semantics/ast/function_definition.h b/executable_semantics/ast/function_definition.h index 2a6c4fad5ee0..f5337a568ced 100644 --- a/executable_semantics/ast/function_definition.h +++ b/executable_semantics/ast/function_definition.h @@ -61,7 +61,7 @@ class FunctionDefinition { auto body() -> std::optional> { return body_; } // The static type of this function. Cannot be called before typechecking. - auto static_type() const -> Nonnull { return *static_type_; } + auto static_type() const -> const Value& { return **static_type_; } // Sets the static type of this expression. Can only be called once, during // typechecking. diff --git a/executable_semantics/ast/pattern.h b/executable_semantics/ast/pattern.h index a6ae1cdc423a..926543b80bd6 100644 --- a/executable_semantics/ast/pattern.h +++ b/executable_semantics/ast/pattern.h @@ -49,7 +49,7 @@ class Pattern { auto source_loc() const -> SourceLocation { return source_loc_; } // The static type of this pattern. Cannot be called before typechecking. - auto static_type() const -> Nonnull { return *static_type_; } + auto static_type() const -> const Value& { return **static_type_; } // Sets the static type of this expression. Can only be called once, during // typechecking. diff --git a/executable_semantics/interpreter/type_checker.cpp b/executable_semantics/interpreter/type_checker.cpp index 5a900e3087b5..b391c8f3e014 100644 --- a/executable_semantics/interpreter/type_checker.cpp +++ b/executable_semantics/interpreter/type_checker.cpp @@ -31,7 +31,7 @@ namespace Carbon { static void SetStaticType(Nonnull expression, Nonnull type) { if (expression->has_static_type()) { - CHECK(TypeEqual(expression->static_type(), type)); + CHECK(TypeEqual(&expression->static_type(), type)); } else { expression->set_static_type(type); } @@ -42,7 +42,7 @@ static void SetStaticType(Nonnull expression, static void SetStaticType(Nonnull pattern, Nonnull type) { if (pattern->has_static_type()) { - CHECK(TypeEqual(pattern->static_type(), type)); + CHECK(TypeEqual(&pattern->static_type(), type)); } else { pattern->set_static_type(type); } @@ -53,7 +53,7 @@ static void SetStaticType(Nonnull pattern, static void SetStaticType(Nonnull definition, Nonnull type) { if (definition->has_static_type()) { - CHECK(TypeEqual(definition->static_type(), type)); + CHECK(TypeEqual(&definition->static_type(), type)); } else { definition->set_static_type(type); } @@ -436,18 +436,18 @@ auto TypeChecker::TypeCheckExp(Nonnull e, TypeEnv types, case Expression::Kind::IndexExpression: { auto& index = cast(*e); auto res = TypeCheckExp(&index.aggregate(), types, values); - Nonnull aggregate_type = index.aggregate().static_type(); - switch (aggregate_type->kind()) { + const Value& aggregate_type = index.aggregate().static_type(); + switch (aggregate_type.kind()) { case Value::Kind::TupleValue: { auto i = cast(*interpreter.InterpExp(values, &index.offset())) .Val(); std::string f = std::to_string(i); std::optional> field_t = - cast(*aggregate_type).FindField(f); + cast(aggregate_type).FindField(f); if (!field_t) { FATAL_COMPILATION_ERROR(e->source_loc()) - << "field " << f << " is not in the tuple " << *aggregate_type; + << "field " << f << " is not in the tuple " << aggregate_type; } SetStaticType(&index, *field_t); return TCResult(res.types); @@ -465,7 +465,7 @@ auto TypeChecker::TypeCheckExp(Nonnull e, TypeEnv types, new_types = arg_res.types; new_args.push_back(FieldInitializer(arg.name(), &arg.expression())); arg_types.push_back( - {.name = arg.name(), .value = arg.expression().static_type()}); + {.name = arg.name(), .value = &arg.expression().static_type()}); } SetStaticType(e, arena->New(std::move(arg_types))); return TCResult(new_types); @@ -478,7 +478,7 @@ auto TypeChecker::TypeCheckExp(Nonnull e, TypeEnv types, auto arg_res = TypeCheckExp(&arg.expression(), new_types, values); new_types = arg_res.types; new_args.push_back(FieldInitializer(arg.name(), &arg.expression())); - arg_types.push_back({arg.name(), arg.expression().static_type()}); + arg_types.push_back({arg.name(), &arg.expression().static_type()}); } SetStaticType(e, arena->New(std::move(arg_types))); return TCResult(new_types); @@ -508,10 +508,10 @@ auto TypeChecker::TypeCheckExp(Nonnull e, TypeEnv types, case Expression::Kind::FieldAccessExpression: { auto& access = cast(*e); auto res = TypeCheckExp(&access.aggregate(), types, values); - Nonnull aggregate_type = access.aggregate().static_type(); - switch (aggregate_type->kind()) { + const Value& aggregate_type = access.aggregate().static_type(); + switch (aggregate_type.kind()) { case Value::Kind::StructType: { - const auto& struct_type = cast(*aggregate_type); + const auto& struct_type = cast(aggregate_type); for (const auto& [field_name, field_type] : struct_type.fields()) { if (access.field() == field_name) { SetStaticType(&access, field_type); @@ -523,7 +523,7 @@ auto TypeChecker::TypeCheckExp(Nonnull e, TypeEnv types, << access.field(); } case Value::Kind::NominalClassType: { - const auto& t_class = cast(*aggregate_type); + const auto& t_class = cast(aggregate_type); // Search for a field for (auto& field : t_class.Fields()) { if (access.field() == field.first) { @@ -543,7 +543,7 @@ auto TypeChecker::TypeCheckExp(Nonnull e, TypeEnv types, << access.field(); } case Value::Kind::TupleValue: { - const auto& tup = cast(*aggregate_type); + const auto& tup = cast(aggregate_type); for (const TupleElement& field : tup.Elements()) { if (access.field() == field.name) { SetStaticType(&access, field.value); @@ -555,12 +555,12 @@ auto TypeChecker::TypeCheckExp(Nonnull e, TypeEnv types, << access.field(); } case Value::Kind::ChoiceType: { - const auto& choice = cast(*aggregate_type); + const auto& choice = cast(aggregate_type); for (const auto& vt : choice.Alternatives()) { if (access.field() == vt.first) { SetStaticType(&access, arena->New( std::vector(), - vt.second, aggregate_type)); + vt.second, &aggregate_type)); return TCResult(res.types); } } @@ -600,7 +600,7 @@ auto TypeChecker::TypeCheckExp(Nonnull e, TypeEnv types, auto res = TypeCheckExp(argument, types, values); new_types = res.types; es.push_back(argument); - ts.push_back(argument->static_type()); + ts.push_back(&argument->static_type()); } switch (op.op()) { case Operator::Neg: @@ -665,17 +665,16 @@ auto TypeChecker::TypeCheckExp(Nonnull e, TypeEnv types, case Expression::Kind::CallExpression: { auto& call = cast(*e); auto fun_res = TypeCheckExp(&call.function(), types, values); - switch (call.function().static_type()->kind()) { + switch (call.function().static_type().kind()) { case Value::Kind::FunctionType: { - const auto& fun_t = - cast(*call.function().static_type()); + const auto& fun_t = cast(call.function().static_type()); auto arg_res = TypeCheckExp(&call.argument(), fun_res.types, values); auto parameter_type = fun_t.Param(); auto return_type = fun_t.Ret(); if (!fun_t.Deduced().empty()) { auto deduced_args = ArgumentDeduction( e->source_loc(), TypeEnv(arena), parameter_type, - call.argument().static_type()); + &call.argument().static_type()); for (auto& deduced_param : fun_t.Deduced()) { // TODO: change the following to a CHECK once the real checking // has been added to the type checking of function signatures. @@ -689,7 +688,7 @@ auto TypeChecker::TypeCheckExp(Nonnull e, TypeEnv types, return_type = Substitute(deduced_args, return_type); } else { ExpectType(e->source_loc(), "call", parameter_type, - call.argument().static_type()); + &call.argument().static_type()); } SetStaticType(&call, return_type); return TCResult(arg_res.types); @@ -808,7 +807,7 @@ auto TypeChecker::TypeCheckPattern( new_types = field_result.types; new_fields.push_back(TuplePattern::Field(field.name, field.pattern)); field_types.push_back( - {.name = field.name, .value = field.pattern->static_type()}); + {.name = field.name, .value = &field.pattern->static_type()}); } SetStaticType(&tuple, arena->New(std::move(field_types))); return TCResult(new_types); @@ -841,7 +840,7 @@ auto TypeChecker::TypeCheckPattern( case Pattern::Kind::ExpressionPattern: { const auto& expression = cast(*p).Expression(); TCResult result = TypeCheckExp(expression, types, values); - SetStaticType(p, expression->static_type()); + SetStaticType(p, &expression->static_type()); return TCResult(result.types); } } @@ -868,7 +867,7 @@ auto TypeChecker::TypeCheckStmt(Nonnull s, TypeEnv types, std::vector new_clauses; for (auto& clause : match.clauses()) { new_clauses.push_back(TypeCheckCase( - match.expression().static_type(), &clause.pattern(), + &match.expression().static_type(), &clause.pattern(), &clause.statement(), types, values, return_type_context)); } return TCResult(types); @@ -877,7 +876,7 @@ auto TypeChecker::TypeCheckStmt(Nonnull s, TypeEnv types, auto& while_stmt = cast(*s); TypeCheckExp(while_stmt.Cond(), types, values); ExpectType(s->source_loc(), "condition of `while`", - arena->New(), while_stmt.Cond()->static_type()); + arena->New(), &while_stmt.Cond()->static_type()); TypeCheckStmt(while_stmt.Body(), types, values, return_type_context); return TCResult(types); } @@ -896,8 +895,8 @@ auto TypeChecker::TypeCheckStmt(Nonnull s, TypeEnv types, case Statement::Kind::VariableDefinition: { auto& var = cast(*s); TypeCheckExp(var.Init(), types, values); - Nonnull rhs_ty = var.Init()->static_type(); - auto lhs_res = TypeCheckPattern(var.Pat(), types, values, rhs_ty); + const Value& rhs_ty = var.Init()->static_type(); + auto lhs_res = TypeCheckPattern(var.Pat(), types, values, &rhs_ty); return TCResult(lhs_res.types); } case Statement::Kind::Sequence: { @@ -916,8 +915,8 @@ auto TypeChecker::TypeCheckStmt(Nonnull s, TypeEnv types, auto& assign = cast(*s); TypeCheckExp(assign.Rhs(), types, values); auto lhs_res = TypeCheckExp(assign.Lhs(), types, values); - ExpectType(s->source_loc(), "assign", assign.Lhs()->static_type(), - assign.Rhs()->static_type()); + ExpectType(s->source_loc(), "assign", &assign.Lhs()->static_type(), + &assign.Rhs()->static_type()); return TCResult(lhs_res.types); } case Statement::Kind::ExpressionStatement: { @@ -928,7 +927,7 @@ auto TypeChecker::TypeCheckStmt(Nonnull s, TypeEnv types, auto& if_stmt = cast(*s); TypeCheckExp(if_stmt.Cond(), types, values); ExpectType(s->source_loc(), "condition of `if`", arena->New(), - if_stmt.Cond()->static_type()); + &if_stmt.Cond()->static_type()); TypeCheckStmt(if_stmt.ThenStmt(), types, values, return_type_context); if (if_stmt.ElseStmt()) { TypeCheckStmt(*if_stmt.ElseStmt(), types, values, return_type_context); @@ -947,12 +946,12 @@ auto TypeChecker::TypeCheckStmt(Nonnull s, TypeEnv types, } else { // Infer the auto return from the first `return` statement. return_type_context->set_deduced_return_type( - ret.Exp()->static_type()); + &ret.Exp()->static_type()); } } else { ExpectType(s->source_loc(), "return", *return_type_context->deduced_return_type(), - ret.Exp()->static_type()); + &ret.Exp()->static_type()); } if (ret.IsOmittedExp() != return_type_context->is_omitted()) { FATAL_COMPILATION_ERROR(s->source_loc()) @@ -972,7 +971,8 @@ auto TypeChecker::TypeCheckStmt(Nonnull s, TypeEnv types, auto& run = cast(*s); TypeCheckExp(run.Argument(), types, values); ExpectType(s->source_loc(), "argument of `run`", - arena->New(), run.Argument()->static_type()); + arena->New(), + &run.Argument()->static_type()); return TCResult(types); } case Statement::Kind::Await: { @@ -1094,7 +1094,7 @@ auto TypeChecker::TypeCheckFunDef(FunctionDefinition* f, TypeEnv types, } ExpectIsConcreteType(f->return_type().source_loc(), return_type); SetStaticType(f, arena->New(f->deduced_parameters(), - f->param_pattern().static_type(), + &f->param_pattern().static_type(), return_type)); return TCResult(types); } @@ -1116,10 +1116,10 @@ auto TypeChecker::TypeOfFunDef(TypeEnv types, Env values, if (ret->kind() == Value::Kind::AutoType) { // FIXME do this unconditionally? TypeCheckFunDef(fun_def, types, values); - return fun_def->static_type(); + return &fun_def->static_type(); } return arena->New(fun_def->deduced_parameters(), - fun_def->param_pattern().static_type(), ret); + &fun_def->param_pattern().static_type(), ret); } auto TypeChecker::TypeOfClassDef(const ClassDefinition* sd, TypeEnv /*types*/, @@ -1200,7 +1200,7 @@ void TypeChecker::TypeCheck(Nonnull d, const TypeEnv& types, Nonnull declared_type = interpreter.InterpExp(values, binding_type->Expression()); ExpectType(var.source_loc(), "initializer of variable", declared_type, - var.initializer().static_type()); + &var.initializer().static_type()); return; } }