diff options
Diffstat (limited to 'type_checker.c')
| -rw-r--r-- | type_checker.c | 215 |
1 files changed, 210 insertions, 5 deletions
diff --git a/type_checker.c b/type_checker.c index deb3345..538f51d 100644 --- a/type_checker.c +++ b/type_checker.c @@ -1,4 +1,4 @@ -#include "ast.h" +#include "type_checker.h" #include "scope.h" #include <stdio.h> #include <stdlib.h> @@ -13,15 +13,211 @@ static struct scope* scope; +static void type_check_expr(struct expr_node* node); +static void type_check_lval(struct lval_node* node); +static void type_check_stmt(struct stmt_node* node); static void type_check_group(struct group_node* node); -static void type_check_expr(struct expr_node* node) {} +static void assert_cast_compatible( + const struct type_def* lval_type, + const struct type_def* rval_type +) { + /* TODO: impl */ + /* TODO: we should also insert cast nodes eventually */ + if (false) { + TYPE_PANIC( + "cannot assign value of type '%s' to '%s'", + lval_type->key.name, + rval_type->key.name); + } +} + +static const struct type_def* resolve_type_ref(struct type_ref_node* node) { + /* TODO: support pointers and such */ + if (node->def_ref == NULL) TYPE_PANIC("use of unresolved type"); + return node->def_ref; +} + +/* TODO: we likely want some shortcut to access char, int, double, long, etc. */ +static const struct type_def* resolve_int_lit(struct int_lit_node* node) { + const struct type_def* resolved_type; + scope_get_type(scope, &resolved_type, &(struct type_key) { + .name = "int", + .how_long = 1, + .marked_signed = false, + .marked_unsigned = false, + }); + + if (resolved_type == NULL) + TYPE_PANIC( + "no integral type exists, " + "ensure install_default_types was called."); + + return resolved_type; +} + +static const struct type_def* resolve_float_lit(struct float_lit_node* node) { + const struct type_def* resolved_type; + scope_get_type(scope, &resolved_type, &(struct type_key) { + .name = "float", + .how_long = 0, + .marked_signed = false, + .marked_unsigned = false, + }); + + if (resolved_type == NULL) + TYPE_PANIC( + "no floating point type exists, " + "ensure install_default_types was called."); + + return resolved_type; +} + +static const struct type_def* resolve_char_lit(struct char_lit_node* node) { + const struct type_def* resolved_type; + scope_get_type(scope, &resolved_type, &(struct type_key) { + .name = "char", + .how_long = 0, + .marked_signed = false, + .marked_unsigned = false, + }); + + if (resolved_type == NULL) + TYPE_PANIC( + "no character type exists, " + "ensure install_default_types was called."); + + return resolved_type; +} -static void type_check_var_decl(struct var_decl_node* node) {} +static const struct type_def* resolve_str_lit(struct str_lit_node* node) { + /* TODO: pointer types */ + const struct type_def* resolved_type; + scope_get_type(scope, &resolved_type, &(struct type_key) { + .name = "int", + .how_long = 1, + .marked_signed = false, + .marked_unsigned = true, + }); -static void type_check_return(struct return_node* node) {} + if (resolved_type == NULL) + TYPE_PANIC( + "no string type exists, " + "ensure install_default_types was called."); -static void type_check_if(struct if_node* node) {} + return resolved_type; +} + +static const struct type_def* resolve_var_ref(struct var_ref_node* node) { + if (node->def_ref->resolved_type == NULL) + TYPE_PANIC("variable '%s' has unresolved type.", node->def_ref->name); + return node->def_ref->resolved_type; +} + +static void type_check_var_decl(struct var_decl_node* node) { + node->def_ref->resolved_type = resolve_type_ref(&node->type); +} + +static const struct type_def* resolve_assign(struct assign_node* node) { + type_check_lval(&node->lval); + type_check_expr(node->rval); + const struct type_def* lval_type = node->lval.resolved_type; + const struct type_def* rval_type = node->rval->resolved_type; + + assert_cast_compatible(lval_type, rval_type); + return lval_type; +} + +static const struct type_def* resolve_call(struct call_node* node) { + struct args_decl_node* arg_decl = node->called_fn_ref->args; + struct args_eval_node* arg_eval = node->args; + while (arg_decl != NULL && arg_eval != NULL) { + type_check_var_decl(arg_decl->decl); + const struct type_def* decl_type = + arg_decl->decl->def_ref->resolved_type; + type_check_expr(arg_eval->expr); + const struct type_def* 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 resolve_type_ref(&node->called_fn_ref->return_type); +} + +static const struct type_def* resolve_unary(struct unary_node* node) { + type_check_expr(node->expr); + /* TODO: delegate by individual operation to prohibit shit like -"yes" */ + return node->expr->resolved_type; +} + +static const struct type_def* resolve_binary(struct binary_node* node) { + /* TODO: math rules and such lol */ + type_check_expr(node->lhs); + type_check_expr(node->rhs); + assert_cast_compatible(node->lhs->resolved_type, node->rhs->resolved_type); + return node->lhs->resolved_type; +} + +static void type_check_lval(struct lval_node* node) { + switch (node->type) { + case LVAL_VAR_DECL: + type_check_var_decl(&node->inner.var_decl); + node->resolved_type = node->inner.var_decl.def_ref->resolved_type; + break; + case LVAL_VAR_REF: + node->resolved_type = resolve_var_ref(&node->inner.var_ref); + break; + } +} + +static void type_check_expr(struct expr_node* node) { + switch (node->type) { + case EXPR_INT_LIT: + node->resolved_type = resolve_int_lit(&node->inner.int_lit); + break; + case EXPR_FLOAT_LIT: + node->resolved_type = resolve_float_lit(&node->inner.float_lit); + break; + case EXPR_CHAR_LIT: + node->resolved_type = resolve_char_lit(&node->inner.char_lit); + break; + case EXPR_STR_LIT: + node->resolved_type = resolve_str_lit(&node->inner.str_lit); + break; + case EXPR_VAR_REF: + node->resolved_type = resolve_var_ref(&node->inner.var_ref); + break; + case EXPR_ASSIGN: + node->resolved_type = resolve_assign(&node->inner.assign); + break; + case EXPR_CALL: + node->resolved_type = resolve_call(&node->inner.call); + break; + case EXPR_UNARY: + node->resolved_type = resolve_unary(&node->inner.unary); + break; + case EXPR_BINARY: + node->resolved_type = resolve_binary(&node->inner.binary); + break; + } +} + +static void type_check_return(struct return_node* node) { + if (node->ret_val != NULL) type_check_expr(node->ret_val); +} + +static void type_check_if(struct if_node* node) { + printf("tc if\n"); + type_check_expr(node->cond); + type_check_stmt(node->true_branch); + if (node->false_branch != NULL) type_check_stmt(node->false_branch); +} static void type_check_stmt(struct stmt_node* node) { switch (node->type) { @@ -59,7 +255,16 @@ static void type_check_group(struct group_node* node) { static void type_check_fn_decl(struct fn_decl_node* node) { if (node->scope->next_out != scope) TYPE_PANIC("scopes are borked"); scope = node->scope; + + node->return_type.def_ref = resolve_type_ref(&node->return_type); + + for (struct args_decl_node* arg_decl = node->args; + arg_decl != NULL; + arg_decl = arg_decl->next) + type_check_var_decl(arg_decl->decl); + type_check_group(&node->body); + scope = scope->next_out; } |
