#include "ccc.h" #include "codegen.h" #include "scope.h" #include "register.h" #include #include #include #define CGEN_PANIC(format, ...) {\ fprintf(\ stderr,\ "ccc: code gen error: " format "\n" __VA_OPT__(,)\ __VA_ARGS__);\ exit(1);\ } struct lval_def { const struct type* type; struct storage_location loc; }; static const struct storage_location RV_LOC = { .type = STO_REG, .reg = &RAX, }; static const struct storage_location MULDIV_LOC = RV_LOC; static const struct storage_location MULDIV_OVERFLOW_LOC = { .type = STO_REG, .reg = &RDX, }; #define RETURN_LABEL_FMT "%s@coda" #define FULL_REG_SZ 8 static struct scope* scope; static const struct fn_decl_node* active_fn; static integral_t branch_counter = 0; static integral_t loop_counter = 0; static void enter_scope( struct scope* child_scope, integral_t bp_offset ) { if (child_scope == NULL || child_scope->next_out != scope) CGEN_PANIC("enter_scope: scopes are misaligned"); scope = child_scope; scope->bp_offset = bp_offset; } static void exit_scope(struct scope* child_scope, bool save_bp_offset) { if (child_scope != scope || child_scope->next_out == NULL) CGEN_PANIC("exit_scope: scopes are misaligned"); scope = child_scope->next_out; if (save_bp_offset) scope->bp_offset = child_scope->bp_offset; } static struct lval_def allocate_register(const struct type* type) { return (struct lval_def) { .loc = { .type = STO_REG, .reg = &RAX, /* TODO: no real register coloring happening LOL */ }, .type = type, }; } static const struct data_type* get_effective_data_type( const struct type* type ) { switch (type->type) { case TP_DATA: return type->data.data_type; case TP_PTR: return &long_long_type; } CGEN_PANIC("unhandled type of type case"); } static struct lval_def allocate_stack( FILE* outfile, const struct type* type ) { integral_t type_sz = get_effective_data_type(type)->sz; fprintf(outfile, "\tsub rsp, %llu\n", type_sz); scope->bp_offset += type_sz; return (struct lval_def) { .loc = { .type = STO_STACK, .bp_offset = scope->bp_offset, }, .type = type, }; } static struct lval_def allocate_temporary( FILE* outfile, const struct type* type ) { return allocate_stack(outfile, type); } static void deallocate_temporary(FILE* outfile, const struct lval_def* tmp) { if (tmp->loc.type == STO_STACK) { integral_t type_sz = get_effective_data_type(tmp->type)->sz; fprintf(outfile, "\tadd rsp, %llu\n", type_sz); scope->bp_offset -= type_sz; } else if (tmp->loc.type == STO_REG) { /* TOOD: release the register back to the algo */ } } static void emit_storage_loc( FILE* outfile, const struct storage_location* loc, integral_t sz ) { switch (loc->type) { case STO_LABEL: fprintf(outfile, "%s", loc->label); break; case STO_FN: fprintf(outfile, "%s", loc->decl->name); break; case STO_REG: if (sz > 4) fprintf(outfile, "%s", loc->reg->qword); else if (sz > 2) fprintf(outfile, "%s", loc->reg->dword); else if (sz > 1) fprintf(outfile, "%s", loc->reg->word); else fprintf(outfile, "%s", loc->reg->byte); break; case STO_STACK: if (loc->bp_offset < 0) fprintf(outfile, "[rbp + %lld]", -loc->bp_offset); else if (loc->bp_offset > 0) fprintf(outfile, "[rbp - %lld]", loc->bp_offset); else fprintf(outfile, "[rbp]"); break; case STO_IMM: fprintf(outfile, "%llu", loc->value); break; case STO_UNRESOLVED: CGEN_PANIC("can't emit unresolved storage location"); } } static bool locs_equal( const struct storage_location* a, const struct storage_location* b ) { if (a->type != b->type) return false; switch (a->type) { case STO_IMM: return a->value == b->value; case STO_REG: return strcmp(a->reg->qword, b->reg->qword) == 0; case STO_STACK: return a->bp_offset == b->bp_offset; case STO_LABEL: return strcmp(a->label, b->label) == 0; case STO_FN: return a->decl == b->decl; case STO_UNRESOLVED: return false; } CGEN_PANIC("unhandled storage type case"); } static void emit_size_const(FILE* outfile, integral_t sz) { if (sz > 4) fprintf(outfile, "qword "); else if (sz > 2) fprintf(outfile, "dword "); else if (sz > 1) fprintf(outfile, "word "); else fprintf(outfile, "byte "); } static void emit_mov( FILE* outfile, const struct lval_def* dst, const struct storage_location* src ) { /* first optimization: if dst == src, emit nothing */ if (locs_equal(&dst->loc, src)) return; integral_t dst_sz = get_effective_data_type(dst->type)->sz; switch (dst->loc.type) { case STO_REG: if (src->type == STO_REG && dst_sz < 4) { fprintf(outfile, "\tmovzx "); emit_storage_loc(outfile, &dst->loc, FULL_REG_SZ); } else { fprintf(outfile, "\tmov "); emit_storage_loc(outfile, &dst->loc, dst_sz); } fprintf(outfile, ", "); emit_storage_loc(outfile, src, dst_sz); break; case STO_STACK: if (src->type == STO_STACK) { /* `mov mem, mem` is illegal in x86_64 */ struct lval_def tmp = allocate_register(dst->type); emit_mov(outfile, &tmp, src); emit_mov(outfile, dst, &tmp.loc); return; } fprintf(outfile, "\tmov "); if (src->type == STO_IMM) emit_size_const(outfile, dst_sz); emit_storage_loc(outfile, &dst->loc, dst_sz); fprintf(outfile, ", "); emit_storage_loc(outfile, src, dst_sz); break; case STO_LABEL: case STO_IMM: case STO_FN: case STO_UNRESOLVED: CGEN_PANIC("can't move value into storage type"); } fprintf(outfile, "\n"); } static void emit_cmp_zero(FILE* outfile, const struct lval_def* lval) { fprintf(outfile, "\tcmp "); integral_t type_sz = get_effective_data_type(lval->type)->sz; switch (lval->loc.type) { case STO_REG: emit_storage_loc(outfile, &lval->loc, type_sz); break; case STO_STACK: emit_size_const(outfile, type_sz); emit_storage_loc(outfile, &lval->loc, type_sz); break; case STO_LABEL: case STO_IMM: case STO_FN: case STO_UNRESOLVED: CGEN_PANIC("can't compare this storage type") } fprintf(outfile, ", 0\n"); } static void emit_expr( FILE* outfile, const struct expr_node* node, const struct lval_def* dst); static void emit_int_lit( FILE* outfile, const struct int_lit_node* node, const struct lval_def* dst ) { if (dst != NULL) emit_mov( outfile, dst, &(struct storage_location) { .type = STO_IMM, .value = node->val, }); } static void emit_float_lit( FILE* outfile, const struct float_lit_node* node, const struct lval_def* dst ) { if (dst != NULL) { CGEN_PANIC("float literals are not implemented"); } } static void emit_char_lit( FILE* outfile, const struct char_lit_node* node, const struct lval_def* dst ) { if (dst != NULL) { emit_mov( outfile, dst, &(struct storage_location) { .type = STO_IMM, .value = (integral_t) node->val }); } } static void emit_str_lit( FILE* outfile, const struct str_lit_node* node, const struct lval_def* dst ) { if (dst != NULL) { CGEN_PANIC("string literals are not implemented"); } } static void emit_var_ref( FILE* outfile, const struct var_ref_node* node, const struct lval_def* dst ) { if (dst != NULL) { emit_mov(outfile, dst, &node->def_ref->loc); } } static void emit_stmt(FILE* outfile, const struct stmt_node* node); static void emit_decl( FILE* outfile, const struct decl_node* node ) { struct lval_def var_dst = allocate_stack(outfile, node->def_ref->type); node->def_ref->loc = var_dst.loc; fprintf(outfile, "\t; '%s' lives in: ", node->def_ref->name); integral_t dst_sz = get_effective_data_type(var_dst.type)->sz; emit_storage_loc(outfile, &var_dst.loc, dst_sz); fprintf(outfile, "\n"); if (node->initial_value != NULL) emit_expr(outfile, node->initial_value, &var_dst); } static void emit_decl_list( FILE* outfile, const struct decl_list_node* node ) { for (const struct decl_node* cur = node->head; cur != NULL; cur = cur->next) emit_decl(outfile, cur); } static void emit_assignment( FILE* outfile, const struct assign_node* node, const struct lval_def* dst ) { struct lval_def lval_def; switch (node->lval->type) { case EXPR_VAR_REF: struct var_ref_node* var_ref = &node->lval->inner.var_ref; lval_def = (struct lval_def) { .type = var_ref->def_ref->type, .loc = var_ref->def_ref->loc, }; break; default: CGEN_PANIC("expression is not assignable"); } emit_expr(outfile, node->rval, &lval_def); if (dst != NULL) emit_mov(outfile, dst, &lval_def.loc); } static void emit_call( FILE* outfile, const struct call_node* node, const struct lval_def* dst ) { integral_t orig_bp_offset = scope->bp_offset; integral_t arg_bp_offset = orig_bp_offset; 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) { struct lval_def arg_dst = allocate_stack(outfile, arg_decl->def_ref->type); emit_expr(outfile, arg_eval->expr, &arg_dst); arg_decl = arg_decl->next; arg_eval = arg_eval->next; } unsigned char arg_regnum = 0; arg_decl = node->fn_ref->args; while (arg_decl != NULL) { const struct type* arg_type = arg_decl->def_ref->type; arg_bp_offset += get_effective_data_type(arg_type)->sz; struct lval_def arg_dst; /* TODO: if the convention register is used, * will need to spill the current value */ if (arg_regnum < CC_N_REGS) arg_dst = (struct lval_def) { .type = arg_type, .loc = (struct storage_location) { .type = STO_REG, .reg = CALLING_CONV[arg_regnum++], }, }; else arg_dst = allocate_stack(outfile, arg_type); emit_mov(outfile, &arg_dst, &(struct storage_location) { .type = STO_STACK, .bp_offset = arg_bp_offset, }); arg_decl = arg_decl->next; } fprintf(outfile, "\tcall %s\n", node->fn_ref->name); if (dst != NULL) { if (get_effective_data_type( &node->fn_ref->return_type) == &void_type) CGEN_PANIC("can't assign the result of a void function"); emit_mov(outfile, dst, &RV_LOC); } /* mass-pop our argument temporaries off the stack */ scope->bp_offset = orig_bp_offset; if (orig_bp_offset > 0) fprintf(outfile, "\tlea rsp, [rbp - %llu]\n", orig_bp_offset); else fprintf(outfile, "\tmov rsp, rbp\n"); } static void emit_unary( FILE* outfile, const struct unary_node* node, const struct lval_def* dst ) { emit_expr(outfile, node->expr, dst); if (dst == NULL) return; integral_t dst_sz = get_effective_data_type(dst->type)->sz; switch (node->op) { case UNARY_NEG: fprintf(outfile, "\tneg "); if (dst->loc.type == STO_STACK) emit_size_const(outfile, dst_sz); emit_storage_loc(outfile, &dst->loc, dst_sz); fprintf(outfile, "\n"); break; } } /* TODO: tighten up/enforce with regard to up-casting smaller operands */ static void emit_binary( FILE* outfile, const struct binary_node* node, const struct lval_def* dst ) { if (dst == NULL) { emit_expr(outfile, node->lhs, NULL); emit_expr(outfile, node->rhs, NULL); return; } /* LHS goes in RAX explicitly because imul and idiv are weird */ struct lval_def rhs_dst = allocate_temporary(outfile, dst->type); emit_expr(outfile, node->rhs, &rhs_dst); struct lval_def lhs_dst = (struct lval_def) { .loc = MULDIV_LOC, .type = dst->type, }; emit_expr(outfile, node->lhs, &lhs_dst); integral_t dst_sz = get_effective_data_type(dst->type)->sz; switch (node->op) { case BINARY_ADD: fprintf(outfile, "\tadd "); emit_storage_loc(outfile, &lhs_dst.loc, dst_sz); fprintf(outfile, ", "); break; case BINARY_SUB: fprintf(outfile, "\tsub "); emit_storage_loc(outfile, &lhs_dst.loc, dst_sz); fprintf(outfile, ", "); break; case BINARY_MUL: fprintf(outfile, "\timul "); if (rhs_dst.loc.type == STO_STACK) emit_size_const(outfile, dst_sz); break; case BINARY_DIV: /* nothing in the top half reg */ fprintf(outfile, "\txor "); emit_storage_loc(outfile, &MULDIV_OVERFLOW_LOC, FULL_REG_SZ); fprintf(outfile, ", "); emit_storage_loc(outfile, &MULDIV_OVERFLOW_LOC, FULL_REG_SZ); fprintf(outfile, "\n"); fprintf(outfile, "\tidiv "); if (rhs_dst.loc.type == STO_STACK) emit_size_const(outfile, dst_sz); break; } emit_storage_loc(outfile, &rhs_dst.loc, dst_sz); fprintf(outfile, "\n"); /* TODO: deal with RDX overflow shit for imul and idiv */ emit_mov(outfile, dst, &lhs_dst.loc); deallocate_temporary(outfile, &lhs_dst); deallocate_temporary(outfile, &rhs_dst); } static void emit_expr_list( FILE* outfile, const struct expr_list_node* node, const struct lval_def* dst ) { for (; node != NULL; node = node->next) emit_expr(outfile, node->expr, node->next == NULL ? dst : NULL); } static void emit_cast( FILE* outfile, const struct cast_node* node, const struct lval_def* dst ) { /* TODO: for anything but a reinterpret cast this is garbage */ emit_expr(outfile, node->expr, dst); } static void emit_expr( FILE* outfile, const struct expr_node* node, const struct lval_def* dst ) { switch (node->type) { case EXPR_INT_LIT: emit_int_lit(outfile, &node->inner.int_lit, dst); break; case EXPR_FLOAT_LIT: emit_float_lit(outfile, &node->inner.float_lit, dst); break; case EXPR_CHAR_LIT: emit_char_lit(outfile, &node->inner.char_lit, dst); break; case EXPR_STR_LIT: emit_str_lit(outfile, &node->inner.str_lit, dst); break; case EXPR_VAR_REF: emit_var_ref(outfile, &node->inner.var_ref, dst); break; case EXPR_ASSIGN: emit_assignment(outfile, &node->inner.assign, dst); break; case EXPR_CALL: emit_call(outfile, &node->inner.call, dst); break; case EXPR_UNARY: emit_unary(outfile, &node->inner.unary, dst); break; case EXPR_BINARY: emit_binary(outfile, &node->inner.binary, dst); break; case EXPR_PAREN: emit_expr_list(outfile, node->inner.paren.expr_list, dst); break; case EXPR_CAST: emit_cast(outfile, &node->inner.cast, dst); break; } } static void emit_return(FILE* outfile, const struct return_node* node) { if (active_fn == NULL) CGEN_PANIC("must be inside a function to return"); if (node->ret_val != NULL) emit_expr_list( outfile, node->ret_val, &(struct lval_def) { .loc = RV_LOC, .type = &active_fn->return_type, }); fprintf(outfile, "\tjmp " RETURN_LABEL_FMT "\n", active_fn->name); } static void emit_group_contents(FILE* outfile, const struct group_node* node) { const struct stmt_node* body_node = node->head; while (body_node != NULL) { emit_stmt(outfile, body_node); body_node = body_node->next; } } static void emit_group(FILE* outfile, const struct group_node* node) { enter_scope(node->scope, scope->bp_offset); emit_group_contents(outfile, node); /* don't reset sp because alloca needs to work */ scope->next_out->bp_offset = scope->bp_offset; exit_scope(node->scope, true); } static void emit_if(FILE* outfile, const struct if_node* node) { enter_scope(node->scope, scope->bp_offset); struct lval_def cond_result = allocate_temporary(outfile, node->cond->resolved_type); emit_expr_list(outfile, node->cond, &cond_result); emit_cmp_zero(outfile, &cond_result); integral_t branch_num = ++branch_counter; fprintf(outfile, "\tjz branch_false@%lld\n", branch_num); emit_stmt(outfile, node->true_branch); if (node->false_branch == NULL) { fprintf(outfile, "branch_false@%lld:\n", branch_num); } else { fprintf(outfile, "\tjmp branch_done@%lld\n", branch_num); fprintf(outfile, "branch_false@%lld:\n", branch_num); emit_stmt(outfile, node->false_branch); fprintf(outfile, "branch_done@%lld:\n", branch_num); } exit_scope(node->scope, true); } static void emit_loop_init(FILE* outfile, const struct loop_init_node* node) { switch (node->type) { case INIT_EXPR_LIST: emit_expr_list(outfile, node->expr_list, NULL); break; case INIT_DECL_LIST: emit_decl_list(outfile, node->decl_list); break; } } static void emit_loop(FILE* outfile, const struct loop_node* node) { enter_scope(node->scope, scope->bp_offset); if (node->init != NULL) emit_loop_init(outfile, node->init); integral_t loop_num = ++loop_counter; if (node->cond != NULL) { struct lval_def cond_dst = allocate_temporary(outfile, node->cond->resolved_type); fprintf(outfile, "loop_head@%lld:\n", loop_num); emit_expr_list(outfile, node->cond, &cond_dst); emit_cmp_zero(outfile, &cond_dst); fprintf(outfile, "\tjz loop_done@%lld\n", loop_num); } else { fprintf(outfile, "loop_head@%lld:\n", loop_num); } emit_stmt(outfile, node->body); if (node->incr != NULL) emit_expr_list(outfile, node->incr, NULL); fprintf(outfile, "\tjmp loop_head@%lld\n", loop_num); fprintf(outfile, "loop_done@%lld:\n", loop_num); exit_scope(node->scope, true); } static void emit_stmt(FILE* outfile, const struct stmt_node* node) { switch (node->type) { case STMT_EMPTY: break; case STMT_DECL_LIST: emit_decl_list(outfile, &node->inner.decl_list); break; case STMT_RETURN: emit_return(outfile, &node->inner.return_); break; case STMT_EXPR_LIST: emit_expr_list(outfile, &node->inner.expr_list, NULL); break; case STMT_GROUP: emit_group(outfile, &node->inner.group); break; case STMT_IF: emit_if(outfile, &node->inner.if_); break; case STMT_LOOP: emit_loop(outfile, &node->inner.loop); break; } } static void emit_fn_decl(FILE* outfile, const struct fn_decl_node* node) { enter_scope(node->scope, 0); if (active_fn != NULL) CGEN_PANIC( "can't define function %s inside function %s", node->name, active_fn->name); active_fn = node; /* TODO: we need to account for the base pointer moving in var locs */ fprintf(outfile, "%s:\n", node->name); fprintf(outfile, "\tpush rbp\n"); fprintf(outfile, "\tmov rbp, rsp\n"); scope->bp_offset = 0; long long spilled_bp_ofs = -16; // return address + old bp unsigned char arg_regnum = 0; struct arg_decl_node* arg_decl = node->args; for (; arg_decl != NULL; arg_decl = arg_decl->next) { struct var_def* arg_def = arg_decl->def_ref; struct lval_def arg_dst = allocate_stack(outfile, arg_def->type); arg_def->loc = arg_dst.loc; struct storage_location arg_src; if (arg_regnum < CC_N_REGS) { arg_src = (struct storage_location) { .type = STO_REG, .reg = CALLING_CONV[arg_regnum++] }; } else { arg_src = (struct storage_location) { .type = STO_STACK, .bp_offset = spilled_bp_ofs, }; integral_t arg_sz = get_effective_data_type(arg_def->type)->sz; spilled_bp_ofs -= arg_sz; } emit_mov(outfile, &arg_dst, &arg_src); } enter_scope(node->body.scope, scope->bp_offset); emit_group_contents(outfile, &node->body); exit_scope(node->body.scope, true); fprintf(outfile, RETURN_LABEL_FMT ":\n", node->name); fprintf(outfile, "\tmov rsp, rbp\n"); fprintf(outfile, "\tpop rbp\n"); fprintf(outfile, "\tret\n"); active_fn = NULL; exit_scope(node->scope, false); } static void emit_root_node(FILE* outfile, const struct root_node* node) { switch (node->type) { case ROOT_FN_DECL: emit_fn_decl(outfile, &node->inner.fn_decl); break; } } void emit_code(struct ast* ast, const char* path) { FILE* outfile = fopen(path, "w"); if (outfile == NULL) CCC_PANIC; scope = ast->root_scope; fprintf(outfile, "section .text\n"); /* output all function declarations in the root scope as globals */ const struct root_node* node = ast->root_node; for (; node != NULL; node = node->next) { if (node->type != ROOT_FN_DECL) continue; fprintf(outfile, "global %s\n", node->inner.fn_decl.name); } fprintf(outfile, "\n"); /* actual code body */ node = ast->root_node; while (node != NULL) { emit_root_node(outfile, node); if (node->next != NULL) fprintf(outfile, "\n"); node = node->next; } fclose(outfile); }