Switch Value to use the inherit/cast model (#703)

This commit is contained in:
Jon Meow
2021-08-05 08:34:08 -07:00
committed by GitHub
parent 3116c6ff5d
commit 7e91fbc276
6 changed files with 841 additions and 828 deletions
+131 -125
View File
@@ -75,28 +75,29 @@ auto EvalPrim(Operator op, const std::vector<const Value*>& args, int line_num)
-> const Value* {
switch (op) {
case Operator::Neg:
return Value::MakeIntValue(-args[0]->GetIntValue());
return global_arena->New<IntValue>(-cast<IntValue>(*args[0]).Val());
case Operator::Add:
return Value::MakeIntValue(args[0]->GetIntValue() +
args[1]->GetIntValue());
return global_arena->New<IntValue>(cast<IntValue>(*args[0]).Val() +
cast<IntValue>(*args[1]).Val());
case Operator::Sub:
return Value::MakeIntValue(args[0]->GetIntValue() -
args[1]->GetIntValue());
return global_arena->New<IntValue>(cast<IntValue>(*args[0]).Val() -
cast<IntValue>(*args[1]).Val());
case Operator::Mul:
return Value::MakeIntValue(args[0]->GetIntValue() *
args[1]->GetIntValue());
return global_arena->New<IntValue>(cast<IntValue>(*args[0]).Val() *
cast<IntValue>(*args[1]).Val());
case Operator::Not:
return Value::MakeBoolValue(!args[0]->GetBoolValue());
return global_arena->New<BoolValue>(!cast<BoolValue>(*args[0]).Val());
case Operator::And:
return Value::MakeBoolValue(args[0]->GetBoolValue() &&
args[1]->GetBoolValue());
return global_arena->New<BoolValue>(cast<BoolValue>(*args[0]).Val() &&
cast<BoolValue>(*args[1]).Val());
case Operator::Or:
return Value::MakeBoolValue(args[0]->GetBoolValue() ||
args[1]->GetBoolValue());
return global_arena->New<BoolValue>(cast<BoolValue>(*args[0]).Val() ||
cast<BoolValue>(*args[1]).Val());
case Operator::Eq:
return Value::MakeBoolValue(ValueEqual(args[0], args[1], line_num));
return global_arena->New<BoolValue>(
ValueEqual(args[0], args[1], line_num));
case Operator::Ptr:
return Value::MakePointerType(args[0]);
return global_arena->New<PointerType>(args[0]);
case Operator::Deref:
llvm::errs() << line_num << ": dereference not implemented yet\n";
exit(-1);
@@ -114,12 +115,13 @@ void InitEnv(const Declaration& d, Env* env) {
Env new_env = *env;
// Bring the deduced parameters into scope.
for (const auto& deduced : func_def.deduced_parameters) {
Address a =
state->heap.AllocateValue(Value::MakeVariableType(deduced.name));
Address a = state->heap.AllocateValue(
global_arena->New<VariableType>(deduced.name));
new_env.Set(deduced.name, a);
}
auto pt = InterpPattern(new_env, func_def.param_pattern);
auto f = Value::MakeFunctionValue(func_def.name, pt, func_def.body);
auto f =
global_arena->New<FunctionValue>(func_def.name, pt, func_def.body);
Address a = state->heap.AllocateValue(f);
env->Set(func_def.name, a);
break;
@@ -141,8 +143,8 @@ void InitEnv(const Declaration& d, Env* env) {
}
}
}
auto st = Value::MakeStructType(struct_def.name, std::move(fields),
std::move(methods));
auto st = global_arena->New<StructType>(
struct_def.name, std::move(fields), std::move(methods));
auto a = state->heap.AllocateValue(st);
env->Set(struct_def.name, a);
break;
@@ -155,7 +157,7 @@ void InitEnv(const Declaration& d, Env* env) {
auto t = InterpExp(Env(), signature);
alts.push_back(make_pair(name, t));
}
auto ct = Value::MakeChoiceType(choice.name, std::move(alts));
auto ct = global_arena->New<ChoiceType>(choice.name, std::move(alts));
auto a = state->heap.AllocateValue(ct);
env->Set(choice.name, a);
break;
@@ -185,35 +187,34 @@ static void InitGlobals(std::list<Declaration>* fs) {
// F is the function
void CallFunction(int line_num, std::vector<const Value*> operas,
State* state) {
switch (operas[0]->tag()) {
case ValKind::FunctionValue: {
switch (operas[0]->Tag()) {
case Value::Kind::FunctionValue: {
const auto& fn = cast<FunctionValue>(*operas[0]);
// Bind arguments to parameters
std::list<std::string> params;
std::optional<Env> matches =
PatternMatch(operas[0]->GetFunctionValue().param, operas[1], globals,
&params, line_num);
PatternMatch(fn.Param(), operas[1], globals, &params, line_num);
CHECK(matches) << "internal error in call_function, pattern match failed";
// Create the new frame and push it on the stack
auto* scope = global_arena->New<Scope>(*matches, params);
auto* frame = global_arena->New<Frame>(
operas[0]->GetFunctionValue().name, Stack(scope),
Stack(
Action::MakeStatementAction(operas[0]->GetFunctionValue().body)));
fn.Name(), Stack(scope),
Stack(Action::MakeStatementAction(fn.Body())));
state->stack.Push(frame);
break;
}
case ValKind::StructType: {
case Value::Kind::StructType: {
const Value* arg = CopyVal(operas[1], line_num);
const Value* sv = Value::MakeStructValue(operas[0], arg);
const Value* sv = global_arena->New<StructValue>(operas[0], arg);
Frame* frame = state->stack.Top();
frame->todo.Push(Action::MakeValAction(sv));
break;
}
case ValKind::AlternativeConstructorValue: {
case Value::Kind::AlternativeConstructorValue: {
const auto& alt = cast<AlternativeConstructorValue>(*operas[0]);
const Value* arg = CopyVal(operas[1], line_num);
const Value* av = Value::MakeAlternativeValue(
operas[0]->GetAlternativeConstructorValue().alt_name,
operas[0]->GetAlternativeConstructorValue().choice_name, arg);
const Value* av = global_arena->New<AlternativeValue>(
alt.AltName(), alt.ChoiceName(), arg);
Frame* frame = state->stack.Top();
frame->todo.Push(Action::MakeValAction(av));
break;
@@ -248,7 +249,7 @@ void CreateTuple(Frame* frame, Action* act, const Expression* exp) {
for (auto i = act->results.begin(); i != act->results.end(); ++i, ++f) {
elements.push_back({.name = f->name, .value = *i});
}
const Value* tv = Value::MakeTupleValue(std::move(elements));
const Value* tv = global_arena->New<TupleValue>(std::move(elements));
frame->todo.Pop(1);
frame->todo.Push(Action::MakeValAction(tv));
}
@@ -261,29 +262,28 @@ void CreateTuple(Frame* frame, Action* act, const Expression* exp) {
auto PatternMatch(const Value* p, const Value* v, Env values,
std::list<std::string>* vars, int line_num)
-> std::optional<Env> {
switch (p->tag()) {
case ValKind::BindingPlaceholderValue: {
const BindingPlaceholderValue& placeholder =
p->GetBindingPlaceholderValue();
if (placeholder.name.has_value()) {
switch (p->Tag()) {
case Value::Kind::BindingPlaceholderValue: {
const auto& placeholder = cast<BindingPlaceholderValue>(*p);
if (placeholder.Name().has_value()) {
Address a = state->heap.AllocateValue(CopyVal(v, line_num));
vars->push_back(*placeholder.name);
values.Set(*placeholder.name, a);
vars->push_back(*placeholder.Name());
values.Set(*placeholder.Name(), a);
}
return values;
}
case ValKind::TupleValue:
switch (v->tag()) {
case ValKind::TupleValue: {
if (p->GetTupleValue().elements.size() !=
v->GetTupleValue().elements.size()) {
case Value::Kind::TupleValue:
switch (v->Tag()) {
case Value::Kind::TupleValue: {
const auto& p_tup = cast<TupleValue>(*p);
const auto& v_tup = cast<TupleValue>(*v);
if (p_tup.Elements().size() != v_tup.Elements().size()) {
FATAL_RUNTIME_ERROR(line_num)
<< "arity mismatch in tuple pattern match";
<< "arity mismatch in tuple pattern match:\n pattern: "
<< p_tup << "\n value: " << v_tup;
}
for (const TupleElement& pattern_element :
p->GetTupleValue().elements) {
const Value* value_field =
v->GetTupleValue().FindField(pattern_element.name);
for (const TupleElement& pattern_element : p_tup.Elements()) {
const Value* value_field = v_tup.FindField(pattern_element.name);
if (value_field == nullptr) {
FATAL_RUNTIME_ERROR(line_num)
<< "field " << pattern_element.name << "not in " << *v;
@@ -303,18 +303,17 @@ auto PatternMatch(const Value* p, const Value* v, Env values,
<< "\n";
exit(-1);
}
case ValKind::AlternativeValue:
switch (v->tag()) {
case ValKind::AlternativeValue: {
if (p->GetAlternativeValue().choice_name !=
v->GetAlternativeValue().choice_name ||
p->GetAlternativeValue().alt_name !=
v->GetAlternativeValue().alt_name) {
case Value::Kind::AlternativeValue:
switch (v->Tag()) {
case Value::Kind::AlternativeValue: {
const auto& p_alt = cast<AlternativeValue>(*p);
const auto& v_alt = cast<AlternativeValue>(*v);
if (p_alt.ChoiceName() != v_alt.ChoiceName() ||
p_alt.AltName() != v_alt.AltName()) {
return std::nullopt;
}
std::optional<Env> matches = PatternMatch(
p->GetAlternativeValue().argument,
v->GetAlternativeValue().argument, values, vars, line_num);
p_alt.Argument(), v_alt.Argument(), values, vars, line_num);
if (!matches) {
return std::nullopt;
}
@@ -327,18 +326,17 @@ auto PatternMatch(const Value* p, const Value* v, Env values,
<< *v << "\n";
exit(-1);
}
case ValKind::FunctionType:
switch (v->tag()) {
case ValKind::FunctionType: {
case Value::Kind::FunctionType:
switch (v->Tag()) {
case Value::Kind::FunctionType: {
const auto& p_fn = cast<FunctionType>(*p);
const auto& v_fn = cast<FunctionType>(*v);
std::optional<Env> matches =
PatternMatch(p->GetFunctionType().param,
v->GetFunctionType().param, values, vars, line_num);
PatternMatch(p_fn.Param(), v_fn.Param(), values, vars, line_num);
if (!matches) {
return std::nullopt;
}
return PatternMatch(p->GetFunctionType().ret,
v->GetFunctionType().ret, *matches, vars,
line_num);
return PatternMatch(p_fn.Ret(), v_fn.Ret(), *matches, vars, line_num);
}
default:
return std::nullopt;
@@ -353,23 +351,23 @@ auto PatternMatch(const Value* p, const Value* v, Env values,
}
void PatternAssignment(const Value* pat, const Value* val, int line_num) {
switch (pat->tag()) {
case ValKind::PointerValue:
state->heap.Write(pat->GetPointerValue(), CopyVal(val, line_num),
switch (pat->Tag()) {
case Value::Kind::PointerValue:
state->heap.Write(cast<PointerValue>(*pat).Val(), CopyVal(val, line_num),
line_num);
break;
case ValKind::TupleValue: {
switch (val->tag()) {
case ValKind::TupleValue: {
if (pat->GetTupleValue().elements.size() !=
val->GetTupleValue().elements.size()) {
case Value::Kind::TupleValue: {
switch (val->Tag()) {
case Value::Kind::TupleValue: {
const auto& pat_tup = cast<TupleValue>(*pat);
const auto& val_tup = cast<TupleValue>(*val);
if (pat_tup.Elements().size() != val_tup.Elements().size()) {
FATAL_RUNTIME_ERROR(line_num)
<< "arity mismatch in tuple pattern match";
<< "arity mismatch in tuple pattern assignment:\n pattern: "
<< pat_tup << "\n value: " << val_tup;
}
for (const TupleElement& pattern_element :
pat->GetTupleValue().elements) {
const Value* value_field =
val->GetTupleValue().FindField(pattern_element.name);
for (const TupleElement& pattern_element : pat_tup.Elements()) {
const Value* value_field = val_tup.FindField(pattern_element.name);
if (value_field == nullptr) {
FATAL_RUNTIME_ERROR(line_num)
<< "field " << pattern_element.name << "not in " << *val;
@@ -387,16 +385,15 @@ void PatternAssignment(const Value* pat, const Value* val, int line_num) {
}
break;
}
case ValKind::AlternativeValue: {
switch (val->tag()) {
case ValKind::AlternativeValue: {
CHECK(pat->GetAlternativeValue().choice_name ==
val->GetAlternativeValue().choice_name &&
pat->GetAlternativeValue().alt_name ==
val->GetAlternativeValue().alt_name)
case Value::Kind::AlternativeValue: {
switch (val->Tag()) {
case Value::Kind::AlternativeValue: {
const auto& pat_alt = cast<AlternativeValue>(*pat);
const auto& val_alt = cast<AlternativeValue>(*val);
CHECK(val_alt.ChoiceName() == pat_alt.ChoiceName() &&
val_alt.AltName() == pat_alt.AltName())
<< "internal error in pattern assignment";
PatternAssignment(pat->GetAlternativeValue().argument,
val->GetAlternativeValue().argument, line_num);
PatternAssignment(pat_alt.Argument(), val_alt.Argument(), line_num);
break;
}
default:
@@ -434,7 +431,7 @@ void StepLvalue() {
<< ": could not find `" << exp->GetIdentifierExpression().name
<< "`";
}
const Value* v = Value::MakePointerValue(*pointer);
const Value* v = global_arena->New<PointerValue>(*pointer);
frame->todo.Pop();
frame->todo.Push(Action::MakeValAction(v));
break;
@@ -449,11 +446,12 @@ void StepLvalue() {
} else {
// { v :: [].f :: C, E, F} :: S, H}
// -> { { &v.f :: C, E, F} :: S, H }
Address aggregate = act->results[0]->GetPointerValue();
Address aggregate = cast<PointerValue>(*act->results[0]).Val();
Address field =
aggregate.SubobjectAddress(exp->GetFieldAccessExpression().field);
frame->todo.Pop(1);
frame->todo.Push(Action::MakeValAction(Value::MakePointerValue(field)));
frame->todo.Push(
Action::MakeValAction(global_arena->New<PointerValue>(field)));
}
break;
}
@@ -471,11 +469,12 @@ void StepLvalue() {
} else if (act->pos == 2) {
// { v :: [][i] :: C, E, F} :: S, H}
// -> { { &v[i] :: C, E, F} :: S, H }
Address aggregate = act->results[0]->GetPointerValue();
std::string f = std::to_string(act->results[1]->GetIntValue());
Address aggregate = cast<PointerValue>(*act->results[0]).Val();
std::string f = std::to_string(cast<IntValue>(*act->results[1]).Val());
Address field = aggregate.SubobjectAddress(f);
frame->todo.Pop(1);
frame->todo.Push(Action::MakeValAction(Value::MakePointerValue(field)));
frame->todo.Push(
Action::MakeValAction(global_arena->New<PointerValue>(field)));
}
break;
}
@@ -539,12 +538,13 @@ void StepExp() {
act->pos++;
} else if (act->pos == 2) {
auto tuple = act->results[0];
switch (tuple->tag()) {
case ValKind::TupleValue: {
switch (tuple->Tag()) {
case Value::Kind::TupleValue: {
// { { v :: [][i] :: C, E, F} :: S, H}
// -> { { v_i :: C, E, F} : S, H}
std::string f = std::to_string(act->results[1]->GetIntValue());
const Value* field = tuple->GetTupleValue().FindField(f);
std::string f =
std::to_string(cast<IntValue>(*act->results[1]).Val());
const Value* field = cast<TupleValue>(*tuple).FindField(f);
if (field == nullptr) {
FATAL_RUNTIME_ERROR_NO_LINE()
<< "field " << f << " not in " << *tuple;
@@ -622,15 +622,15 @@ void StepExp() {
CHECK(act->pos == 0);
// { {n :: C, E, F} :: S, H} -> { {n' :: C, E, F} :: S, H}
frame->todo.Pop(1);
frame->todo.Push(
Action::MakeValAction(Value::MakeIntValue(exp->GetIntLiteral())));
frame->todo.Push(Action::MakeValAction(
global_arena->New<IntValue>(exp->GetIntLiteral())));
break;
case ExpressionKind::BoolLiteral:
CHECK(act->pos == 0);
// { {n :: C, E, F} :: S, H} -> { {n' :: C, E, F} :: S, H}
frame->todo.Pop(1);
frame->todo.Push(
Action::MakeValAction(Value::MakeBoolValue(exp->GetBoolLiteral())));
frame->todo.Push(Action::MakeValAction(
global_arena->New<BoolValue>(exp->GetBoolLiteral())));
break;
case ExpressionKind::PrimitiveOperatorExpression:
if (act->pos !=
@@ -676,21 +676,21 @@ void StepExp() {
break;
case ExpressionKind::IntTypeLiteral: {
CHECK(act->pos == 0);
const Value* v = Value::MakeIntType();
const Value* v = global_arena->New<IntType>();
frame->todo.Pop(1);
frame->todo.Push(Action::MakeValAction(v));
break;
}
case ExpressionKind::BoolTypeLiteral: {
CHECK(act->pos == 0);
const Value* v = Value::MakeBoolType();
const Value* v = global_arena->New<BoolType>();
frame->todo.Pop(1);
frame->todo.Push(Action::MakeValAction(v));
break;
}
case ExpressionKind::TypeTypeLiteral: {
CHECK(act->pos == 0);
const Value* v = Value::MakeTypeType();
const Value* v = global_arena->New<TypeType>();
frame->todo.Pop(1);
frame->todo.Push(Action::MakeValAction(v));
break;
@@ -709,8 +709,8 @@ void StepExp() {
} else if (act->pos == 2) {
// { { rt :: fn pt -> [] :: C, E, F} :: S, H}
// -> { fn pt -> rt :: {C, E, F} :: S, H}
const Value* v =
Value::MakeFunctionType({}, act->results[0], act->results[1]);
const Value* v = global_arena->New<FunctionType>(
std::vector<GenericBinding>(), act->results[0], act->results[1]);
frame->todo.Pop(1);
frame->todo.Push(Action::MakeValAction(v));
}
@@ -718,7 +718,7 @@ void StepExp() {
}
case ExpressionKind::ContinuationTypeLiteral: {
CHECK(act->pos == 0);
const Value* v = Value::MakeContinuationType();
const Value* v = global_arena->New<ContinuationType>();
frame->todo.Pop(1);
frame->todo.Push(Action::MakeValAction(v));
break;
@@ -736,7 +736,7 @@ void StepPattern() {
switch (pattern->Tag()) {
case Pattern::Kind::AutoPattern: {
CHECK(act->pos == 0);
const Value* v = Value::MakeAutoType();
const Value* v = global_arena->New<AutoType>();
frame->todo.Pop(1);
frame->todo.Push(Action::MakeValAction(v));
break;
@@ -747,8 +747,8 @@ void StepPattern() {
frame->todo.Push(Action::MakePatternAction(binding.Type()));
act->pos++;
} else {
auto v =
Value::MakeBindingPlaceholderValue(binding.Name(), act->results[0]);
auto v = global_arena->New<BindingPlaceholderValue>(binding.Name(),
act->results[0]);
frame->todo.Pop(1);
frame->todo.Push(Action::MakeValAction(v));
}
@@ -759,7 +759,8 @@ void StepPattern() {
if (act->pos == 0) {
if (tuple.Fields().empty()) {
frame->todo.Pop(1);
frame->todo.Push(Action::MakeValAction(Value::MakeTupleValue({})));
frame->todo.Push(Action::MakeValAction(
global_arena->New<TupleValue>(std::vector<TupleElement>())));
} else {
const Pattern* p1 = tuple.Fields()[0].pattern;
frame->todo.Push(Action::MakePatternAction(p1));
@@ -779,7 +780,8 @@ void StepPattern() {
elements.push_back(
{.name = tuple.Fields()[i].name, .value = act->results[i]});
}
const Value* tuple_value = Value::MakeTupleValue(std::move(elements));
const Value* tuple_value =
global_arena->New<TupleValue>(std::move(elements));
frame->todo.Pop(1);
frame->todo.Push(Action::MakeValAction(tuple_value));
}
@@ -796,10 +798,12 @@ void StepPattern() {
act->pos++;
} else {
CHECK(act->pos == 2);
const auto& choice_type = act->results[0]->GetChoiceType();
const auto& choice_type = cast<ChoiceType>(*act->results[0]);
frame->todo.Pop(1);
frame->todo.Push(Action::MakeValAction(Value::MakeAlternativeValue(
alternative.AlternativeName(), choice_type.name, act->results[1])));
frame->todo.Push(
Action::MakeValAction(global_arena->New<AlternativeValue>(
alternative.AlternativeName(), choice_type.Name(),
act->results[1])));
}
break;
}
@@ -917,7 +921,7 @@ void StepStmt() {
// -> { { e :: (while ([]) s) :: C, E, F} :: S, H}
frame->todo.Push(Action::MakeExpressionAction(stmt->GetWhile().cond));
act->pos++;
} else if (act->results[0]->GetBoolValue()) {
} else if (cast<BoolValue>(*act->results[0]).Val()) {
// { {true :: (while ([]) s) :: C, E, F} :: S, H}
// -> { { s :: (while (e) s) :: C, E, F } :: S, H}
frame->todo.Top()->pos = 0;
@@ -1042,7 +1046,7 @@ void StepStmt() {
// -> { { e :: (if ([]) then_stmt else else_stmt) :: C, E, F} :: S, H}
frame->todo.Push(Action::MakeExpressionAction(stmt->GetIf().cond));
act->pos++;
} else if (act->results[0]->GetBoolValue()) {
} else if (cast<BoolValue>(*act->results[0]).Val()) {
// { {true :: if ([]) then_stmt else else_stmt :: C, E, F} ::
// S, H}
// -> { { then_stmt :: C, E, F } :: S, H}
@@ -1099,8 +1103,9 @@ void StepStmt() {
todo.Push(Action::MakeStatementAction(stmt->GetContinuation().body));
Frame* continuation_frame =
global_arena->New<Frame>("__continuation", scopes, todo);
Address continuation_address = state->heap.AllocateValue(
Value::MakeContinuationValue({continuation_frame}));
Address continuation_address =
state->heap.AllocateValue(global_arena->New<ContinuationValue>(
std::vector<Frame*>({continuation_frame})));
// Store the continuation's address in the frame.
continuation_frame->continuation = continuation_address;
// Bind the continuation object to the continuation variable
@@ -1127,7 +1132,7 @@ void StepStmt() {
frame->todo.Push(ignore_result);
// Push the continuation onto the current stack.
const std::vector<Frame*>& continuation_vector =
act->results[0]->GetContinuationValue().stack;
cast<ContinuationValue>(*act->results[0]).Stack();
for (auto frame_iter = continuation_vector.rbegin();
frame_iter != continuation_vector.rend(); ++frame_iter) {
state->stack.Push(*frame_iter);
@@ -1144,7 +1149,8 @@ void StepStmt() {
} while (paused.back()->continuation == std::nullopt);
// Update the continuation with the paused stack.
state->heap.Write(*paused.back()->continuation,
Value::MakeContinuationValue(paused), stmt->line_num);
global_arena->New<ContinuationValue>(paused),
stmt->line_num);
break;
}
}
@@ -1209,7 +1215,7 @@ auto InterpProgram(std::list<Declaration>* fs) -> int {
}
}
const Value* v = state->stack.Top()->todo.Top()->GetValAction().val;
return v->GetIntValue();
return cast<IntValue>(*v).Val();
}
// Interpret an expression at compile-time.