From bfb8020c1a64f505537163caa465bd7f06e43f18 Mon Sep 17 00:00:00 2001 From: Carson Fleming Date: Fri, 17 Jul 2026 17:54:49 -0700 Subject: functional type checker --- ast.h | 2 + main.c | 3 + parser.c | 7 +- scope.c | 16 +---- type_checker.c | 215 +++++++++++++++++++++++++++++++++++++++++++++++++++++++-- type_checker.h | 8 +++ 6 files changed, 226 insertions(+), 25 deletions(-) create mode 100644 type_checker.h diff --git a/ast.h b/ast.h index aac846c..89fd0f8 100644 --- a/ast.h +++ b/ast.h @@ -45,6 +45,8 @@ struct lval_node { struct var_ref_node var_ref; struct var_decl_node var_decl; } inner; + + const struct type_def* resolved_type; }; struct assign_node { diff --git a/main.c b/main.c index 90d7460..9a3a725 100644 --- a/main.c +++ b/main.c @@ -1,5 +1,6 @@ #include "lexer.h" #include "parser.h" +#include "type_checker.h" #include "codegen.h" #include #include @@ -41,6 +42,8 @@ void test_parser(int argc, char** argv) { for (int i = 1; i < argc; i++) { struct ast ast; parse(argv[i], &ast); + type_check(&ast); + unsigned int fn_sz = strlen(argv[i]); char asm_file[fn_sz + 1]; strcpy(asm_file, argv[i]); diff --git a/parser.c b/parser.c index a213ee4..0d403a6 100644 --- a/parser.c +++ b/parser.c @@ -138,7 +138,6 @@ static void parse_expr_assign(struct expr_node* p_node) { expect(TK_ASSIGN); parse_expr(p_node->inner.assign.rval); - p_node->resolved_type = p_node->inner.assign.rval->resolved_type; } static void parse_args_eval(struct args_eval_node** pp_arg) { @@ -158,12 +157,11 @@ static void parse_expr_call(struct expr_node* p_node) { switch (p_node->type) { case EXPR_VAR_REF: struct var_def* var_def = p_node->inner.var_ref.def_ref; - /* TODO: I would like to include functions in the resolved_type rather than checking storage */ + /* TODO: I would like to include functions in the type model rather than checking storage */ if (var_def->loc.type != STO_FN) PARSER_PANIC("called object is not a function"); p_node->inner.call.called_fn_ref = var_def->loc.decl; - p_node->resolved_type = var_def->loc.decl->return_type.def_ref; break; default: PARSER_PANIC("expression is not callable"); @@ -252,8 +250,6 @@ static void parse_expr(struct expr_node* p_node) { case TK_IDENT: p_node->type = EXPR_VAR_REF; parse_var_ref(&p_node->inner.var_ref); - p_node->resolved_type = - p_node->inner.var_ref.def_ref->resolved_type; break; default: PARSER_PANIC("expected expression"); @@ -284,7 +280,6 @@ static void parse_var_decl(struct var_decl_node* p_node) { p_node->def_ref = scope_define_var(scope, (struct var_def) { .name = tok.data.ident, .loc.type = STO_UNRESOLVED, - .resolved_type = p_node->type.def_ref, }); if (p_node->def_ref == NULL) PARSER_PANIC("redefinition of '%s'", tok.data.ident); diff --git a/scope.c b/scope.c index 2fdadbd..db5c240 100644 --- a/scope.c +++ b/scope.c @@ -173,18 +173,6 @@ void scope_install_default_types(struct scope* scope) { .sz = 0, }); - scope_define_type(scope, (struct type_def) { - .key = { - .name = strdup("bool"), - .how_long = 0, - .marked_signed = false, - .marked_unsigned = false, - }, - .sz = 1, - .is_signed = false, - .is_floating = false, - }); - scope_define_type(scope, (struct type_def) { .key = { .name = strdup("char"), @@ -223,8 +211,8 @@ void scope_install_default_types(struct scope* scope) { scope_define_type(scope, (struct type_def) { .key = { - .name = strdup("long"), - .how_long = 0, + .name = strdup("int"), + .how_long = 1, .marked_signed = false, .marked_unsigned = false, }, 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 #include @@ -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; } diff --git a/type_checker.h b/type_checker.h new file mode 100644 index 0000000..1276927 --- /dev/null +++ b/type_checker.h @@ -0,0 +1,8 @@ +#ifndef TYPE_CHECKER_H +#define TYPE_CHECKER_H + +#include "ast.h" + +void type_check(struct ast* ast); + +#endif -- cgit v1.2.3