Factor out a Pattern sum type from Expression (#685)

`Pattern` is intended to pilot some changes I would like to apply to all our sum types:
- The alternatives are expressed as derived classes rather than members of a `std::variant`.
- The alternatives are classes in the [style guide sense](https://google.github.io/styleguide/cppguide.html#Structs_vs._Classes), meaning they can have invariants, but can't have public data members.
- Creating an object is expressed using a constructor rather than a factory function.
- Accessing an alternative is expressed as a cast (using LLVM's RTTI system) rather than `std::get` or a `Get` method.

Co-authored-by: Jon Meow <46229924+jonmeow@users.noreply.github.com>
This commit is contained in:
Geoff Romer
2021-07-30 12:24:12 -07:00
committed by GitHub
co-authored by Jon Meow
parent b08f6bb0f1
commit 6ac3adfa53
31 changed files with 1302 additions and 575 deletions
+2
View File
@@ -98,6 +98,7 @@ cc_library(
"//executable_semantics/ast:expression",
"//executable_semantics/ast:function_definition",
"//executable_semantics/common:tracing_flag",
"@llvm-project//llvm:Support",
],
)
@@ -113,6 +114,7 @@ cc_library(
"//executable_semantics/ast:function_definition",
"//executable_semantics/ast:statement",
"//executable_semantics/common:tracing_flag",
"@llvm-project//llvm:Support",
],
)
@@ -29,6 +29,12 @@ auto Action::MakeExpressionAction(const Expression* e) -> Action* {
return act;
}
auto Action::MakePatternAction(const Pattern* p) -> Action* {
auto* act = new Action();
act->value = PatternAction({.pattern = p});
return act;
}
auto Action::MakeStatementAction(const Statement* s) -> Action* {
auto* act = new Action();
act->value = StatementAction({.stmt = s});
@@ -49,6 +55,10 @@ auto Action::GetExpressionAction() const -> const ExpressionAction& {
return std::get<ExpressionAction>(value);
}
auto Action::GetPatternAction() const -> const PatternAction& {
return std::get<PatternAction>(value);
}
auto Action::GetStatementAction() const -> const StatementAction& {
return std::get<StatementAction>(value);
}
@@ -65,6 +75,9 @@ void Action::Print(llvm::raw_ostream& out) const {
case ActionKind::ExpressionAction:
out << *GetExpressionAction().exp;
break;
case ActionKind::PatternAction:
out << *GetPatternAction().pattern;
break;
case ActionKind::StatementAction:
GetStatementAction().stmt->PrintDepth(1, out);
break;
+12 -1
View File
@@ -9,6 +9,7 @@
#include "common/ostream.h"
#include "executable_semantics/ast/expression.h"
#include "executable_semantics/ast/pattern.h"
#include "executable_semantics/ast/statement.h"
#include "executable_semantics/interpreter/stack.h"
#include "executable_semantics/interpreter/value.h"
@@ -19,6 +20,7 @@ namespace Carbon {
enum class ActionKind {
LValAction,
ExpressionAction,
PatternAction,
StatementAction,
ValAction,
};
@@ -33,6 +35,11 @@ struct ExpressionAction {
const Expression* exp;
};
struct PatternAction {
static constexpr ActionKind Kind = ActionKind::PatternAction;
const Pattern* pattern;
};
struct StatementAction {
static constexpr ActionKind Kind = ActionKind::StatementAction;
const Statement* stmt;
@@ -46,6 +53,7 @@ struct ValAction {
struct Action {
static auto MakeLValAction(const Expression* e) -> Action*;
static auto MakeExpressionAction(const Expression* e) -> Action*;
static auto MakePatternAction(const Pattern* p) -> Action*;
static auto MakeStatementAction(const Statement* s) -> Action*;
static auto MakeValAction(const Value* v) -> Action*;
@@ -53,6 +61,7 @@ struct Action {
auto GetLValAction() const -> const LValAction&;
auto GetExpressionAction() const -> const ExpressionAction&;
auto GetPatternAction() const -> const PatternAction&;
auto GetStatementAction() const -> const StatementAction&;
auto GetValAction() const -> const ValAction&;
@@ -75,7 +84,9 @@ struct Action {
std::vector<const Value*> results;
private:
std::variant<LValAction, ExpressionAction, StatementAction, ValAction> value;
std::variant<LValAction, ExpressionAction, PatternAction, StatementAction,
ValAction>
value;
};
} // namespace Carbon
+116 -30
View File
@@ -20,6 +20,9 @@
#include "executable_semantics/interpreter/frame.h"
#include "executable_semantics/interpreter/stack.h"
#include "llvm/ADT/StringExtras.h"
#include "llvm/Support/Casting.h"
using llvm::cast;
namespace Carbon {
@@ -114,7 +117,7 @@ void InitEnv(const Declaration& d, Env* env) {
state->heap.AllocateValue(Value::MakeVariableType(deduced.name));
new_env.Set(deduced.name, a);
}
auto pt = InterpExp(new_env, func_def.param_pattern);
auto pt = InterpPattern(new_env, func_def.param_pattern);
auto f = Value::MakeFunctionValue(func_def.name, pt, func_def.body);
Address a = state->heap.AllocateValue(f);
env->Set(func_def.name, a);
@@ -128,9 +131,11 @@ void InitEnv(const Declaration& d, Env* env) {
for (const Member* m : struct_def.members) {
switch (m->tag()) {
case MemberKind::FieldMember: {
const auto& field = m->GetFieldMember();
auto t = InterpExp(Env(), field.type);
fields.push_back(make_pair(field.name, t));
const BindingPattern* binding = m->GetFieldMember().binding;
const Expression* type_expression =
cast<ExpressionPattern>(binding->Type())->Expression();
auto type = InterpExp(Env(), type_expression);
fields.push_back(make_pair(*binding->Name(), type));
break;
}
}
@@ -161,7 +166,7 @@ void InitEnv(const Declaration& d, Env* env) {
// result of evaluating the initializer.
auto v = InterpExp(*env, var.initializer);
Address a = state->heap.AllocateValue(v);
env->Set(var.name, a);
env->Set(*var.binding->Name(), a);
break;
}
}
@@ -502,9 +507,7 @@ void StepLvalue() {
case ExpressionKind::BoolTypeLiteral:
case ExpressionKind::TypeTypeLiteral:
case ExpressionKind::FunctionTypeLiteral:
case ExpressionKind::AutoTypeLiteral:
case ExpressionKind::ContinuationTypeLiteral:
case ExpressionKind::BindingExpression: {
case ExpressionKind::ContinuationTypeLiteral: {
FATAL_RUNTIME_ERROR_NO_LINE()
<< "Can't treat expression as lvalue: " << *exp;
}
@@ -521,19 +524,6 @@ void StepExp() {
llvm::outs() << "--- step exp " << *exp << " --->\n";
}
switch (exp->tag()) {
case ExpressionKind::BindingExpression: {
if (act->pos == 0) {
frame->todo.Push(
Action::MakeExpressionAction(exp->GetBindingExpression().type));
act->pos++;
} else {
auto v = Value::MakeBindingPlaceholderValue(
exp->GetBindingExpression().name, act->results[0]);
frame->todo.Pop(1);
frame->todo.Push(Action::MakeValAction(v));
}
break;
}
case ExpressionKind::IndexExpression: {
if (act->pos == 0) {
// { { e[i] :: C, E, F} :: S, H}
@@ -696,13 +686,6 @@ void StepExp() {
frame->todo.Push(Action::MakeValAction(v));
break;
}
case ExpressionKind::AutoTypeLiteral: {
CHECK(act->pos == 0);
const Value* v = Value::MakeAutoType();
frame->todo.Pop(1);
frame->todo.Push(Action::MakeValAction(v));
break;
}
case ExpressionKind::TypeTypeLiteral: {
CHECK(act->pos == 0);
const Value* v = Value::MakeTypeType();
@@ -741,6 +724,91 @@ void StepExp() {
} // switch (exp->tag)
}
void StepPattern() {
Frame* frame = state->stack.Top();
Action* act = frame->todo.Top();
const Pattern* pattern = act->GetPatternAction().pattern;
if (tracing_output) {
llvm::outs() << "--- step pattern " << *pattern << " --->\n";
}
switch (pattern->Tag()) {
case Pattern::Kind::AutoPattern: {
CHECK(act->pos == 0);
const Value* v = Value::MakeAutoType();
frame->todo.Pop(1);
frame->todo.Push(Action::MakeValAction(v));
break;
}
case Pattern::Kind::BindingPattern: {
const auto& binding = cast<BindingPattern>(*pattern);
if (act->pos == 0) {
frame->todo.Push(Action::MakePatternAction(binding.Type()));
act->pos++;
} else {
auto v =
Value::MakeBindingPlaceholderValue(binding.Name(), act->results[0]);
frame->todo.Pop(1);
frame->todo.Push(Action::MakeValAction(v));
}
break;
}
case Pattern::Kind::TuplePattern: {
const auto& tuple = cast<TuplePattern>(*pattern);
if (act->pos == 0) {
if (tuple.Fields().empty()) {
frame->todo.Pop(1);
frame->todo.Push(Action::MakeValAction(Value::MakeTupleValue({})));
} else {
const Pattern* p1 = tuple.Fields()[0].pattern;
frame->todo.Push(Action::MakePatternAction(p1));
act->pos++;
}
} else if (act->pos != static_cast<int>(tuple.Fields().size())) {
// { { vk :: (f1=v1,..., fk=[],fk+1=ek+1,...) :: C, E, F} :: S,
// H}
// -> { { ek+1 :: (f1=v1,..., fk=vk, fk+1=[],...) :: C, E, F} :: S,
// H}
const Pattern* elt = tuple.Fields()[act->pos].pattern;
frame->todo.Push(Action::MakePatternAction(elt));
act->pos++;
} else {
std::vector<TupleElement> elements;
for (size_t i = 0; i < tuple.Fields().size(); ++i) {
elements.push_back(
{.name = tuple.Fields()[i].name, .value = act->results[i]});
}
const Value* tuple_value = Value::MakeTupleValue(std::move(elements));
frame->todo.Pop(1);
frame->todo.Push(Action::MakeValAction(tuple_value));
}
break;
}
case Pattern::Kind::AlternativePattern: {
const auto& alternative = cast<AlternativePattern>(*pattern);
if (act->pos == 0) {
frame->todo.Push(
Action::MakeExpressionAction(alternative.ChoiceType()));
act->pos++;
} else if (act->pos == 1) {
frame->todo.Push(Action::MakePatternAction(alternative.Arguments()));
act->pos++;
} else {
CHECK(act->pos == 2);
const auto& choice_type = act->results[0]->GetChoiceType();
frame->todo.Pop(1);
frame->todo.Push(Action::MakeValAction(Value::MakeAlternativeValue(
alternative.AlternativeName(), choice_type.name, act->results[1])));
}
break;
}
case Pattern::Kind::ExpressionPattern:
frame->todo.Pop(1);
frame->todo.Push(Action::MakeExpressionAction(
cast<ExpressionPattern>(pattern)->Expression()));
break;
}
}
auto IsWhileAct(Action* act) -> bool {
switch (act->tag()) {
case ActionKind::StatementAction:
@@ -810,7 +878,7 @@ void StepStmt() {
// start interpreting the pattern of the clause
// { {v :: (match ([]) ...) :: C, E, F} :: S, H}
// -> { {pi :: (match ([]) ...) :: C, E, F} :: S, H}
frame->todo.Push(Action::MakeExpressionAction(c->first));
frame->todo.Push(Action::MakePatternAction(c->first));
act->pos++;
} else { // try to match
auto v = act->results[0];
@@ -916,7 +984,7 @@ void StepStmt() {
act->pos++;
} else if (act->pos == 1) {
frame->todo.Push(
Action::MakeExpressionAction(stmt->GetVariableDefinition().pat));
Action::MakePatternAction(stmt->GetVariableDefinition().pat));
act->pos++;
} else if (act->pos == 2) {
// { { v :: (x = []) :: C, E, F} :: S, H}
@@ -1097,6 +1165,9 @@ void Step() {
case ActionKind::ExpressionAction:
StepExp();
break;
case ActionKind::PatternAction:
StepPattern();
break;
case ActionKind::StatementAction:
StepStmt();
break;
@@ -1150,4 +1221,19 @@ auto InterpExp(Env values, const Expression* e) -> const Value* {
return v;
}
// Interpret a pattern at compile-time.
auto InterpPattern(Env values, const Pattern* p) -> const Value* {
auto todo = Stack(Action::MakePatternAction(p));
auto* scope = new Scope(values, std::list<std::string>());
auto* frame = new Frame("InterpPattern", Stack(scope), todo);
state->stack = Stack(frame);
while (state->stack.Count() > 1 || state->stack.Top()->todo.Count() > 1 ||
state->stack.Top()->todo.Top()->tag() != ActionKind::ValAction) {
Step();
}
const Value* v = state->stack.Top()->todo.Top()->GetValAction().val;
return v;
}
} // namespace Carbon
@@ -11,6 +11,8 @@
#include "common/ostream.h"
#include "executable_semantics/ast/declaration.h"
#include "executable_semantics/ast/expression.h"
#include "executable_semantics/ast/pattern.h"
#include "executable_semantics/interpreter/frame.h"
#include "executable_semantics/interpreter/heap.h"
#include "executable_semantics/interpreter/stack.h"
@@ -35,6 +37,7 @@ void PrintEnv(Env values);
auto InterpProgram(std::list<Declaration>* fs) -> int;
auto InterpExp(Env values, const Expression* e) -> const Value*;
auto InterpPattern(Env values, const Pattern* p) -> const Value*;
} // namespace Carbon
+242 -174
View File
@@ -10,10 +10,16 @@
#include <set>
#include <vector>
#include "common/ostream.h"
#include "executable_semantics/ast/function_definition.h"
#include "executable_semantics/common/error.h"
#include "executable_semantics/common/tracing_flag.h"
#include "executable_semantics/interpreter/interpreter.h"
#include "executable_semantics/interpreter/value.h"
#include "llvm/Support/Casting.h"
using llvm::cast;
using llvm::dyn_cast;
namespace Carbon {
@@ -53,8 +59,8 @@ auto ReifyType(const Value* t, int line_num) -> const Expression* {
case ValKind::TupleValue: {
std::vector<FieldInitializer> args;
for (const TupleElement& field : t->GetTupleValue().elements) {
args.push_back({.name = field.name,
.expression = ReifyType(field.value, line_num)});
args.push_back(
FieldInitializer(field.name, ReifyType(field.value, line_num)));
}
return Expression::MakeTupleLiteral(0, args);
}
@@ -230,58 +236,14 @@ auto Substitute(TypeEnv dict, const Value* type) -> const Value* {
// types maps variable names to the type of their run-time value.
// values maps variable names to their compile-time values. It is not
// directly used in this function but is passed to InterExp.
// expected is the type that this expression is expected to have.
// This parameter is non-null when the expression is in a pattern context
// and it is used to implement `auto`, otherwise it is null.
// context says what kind of position this expression is nested in,
// whether it's a position that expects a value, a pattern, or a type.
auto TypeCheckExp(const Expression* e, TypeEnv types, Env values,
const Value* expected, TCContext context) -> TCResult {
auto TypeCheckExp(const Expression* e, TypeEnv types, Env values)
-> TCExpression {
if (tracing_output) {
switch (context) {
case TCContext::ValueContext:
llvm::outs() << "checking expression ";
break;
case TCContext::PatternContext:
llvm::outs() << "checking pattern, ";
if (expected) {
llvm::outs() << "expecting " << *expected;
}
llvm::outs() << ", ";
break;
case TCContext::TypeContext:
llvm::outs() << "checking type ";
break;
}
llvm::outs() << *e << "\n";
llvm::outs() << "checking expression " << *e << "\n";
}
switch (e->tag()) {
case ExpressionKind::BindingExpression: {
if (context != TCContext::PatternContext) {
FATAL_COMPILATION_ERROR(e->line_num)
<< "pattern variables are only allowed in pattern context";
}
auto t = InterpExp(values, e->GetBindingExpression().type);
if (t->tag() == ValKind::AutoType) {
if (expected == nullptr) {
FATAL_COMPILATION_ERROR(e->line_num) << " auto not allowed here";
} else {
t = expected;
}
} else if (expected) {
ExpectType(e->line_num, "pattern variable", t, expected);
}
const std::optional<std::string>& name = e->GetBindingExpression().name;
auto new_e = Expression::MakeBindingExpression(e->line_num, name,
ReifyType(t, e->line_num));
if (name.has_value()) {
types.Set(*name, t);
}
return TCResult(new_e, t, types);
}
case ExpressionKind::IndexExpression: {
auto res = TypeCheckExp(e->GetIndexExpression().aggregate, types, values,
nullptr, TCContext::ValueContext);
auto res = TypeCheckExp(e->GetIndexExpression().aggregate, types, values);
auto t = res.type;
switch (t->tag()) {
case ValKind::TupleValue: {
@@ -295,7 +257,7 @@ auto TypeCheckExp(const Expression* e, TypeEnv types, Env values,
}
auto new_e = Expression::MakeIndexExpression(
e->line_num, res.exp, Expression::MakeIntLiteral(e->line_num, i));
return TCResult(new_e, field_t, res.types);
return TCExpression(new_e, field_t, res.types);
}
default:
FATAL_COMPILATION_ERROR(e->line_num) << "expected a tuple";
@@ -305,39 +267,21 @@ auto TypeCheckExp(const Expression* e, TypeEnv types, Env values,
std::vector<FieldInitializer> new_args;
std::vector<TupleElement> arg_types;
auto new_types = types;
if (expected && expected->tag() != ValKind::TupleValue) {
FATAL_COMPILATION_ERROR(e->line_num) << "didn't expect a tuple";
}
if (expected && e->GetTupleLiteral().fields.size() !=
expected->GetTupleValue().elements.size()) {
FATAL_COMPILATION_ERROR(e->line_num) << "tuples of different length";
}
int i = 0;
for (auto arg = e->GetTupleLiteral().fields.begin();
arg != e->GetTupleLiteral().fields.end(); ++arg, ++i) {
const Value* arg_expected = nullptr;
if (expected && expected->tag() == ValKind::TupleValue) {
if (expected->GetTupleValue().elements[i].name != arg->name) {
FATAL_COMPILATION_ERROR(e->line_num)
<< "field names do not match, expected "
<< expected->GetTupleValue().elements[i].name << " but got "
<< arg->name;
}
arg_expected = expected->GetTupleValue().elements[i].value;
}
auto arg_res = TypeCheckExp(arg->expression, new_types, values,
arg_expected, context);
auto arg_res = TypeCheckExp(arg->expression, new_types, values);
new_types = arg_res.types;
new_args.push_back({.name = arg->name, .expression = arg_res.exp});
new_args.push_back(FieldInitializer(arg->name, arg_res.exp));
arg_types.push_back({.name = arg->name, .value = arg_res.type});
}
auto tuple_e = Expression::MakeTupleLiteral(e->line_num, new_args);
auto tuple_t = Value::MakeTupleValue(std::move(arg_types));
return TCResult(tuple_e, tuple_t, new_types);
return TCExpression(tuple_e, tuple_t, new_types);
}
case ExpressionKind::FieldAccessExpression: {
auto res = TypeCheckExp(e->GetFieldAccessExpression().aggregate, types,
values, nullptr, TCContext::ValueContext);
auto res =
TypeCheckExp(e->GetFieldAccessExpression().aggregate, types, values);
auto t = res.type;
switch (t->tag()) {
case ValKind::StructType:
@@ -346,7 +290,7 @@ auto TypeCheckExp(const Expression* e, TypeEnv types, Env values,
if (e->GetFieldAccessExpression().field == field.first) {
const Expression* new_e = Expression::MakeFieldAccessExpression(
e->line_num, res.exp, e->GetFieldAccessExpression().field);
return TCResult(new_e, field.second, res.types);
return TCExpression(new_e, field.second, res.types);
}
}
// Search for a method
@@ -354,7 +298,7 @@ auto TypeCheckExp(const Expression* e, TypeEnv types, Env values,
if (e->GetFieldAccessExpression().field == method.first) {
const Expression* new_e = Expression::MakeFieldAccessExpression(
e->line_num, res.exp, e->GetFieldAccessExpression().field);
return TCResult(new_e, method.second, res.types);
return TCExpression(new_e, method.second, res.types);
}
}
FATAL_COMPILATION_ERROR(e->line_num)
@@ -366,7 +310,7 @@ auto TypeCheckExp(const Expression* e, TypeEnv types, Env values,
if (e->GetFieldAccessExpression().field == field.name) {
auto new_e = Expression::MakeFieldAccessExpression(
e->line_num, res.exp, e->GetFieldAccessExpression().field);
return TCResult(new_e, field.value, res.types);
return TCExpression(new_e, field.value, res.types);
}
}
FATAL_COMPILATION_ERROR(e->line_num)
@@ -380,7 +324,7 @@ auto TypeCheckExp(const Expression* e, TypeEnv types, Env values,
const Expression* new_e = Expression::MakeFieldAccessExpression(
e->line_num, res.exp, e->GetFieldAccessExpression().field);
auto fun_ty = Value::MakeFunctionType({}, vt->second, t);
return TCResult(new_e, fun_ty, res.types);
return TCExpression(new_e, fun_ty, res.types);
}
}
FATAL_COMPILATION_ERROR(e->line_num)
@@ -398,24 +342,23 @@ auto TypeCheckExp(const Expression* e, TypeEnv types, Env values,
std::optional<const Value*> type =
types.Get(e->GetIdentifierExpression().name);
if (type) {
return TCResult(e, *type, types);
return TCExpression(e, *type, types);
} else {
FATAL_COMPILATION_ERROR(e->line_num)
<< "could not find `" << e->GetIdentifierExpression().name << "`";
}
}
case ExpressionKind::IntLiteral:
return TCResult(e, Value::MakeIntType(), types);
return TCExpression(e, Value::MakeIntType(), types);
case ExpressionKind::BoolLiteral:
return TCResult(e, Value::MakeBoolType(), types);
return TCExpression(e, Value::MakeBoolType(), types);
case ExpressionKind::PrimitiveOperatorExpression: {
std::vector<const Expression*> es;
std::vector<const Value*> ts;
auto new_types = types;
for (const Expression* argument :
e->GetPrimitiveOperatorExpression().arguments) {
auto res = TypeCheckExp(argument, types, values, nullptr,
TCContext::ValueContext);
auto res = TypeCheckExp(argument, types, values);
new_types = res.types;
es.push_back(res.exp);
ts.push_back(res.type);
@@ -425,55 +368,54 @@ auto TypeCheckExp(const Expression* e, TypeEnv types, Env values,
switch (e->GetPrimitiveOperatorExpression().op) {
case Operator::Neg:
ExpectType(e->line_num, "negation", Value::MakeIntType(), ts[0]);
return TCResult(new_e, Value::MakeIntType(), new_types);
return TCExpression(new_e, Value::MakeIntType(), new_types);
case Operator::Add:
ExpectType(e->line_num, "addition(1)", Value::MakeIntType(), ts[0]);
ExpectType(e->line_num, "addition(2)", Value::MakeIntType(), ts[1]);
return TCResult(new_e, Value::MakeIntType(), new_types);
return TCExpression(new_e, Value::MakeIntType(), new_types);
case Operator::Sub:
ExpectType(e->line_num, "subtraction(1)", Value::MakeIntType(),
ts[0]);
ExpectType(e->line_num, "subtraction(2)", Value::MakeIntType(),
ts[1]);
return TCResult(new_e, Value::MakeIntType(), new_types);
return TCExpression(new_e, Value::MakeIntType(), new_types);
case Operator::Mul:
ExpectType(e->line_num, "multiplication(1)", Value::MakeIntType(),
ts[0]);
ExpectType(e->line_num, "multiplication(2)", Value::MakeIntType(),
ts[1]);
return TCResult(new_e, Value::MakeIntType(), new_types);
return TCExpression(new_e, Value::MakeIntType(), new_types);
case Operator::And:
ExpectType(e->line_num, "&&(1)", Value::MakeBoolType(), ts[0]);
ExpectType(e->line_num, "&&(2)", Value::MakeBoolType(), ts[1]);
return TCResult(new_e, Value::MakeBoolType(), new_types);
return TCExpression(new_e, Value::MakeBoolType(), new_types);
case Operator::Or:
ExpectType(e->line_num, "||(1)", Value::MakeBoolType(), ts[0]);
ExpectType(e->line_num, "||(2)", Value::MakeBoolType(), ts[1]);
return TCResult(new_e, Value::MakeBoolType(), new_types);
return TCExpression(new_e, Value::MakeBoolType(), new_types);
case Operator::Not:
ExpectType(e->line_num, "!", Value::MakeBoolType(), ts[0]);
return TCResult(new_e, Value::MakeBoolType(), new_types);
return TCExpression(new_e, Value::MakeBoolType(), new_types);
case Operator::Eq:
ExpectType(e->line_num, "==", ts[0], ts[1]);
return TCResult(new_e, Value::MakeBoolType(), new_types);
return TCExpression(new_e, Value::MakeBoolType(), new_types);
case Operator::Deref:
ExpectPointerType(e->line_num, "*", ts[0]);
return TCResult(new_e, ts[0]->GetPointerType().type, new_types);
return TCExpression(new_e, ts[0]->GetPointerType().type, new_types);
case Operator::Ptr:
ExpectType(e->line_num, "*", Value::MakeTypeType(), ts[0]);
return TCResult(new_e, Value::MakeTypeType(), new_types);
return TCExpression(new_e, Value::MakeTypeType(), new_types);
}
break;
}
case ExpressionKind::CallExpression: {
auto fun_res = TypeCheckExp(e->GetCallExpression().function, types,
values, nullptr, TCContext::ValueContext);
auto fun_res =
TypeCheckExp(e->GetCallExpression().function, types, values);
switch (fun_res.type->tag()) {
case ValKind::FunctionType: {
auto fun_t = fun_res.type;
auto arg_res =
TypeCheckExp(e->GetCallExpression().argument, fun_res.types,
values, fun_t->GetFunctionType().param, context);
auto arg_res = TypeCheckExp(e->GetCallExpression().argument,
fun_res.types, values);
auto parameter_type = fun_t->GetFunctionType().param;
auto return_type = fun_t->GetFunctionType().ret;
if (fun_t->GetFunctionType().deduced.size() > 0) {
@@ -497,7 +439,7 @@ auto TypeCheckExp(const Expression* e, TypeEnv types, Env values,
}
auto new_e = Expression::MakeCallExpression(e->line_num, fun_res.exp,
arg_res.exp);
return TCResult(new_e, return_type, arg_res.types);
return TCExpression(new_e, return_type, arg_res.types);
}
default: {
FATAL_COMPILATION_ERROR(e->line_num)
@@ -508,48 +450,157 @@ auto TypeCheckExp(const Expression* e, TypeEnv types, Env values,
break;
}
case ExpressionKind::FunctionTypeLiteral: {
switch (context) {
case TCContext::ValueContext:
case TCContext::TypeContext: {
auto pt = InterpExp(values, e->GetFunctionTypeLiteral().parameter);
auto rt = InterpExp(values, e->GetFunctionTypeLiteral().return_type);
auto new_e = Expression::MakeFunctionTypeLiteral(
e->line_num, ReifyType(pt, e->line_num),
ReifyType(rt, e->line_num));
return TCResult(new_e, Value::MakeTypeType(), types);
}
case TCContext::PatternContext: {
auto param_res = TypeCheckExp(e->GetFunctionTypeLiteral().parameter,
types, values, nullptr, context);
auto ret_res =
TypeCheckExp(e->GetFunctionTypeLiteral().return_type,
param_res.types, values, nullptr, context);
auto new_e = Expression::MakeFunctionTypeLiteral(
e->line_num, ReifyType(param_res.type, e->line_num),
ReifyType(ret_res.type, e->line_num));
return TCResult(new_e, Value::MakeTypeType(), ret_res.types);
}
}
auto pt = InterpExp(values, e->GetFunctionTypeLiteral().parameter);
auto rt = InterpExp(values, e->GetFunctionTypeLiteral().return_type);
auto new_e = Expression::MakeFunctionTypeLiteral(
e->line_num, ReifyType(pt, e->line_num), ReifyType(rt, e->line_num));
return TCExpression(new_e, Value::MakeTypeType(), types);
}
case ExpressionKind::IntTypeLiteral:
return TCResult(e, Value::MakeTypeType(), types);
return TCExpression(e, Value::MakeTypeType(), types);
case ExpressionKind::BoolTypeLiteral:
return TCResult(e, Value::MakeTypeType(), types);
return TCExpression(e, Value::MakeTypeType(), types);
case ExpressionKind::TypeTypeLiteral:
return TCResult(e, Value::MakeTypeType(), types);
case ExpressionKind::AutoTypeLiteral:
return TCResult(e, Value::MakeTypeType(), types);
return TCExpression(e, Value::MakeTypeType(), types);
case ExpressionKind::ContinuationTypeLiteral:
return TCResult(e, Value::MakeTypeType(), types);
return TCExpression(e, Value::MakeTypeType(), types);
}
}
auto TypecheckCase(const Value* expected, const Expression* pat,
// Equivalent to TypeCheckExp, but operates on Patterns instead of Expressions.
// `expected` is the type that this pattern is expected to have, if the
// surrounding context gives us that information. Otherwise, it is null.
auto TypeCheckPattern(const Pattern* p, TypeEnv types, Env values,
const Value* expected) -> TCPattern {
if (tracing_output) {
llvm::outs() << "checking pattern, ";
if (expected) {
llvm::outs() << "expecting " << *expected;
}
llvm::outs() << ", " << *p << "\n";
}
switch (p->Tag()) {
case Pattern::Kind::AutoPattern: {
return {.pattern = p, .type = Value::MakeTypeType(), .types = types};
}
case Pattern::Kind::BindingPattern: {
const auto& binding = cast<BindingPattern>(*p);
const Value* type;
switch (binding.Type()->Tag()) {
case Pattern::Kind::AutoPattern: {
if (expected == nullptr) {
FATAL_COMPILATION_ERROR(binding.LineNumber())
<< "auto not allowed here";
} else {
type = expected;
}
break;
}
case Pattern::Kind::ExpressionPattern: {
type = InterpExp(
values, cast<ExpressionPattern>(binding.Type())->Expression());
CHECK(type->tag() != ValKind::AutoType);
if (expected != nullptr) {
ExpectType(binding.LineNumber(), "pattern variable", type,
expected);
}
break;
}
case Pattern::Kind::TuplePattern:
case Pattern::Kind::BindingPattern:
case Pattern::Kind::AlternativePattern:
FATAL_COMPILATION_ERROR(binding.LineNumber())
<< "Unsupported type pattern";
}
auto new_p = new BindingPattern(
binding.LineNumber(), binding.Name(),
new ExpressionPattern(ReifyType(type, binding.LineNumber())));
if (binding.Name().has_value()) {
types.Set(*binding.Name(), type);
}
return {.pattern = new_p, .type = type, .types = types};
}
case Pattern::Kind::TuplePattern: {
const auto& tuple = cast<TuplePattern>(*p);
std::vector<TuplePattern::Field> new_fields;
std::vector<TupleElement> field_types;
auto new_types = types;
if (expected && expected->tag() != ValKind::TupleValue) {
FATAL_COMPILATION_ERROR(p->LineNumber()) << "didn't expect a tuple";
}
if (expected &&
tuple.Fields().size() != expected->GetTupleValue().elements.size()) {
FATAL_COMPILATION_ERROR(tuple.LineNumber())
<< "tuples of different length";
}
for (size_t i = 0; i < tuple.Fields().size(); ++i) {
const TuplePattern::Field& field = tuple.Fields()[i];
const Value* expected_field_type = nullptr;
if (expected != nullptr) {
const TupleElement& expected_element =
expected->GetTupleValue().elements[i];
if (expected_element.name != field.name) {
FATAL_COMPILATION_ERROR(tuple.LineNumber())
<< "field names do not match, expected "
<< expected_element.name << " but got " << field.name;
}
expected_field_type = expected_element.value;
}
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));
field_types.push_back({.name = field.name, .value = field_result.type});
}
auto new_tuple = new TuplePattern(tuple.LineNumber(), new_fields);
auto tuple_t = Value::MakeTupleValue(std::move(field_types));
return {.pattern = new_tuple, .type = tuple_t, .types = new_types};
}
case Pattern::Kind::AlternativePattern: {
const auto& alternative = cast<AlternativePattern>(*p);
const Value* choice_type = InterpExp(values, alternative.ChoiceType());
if (choice_type->tag() != ValKind::ChoiceType) {
FATAL_COMPILATION_ERROR(alternative.LineNumber())
<< "alternative pattern does not name a choice type.";
}
if (expected != nullptr) {
ExpectType(alternative.LineNumber(), "alternative pattern", expected,
choice_type);
}
const Value* parameter_types =
FindInVarValues(alternative.AlternativeName(),
choice_type->GetChoiceType().alternatives);
if (parameter_types == nullptr) {
FATAL_COMPILATION_ERROR(alternative.LineNumber())
<< "'" << alternative.AlternativeName()
<< "' is not an alternative of " << choice_type;
}
TCPattern arg_results = TypeCheckPattern(alternative.Arguments(), types,
values, parameter_types);
return {.pattern = new AlternativePattern(
alternative.LineNumber(),
ReifyType(choice_type, alternative.LineNumber()),
alternative.AlternativeName(),
cast<TuplePattern>(arg_results.pattern)),
.type = choice_type,
.types = arg_results.types};
}
case Pattern::Kind::ExpressionPattern: {
TCExpression result =
TypeCheckExp(cast<ExpressionPattern>(p)->Expression(), types, values);
return {.pattern = new ExpressionPattern(result.exp),
.type = result.type,
.types = result.types};
}
}
}
auto TypecheckCase(const Value* expected, const Pattern* pat,
const Statement* body, TypeEnv types, Env values,
const Value*& ret_type)
-> std::pair<const Expression*, const Statement*> {
auto pat_res =
TypeCheckExp(pat, types, values, expected, TCContext::PatternContext);
-> std::pair<const Pattern*, const Statement*> {
auto pat_res = TypeCheckPattern(pat, types, values, expected);
auto res = TypeCheckStmt(body, pat_res.types, values, ret_type);
return std::make_pair(pat, res.stmt);
}
@@ -568,11 +619,10 @@ auto TypeCheckStmt(const Statement* s, TypeEnv types, Env values,
}
switch (s->tag()) {
case StatementKind::Match: {
auto res = TypeCheckExp(s->GetMatch().exp, types, values, nullptr,
TCContext::ValueContext);
auto res = TypeCheckExp(s->GetMatch().exp, types, values);
auto res_type = res.type;
auto new_clauses =
new std::list<std::pair<const Expression*, const Statement*>>();
new std::list<std::pair<const Pattern*, const Statement*>>();
for (auto& clause : *s->GetMatch().clauses) {
new_clauses->push_back(TypecheckCase(
res_type, clause.first, clause.second, types, values, ret_type));
@@ -582,8 +632,7 @@ auto TypeCheckStmt(const Statement* s, TypeEnv types, Env values,
return TCStatement(new_s, types);
}
case StatementKind::While: {
auto cnd_res = TypeCheckExp(s->GetWhile().cond, types, values, nullptr,
TCContext::ValueContext);
auto cnd_res = TypeCheckExp(s->GetWhile().cond, types, values);
ExpectType(s->line_num, "condition of `while`", Value::MakeBoolType(),
cnd_res.type);
auto body_res =
@@ -602,11 +651,10 @@ auto TypeCheckStmt(const Statement* s, TypeEnv types, Env values,
types);
}
case StatementKind::VariableDefinition: {
auto res = TypeCheckExp(s->GetVariableDefinition().init, types, values,
nullptr, TCContext::ValueContext);
auto res = TypeCheckExp(s->GetVariableDefinition().init, types, values);
const Value* rhs_ty = res.type;
auto lhs_res = TypeCheckExp(s->GetVariableDefinition().pat, types, values,
rhs_ty, TCContext::PatternContext);
auto lhs_res = TypeCheckPattern(s->GetVariableDefinition().pat, types,
values, rhs_ty);
const Statement* new_s = Statement::MakeVariableDefinition(
s->line_num, s->GetVariableDefinition().pat, res.exp);
return TCStatement(new_s, lhs_res.types);
@@ -623,25 +671,21 @@ auto TypeCheckStmt(const Statement* s, TypeEnv types, Env values,
types3);
}
case StatementKind::Assign: {
auto rhs_res = TypeCheckExp(s->GetAssign().rhs, types, values, nullptr,
TCContext::ValueContext);
auto rhs_res = TypeCheckExp(s->GetAssign().rhs, types, values);
auto rhs_t = rhs_res.type;
auto lhs_res = TypeCheckExp(s->GetAssign().lhs, types, values, rhs_t,
TCContext::ValueContext);
auto lhs_res = TypeCheckExp(s->GetAssign().lhs, types, values);
auto lhs_t = lhs_res.type;
ExpectType(s->line_num, "assign", lhs_t, rhs_t);
auto new_s = Statement::MakeAssign(s->line_num, lhs_res.exp, rhs_res.exp);
return TCStatement(new_s, lhs_res.types);
}
case StatementKind::ExpressionStatement: {
auto res = TypeCheckExp(s->GetExpressionStatement().exp, types, values,
nullptr, TCContext::ValueContext);
auto res = TypeCheckExp(s->GetExpressionStatement().exp, types, values);
auto new_s = Statement::MakeExpressionStatement(s->line_num, res.exp);
return TCStatement(new_s, types);
}
case StatementKind::If: {
auto cnd_res = TypeCheckExp(s->GetIf().cond, types, values, nullptr,
TCContext::ValueContext);
auto cnd_res = TypeCheckExp(s->GetIf().cond, types, values);
ExpectType(s->line_num, "condition of `if`", Value::MakeBoolType(),
cnd_res.type);
auto thn_res =
@@ -653,8 +697,7 @@ auto TypeCheckStmt(const Statement* s, TypeEnv types, Env values,
return TCStatement(new_s, types);
}
case StatementKind::Return: {
auto res = TypeCheckExp(s->GetReturn().exp, types, values, nullptr,
TCContext::ValueContext);
auto res = TypeCheckExp(s->GetReturn().exp, types, values);
if (ret_type->tag() == ValKind::AutoType) {
// The following infers the return type from the first 'return'
// statement. This will get more difficult with subtyping, when we
@@ -676,9 +719,8 @@ auto TypeCheckStmt(const Statement* s, TypeEnv types, Env values,
return TCStatement(new_continuation, types);
}
case StatementKind::Run: {
TCResult argument_result =
TypeCheckExp(s->GetRun().argument, types, values, nullptr,
TCContext::ValueContext);
TCExpression argument_result =
TypeCheckExp(s->GetRun().argument, types, values);
ExpectType(s->line_num, "argument of `run`",
Value::MakeContinuationType(), argument_result.type);
const Statement* new_run =
@@ -706,7 +748,7 @@ auto CheckOrEnsureReturn(const Statement* stmt, bool void_return, int line_num)
switch (stmt->tag()) {
case StatementKind::Match: {
auto new_clauses =
new std::list<std::pair<const Expression*, const Statement*>>();
new std::list<std::pair<const Pattern*, const Statement*>>();
for (auto i = stmt->GetMatch().clauses->begin();
i != stmt->GetMatch().clauses->end(); ++i) {
auto s = CheckOrEnsureReturn(i->second, void_return, stmt->line_num);
@@ -774,10 +816,9 @@ auto TypeCheckFunDef(const FunctionDefinition* f, TypeEnv types, Env values)
values.Set(deduced.name, a);
}
// Type check the parameter pattern
auto param_res = TypeCheckExp(f->param_pattern, types, values, nullptr,
TCContext::PatternContext);
auto param_res = TypeCheckPattern(f->param_pattern, types, values, nullptr);
// Evaluate the return type expression
auto return_type = InterpExp(values, f->return_type);
auto return_type = InterpPattern(values, f->return_type);
if (f->name == "main") {
ExpectType(f->line_num, "return type of `main`", Value::MakeIntType(),
return_type);
@@ -786,9 +827,9 @@ auto TypeCheckFunDef(const FunctionDefinition* f, TypeEnv types, Env values)
auto res = TypeCheckStmt(f->body, param_res.types, values, return_type);
bool void_return = TypeEqual(return_type, Value::MakeUnitTypeVal());
auto body = CheckOrEnsureReturn(res.stmt, void_return, f->line_num);
return new FunctionDefinition(f->line_num, f->name, f->deduced_parameters,
f->param_pattern,
ReifyType(return_type, f->line_num), body);
return new FunctionDefinition(
f->line_num, f->name, f->deduced_parameters, f->param_pattern,
new ExpressionPattern(ReifyType(return_type, f->line_num)), body);
}
auto TypeOfFunDef(TypeEnv types, Env values, const FunctionDefinition* fun_def)
@@ -801,13 +842,13 @@ auto TypeOfFunDef(TypeEnv types, Env values, const FunctionDefinition* fun_def)
values.Set(deduced.name, a);
}
// Type check the parameter pattern
auto param_res = TypeCheckExp(fun_def->param_pattern, types, values, nullptr,
TCContext::PatternContext);
auto param_res =
TypeCheckPattern(fun_def->param_pattern, types, values, nullptr);
// Evaluate the return type expression
auto ret = InterpExp(values, fun_def->return_type);
auto ret = InterpPattern(values, fun_def->return_type);
if (ret->tag() == ValKind::AutoType) {
auto f = TypeCheckFunDef(fun_def, types, values);
ret = InterpExp(values, f->return_type);
ret = InterpPattern(values, f->return_type);
}
return Value::MakeFunctionType(fun_def->deduced_parameters, param_res.type,
ret);
@@ -819,10 +860,22 @@ auto TypeOfStructDef(const StructDefinition* sd, TypeEnv /*types*/, Env ct_top)
VarValues methods;
for (const Member* m : sd->members) {
switch (m->tag()) {
case MemberKind::FieldMember:
auto t = InterpExp(ct_top, m->GetFieldMember().type);
fields.push_back(std::make_pair(m->GetFieldMember().name, t));
case MemberKind::FieldMember: {
const BindingPattern* binding = m->GetFieldMember().binding;
if (!binding->Name().has_value()) {
FATAL_COMPILATION_ERROR(binding->LineNumber())
<< "Struct members must have names";
}
const Expression* type_expression =
dyn_cast<ExpressionPattern>(binding->Type())->Expression();
if (type_expression == nullptr) {
FATAL_COMPILATION_ERROR(binding->LineNumber())
<< "Struct members must have explicit types";
}
auto type = InterpExp(ct_top, type_expression);
fields.push_back(std::make_pair(*binding->Name(), type));
break;
}
}
}
return Value::MakeStructType(sd->name, std::move(fields), std::move(methods));
@@ -836,8 +889,14 @@ static auto GetName(const Declaration& d) -> const std::string& {
return d.GetStructDeclaration().definition.name;
case DeclarationKind::ChoiceDeclaration:
return d.GetChoiceDeclaration().name;
case DeclarationKind::VariableDeclaration:
return d.GetVariableDeclaration().name;
case DeclarationKind::VariableDeclaration: {
const BindingPattern* binding = d.GetVariableDeclaration().binding;
if (!binding->Name().has_value()) {
FATAL_COMPILATION_ERROR(binding->LineNumber())
<< "Top-level variable declarations must have names";
}
return *binding->Name();
}
}
}
@@ -872,9 +931,16 @@ auto MakeTypeChecked(const Declaration& d, const TypeEnv& types,
// Signals a type error if the initializing expression does not have
// the declared type of the variable, otherwise returns this
// declaration with annotated types.
TCResult type_checked_initializer = TypeCheckExp(
var.initializer, types, values, nullptr, TCContext::ValueContext);
const Value* declared_type = InterpExp(values, var.type);
TCExpression type_checked_initializer =
TypeCheckExp(var.initializer, types, values);
const Expression* type =
dyn_cast<ExpressionPattern>(var.binding->Type())->Expression();
if (type == nullptr) {
// TODO: consider adding support for `auto`
FATAL_COMPILATION_ERROR(var.source_location)
<< "Type of a top-level variable must be an expression.";
}
const Value* declared_type = InterpExp(values, type);
ExpectType(var.source_location, "initializer of variable", declared_type,
type_checked_initializer.type);
return d;
@@ -926,8 +992,10 @@ static void TopLevel(const Declaration& d, TypeCheckContext* tops) {
const auto& var = d.GetVariableDeclaration();
// Associate the variable name with it's declared type in the
// compile-time symbol table.
const Value* declared_type = InterpExp(tops->values, var.type);
tops->types.Set(var.name, declared_type);
const Expression* type =
cast<ExpressionPattern>(var.binding->Type())->Expression();
const Value* declared_type = InterpExp(tops->values, type);
tops->types.Set(*var.binding->Name(), declared_type);
break;
}
}
+12 -6
View File
@@ -17,10 +17,8 @@ namespace Carbon {
using TypeEnv = Dictionary<std::string, const Value*>;
enum class TCContext { ValueContext, PatternContext, TypeContext };
struct TCResult {
TCResult(const Expression* e, const Value* t, TypeEnv types)
struct TCExpression {
TCExpression(const Expression* e, const Value* t, TypeEnv types)
: exp(e), type(t), types(types) {}
const Expression* exp;
@@ -28,6 +26,12 @@ struct TCResult {
TypeEnv types;
};
struct TCPattern {
const Pattern* pattern;
const Value* type;
TypeEnv types;
};
struct TCStatement {
TCStatement(const Statement* s, TypeEnv types) : stmt(s), types(types) {}
@@ -35,8 +39,10 @@ struct TCStatement {
TypeEnv types;
};
auto TypeCheckExp(const Expression* e, TypeEnv types, Env values,
const Value* expected, TCContext context) -> TCResult;
auto TypeCheckExp(const Expression* e, TypeEnv types, Env values)
-> TCExpression;
auto TypeCheckPattern(const Pattern* p, TypeEnv types, Env values,
const Value* expected) -> TCPattern;
auto TypeCheckStmt(const Statement*, TypeEnv, Env, Value const*&)
-> TCStatement;