summaryrefslogtreecommitdiff
path: root/type_checker.c
diff options
context:
space:
mode:
Diffstat (limited to 'type_checker.c')
-rw-r--r--type_checker.c91
1 files changed, 59 insertions, 32 deletions
diff --git a/type_checker.c b/type_checker.c
index ccd816f..8b4763c 100644
--- a/type_checker.c
+++ b/type_checker.c
@@ -95,14 +95,23 @@ static struct type* resolve_var_ref(struct var_ref_node* node) {
return copy_type(node->def_ref->type);
}
-static void type_check_var_decl(struct var_decl_node* node) {
- node->def_ref->type = &node->type;
+static void type_check_decl(
+ const struct type* decl_type,
+ struct decl_node* node
+) {
+ node->def_ref->type = decl_type;
+
if (node->initial_value != NULL) {
type_check_expr(node->initial_value);
- assert_cast_compatible(&node->type, node->initial_value->resolved_type);
+ assert_cast_compatible(decl_type, node->initial_value->resolved_type);
}
}
+static void type_check_decl_list(struct decl_list_node* node) {
+ for (struct decl_node* cur = node->head; cur != NULL; cur = cur->next)
+ type_check_decl(&node->type, cur);
+}
+
static struct type* resolve_assign(struct assign_node* node) {
switch (node->lval->type) {
case EXPR_VAR_REF:
@@ -120,29 +129,35 @@ static struct type* resolve_assign(struct assign_node* node) {
return copy_type(lval_type);
}
+static void type_check_expr_list(struct expr_list_node* node) {
+ const struct type* last_item_type;
+ for (; node != NULL; node = node->next) {
+ type_check_expr(node->expr);
+ last_item_type = node->expr->resolved_type;
+ }
+ node->resolved_type = copy_type(last_item_type);
+}
+
static struct type* resolve_call(struct call_node* node) {
- struct decl_list_node* arg_decl = node->called_fn_ref->args;
+ type_check_expr_list(node->args);
+ /* TODO: enforce that the function has been type checked prior to this */
+ struct arg_decl_node* arg_decl = node->fn_ref->args;
struct expr_list_node* arg_eval = node->args;
while (arg_decl != NULL && arg_eval != NULL) {
- if (arg_decl->decl->initial_value != NULL)
- TYPE_PANIC("C does not support default arguments");
-
- type_check_var_decl(arg_decl->decl);
- const struct type* decl_type =
- arg_decl->decl->def_ref->type;
- type_check_expr(arg_eval->expr);
+ const struct type* decl_type = &arg_decl->type;
const struct type* eval_type = arg_eval->expr->resolved_type;
assert_cast_compatible(decl_type, eval_type);
arg_decl = arg_decl->next;
arg_eval = arg_eval->next;
}
- if (arg_decl != NULL || arg_eval != NULL)
- TYPE_PANIC(
- "mismatched argument count in call to '%s'",
- node->called_fn_ref->name);
- return copy_type(&node->called_fn_ref->return_type);
+ if (arg_decl != NULL)
+ TYPE_PANIC("missing arguments to function '%s'", node->fn_ref->name);
+ if (arg_eval != NULL)
+ TYPE_PANIC("too many arguments to function '%s'", node->fn_ref->name);
+
+ return copy_type(&node->fn_ref->return_type);
}
static struct type* resolve_unary(struct unary_node* node) {
@@ -159,6 +174,11 @@ static struct type* resolve_binary(struct binary_node* node) {
return copy_type(node->lhs->resolved_type);
}
+static struct type* resolve_paren(struct paren_node* node) {
+ type_check_expr_list(node->expr_list);
+ return copy_type(node->expr_list->resolved_type);
+}
+
static void type_check_expr(struct expr_node* node) {
switch (node->type) {
case EXPR_INT_LIT:
@@ -188,34 +208,41 @@ static void type_check_expr(struct expr_node* node) {
case EXPR_BINARY:
node->resolved_type = resolve_binary(&node->inner.binary);
break;
+ case EXPR_PAREN:
+ node->resolved_type = resolve_paren(&node->inner.paren);
+ break;
}
}
static void type_check_return(struct return_node* node) {
- if (node->ret_val != NULL) type_check_expr(node->ret_val);
+ const struct type* return_type = &node->fn_ref->return_type;
+ if (node->ret_val != NULL) {
+ type_check_expr_list(node->ret_val);
+ assert_cast_compatible(return_type, node->ret_val->resolved_type);
+ } else if (return_type->type != TP_DATA
+ || return_type->data.data_type != &void_type)
+ TYPE_PANIC(
+ "void function '%s' should not return a value",
+ node->fn_ref->name);
}
static void type_check_if(struct if_node* node) {
enter_scope(node->scope);
- type_check_expr(node->cond);
+ type_check_expr_list(node->cond);
type_check_stmt(node->true_branch);
if (node->false_branch != NULL) type_check_stmt(node->false_branch);
exit_scope(node->scope);
}
-static void type_check_expr_list(struct expr_list_node* node) {
- for (; node != NULL; node = node->next) type_check_expr(node->expr);
-}
-
static void type_check_loop_init(struct loop_init_node* node) {
switch (node->type) {
case INIT_EXPR_LIST:
type_check_expr_list(node->expr_list);
break;
- case INIT_DECL:
- type_check_var_decl(node->decl);
+ case INIT_DECL_LIST:
+ type_check_decl_list(node->decl_list);
break;
}
}
@@ -224,8 +251,8 @@ static void type_check_loop(struct loop_node* node) {
enter_scope(node->scope);
if (node->init != NULL) type_check_loop_init(node->init);
- if (node->cond != NULL) type_check_expr(node->cond);
- if (node->incr != NULL) type_check_expr(node->incr);
+ if (node->cond != NULL) type_check_expr_list(node->cond);
+ if (node->incr != NULL) type_check_expr_list(node->incr);
type_check_stmt(node->body);
exit_scope(node->scope);
@@ -235,11 +262,11 @@ static void type_check_stmt(struct stmt_node* node) {
switch (node->type) {
case STMT_EMPTY:
break;
- case STMT_EXPR:
- type_check_expr(&node->inner.expr);
+ case STMT_EXPR_LIST:
+ type_check_expr_list(&node->inner.expr_list);
break;
- case STMT_VAR_DECL:
- type_check_var_decl(&node->inner.var_decl);
+ case STMT_DECL_LIST:
+ type_check_decl_list(&node->inner.decl_list);
break;
case STMT_RETURN:
type_check_return(&node->inner.return_);
@@ -271,10 +298,10 @@ static void type_check_group(struct group_node* node) {
static void type_check_fn_decl(struct fn_decl_node* node) {
enter_scope(node->scope);
- for (struct decl_list_node* arg_decl = node->args;
+ for (struct arg_decl_node* arg_decl = node->args;
arg_decl != NULL;
arg_decl = arg_decl->next)
- type_check_var_decl(arg_decl->decl);
+ arg_decl->def_ref->type = &arg_decl->type;
type_check_group(&node->body);