diff options
Diffstat (limited to 'src')
| -rw-r--r-- | src/codegen/nasm.c | 56 | ||||
| -rw-r--r-- | src/lexer/lexer.c | 2 | ||||
| -rw-r--r-- | src/lexer/lexer.h | 1 | ||||
| -rw-r--r-- | src/parser/ast.c | 8 | ||||
| -rw-r--r-- | src/parser/ast.h | 4 | ||||
| -rw-r--r-- | src/parser/parser.c | 81 | ||||
| -rw-r--r-- | src/sema/sema.c | 7 |
7 files changed, 132 insertions, 27 deletions
diff --git a/src/codegen/nasm.c b/src/codegen/nasm.c index 6df22a4..8bcdda3 100644 --- a/src/codegen/nasm.c +++ b/src/codegen/nasm.c @@ -908,6 +908,12 @@ static void emit_call(struct Emitter* emitter, struct CallStatement* call) static void emit_statement(struct Emitter* emitter, struct Statement* statement); +static void emit_block(struct Emitter* emitter, struct Statement* body, size_t count) +{ + for (size_t i = 0; i < count; i += 1) + emit_statement(emitter, &body[i]); +} + // ucomisd sets the flags like an unsigned compare, so float branches use the // unsigned jump family (ja/jae/jb/jbe) rather than the signed one static const char* float_jump_if_false(enum TokenType comparison) @@ -939,10 +945,13 @@ static void emit_float_operand(struct Emitter* emitter, const struct Expr* expr) static void emit_if(struct Emitter* emitter, struct IfStatement* branch) { bool is_float = value_is_float(emitter, branch->left) || value_is_float(emitter, branch->right); + bool has_else = branch->else_count > 0; uint32_t id = emitter->label_id; emitter->label_id += 1; + const char* target = has_else ? ".if_else_" : ".if_end_"; + if (is_float) { const char* jump = float_jump_if_false(branch->comparison.type); @@ -958,30 +967,34 @@ static void emit_if(struct Emitter* emitter, struct IfStatement* branch) emit_float_operand(emitter, branch->left); fprintf(emitter->out, ", "); emit_float_operand(emitter, branch->right); - fprintf(emitter->out, "\n\t%s .if_end_%u\n", jump, id); - - emit_statement(emitter, branch->body); - fprintf(emitter->out, ".if_end_%u:\n", id); - return; + fprintf(emitter->out, "\n\t%s %s%u\n", jump, target, id); } - - const char* jump = jump_if_false(branch->comparison.type); - if (jump == NULL - || branch->left->kind == EXPR_BINARY || branch->left->kind == EXPR_DEREF - || branch->right->kind == EXPR_BINARY || branch->right->kind == EXPR_DEREF) + else { - fprintf(emitter->out, "\t; TODO: unsupported if\n"); - return; + const char* jump = jump_if_false(branch->comparison.type); + if (jump == NULL + || branch->left->kind == EXPR_BINARY || branch->left->kind == EXPR_DEREF + || branch->right->kind == EXPR_BINARY || branch->right->kind == EXPR_DEREF) + { + fprintf(emitter->out, "\t; TODO: unsupported if\n"); + return; + } + + fprintf(emitter->out, "\tcmp "); + emit_operand(emitter, branch->left); + fprintf(emitter->out, ", "); + emit_operand(emitter, branch->right); + fprintf(emitter->out, "\n\t%s %s%u\n", jump, target, id); } - fprintf(emitter->out, "\tcmp "); - emit_operand(emitter, branch->left); - fprintf(emitter->out, ", "); - emit_operand(emitter, branch->right); - fprintf(emitter->out, "\n"); - fprintf(emitter->out, "\t%s .if_end_%u\n", jump, id); + emit_block(emitter, branch->body, branch->body_count); - emit_statement(emitter, branch->body); + if (has_else) + { + fprintf(emitter->out, "\tjmp .if_end_%u\n", id); + fprintf(emitter->out, ".if_else_%u:\n", id); + emit_block(emitter, branch->else_body, branch->else_count); + } fprintf(emitter->out, ".if_end_%u:\n", id); } @@ -1083,7 +1096,10 @@ static void collect_floats_statement(struct FloatTable* floats, struct Statement case STATEMENT_IF: collect_floats_expr(floats, statement->branch.left); collect_floats_expr(floats, statement->branch.right); - collect_floats_statement(floats, statement->branch.body); + for (size_t i = 0; i < statement->branch.body_count; i += 1) + collect_floats_statement(floats, &statement->branch.body[i]); + for (size_t i = 0; i < statement->branch.else_count; i += 1) + collect_floats_statement(floats, &statement->branch.else_body[i]); break; case STATEMENT_CALL: for (size_t i = 0; i < statement->call.arg_count; i += 1) diff --git a/src/lexer/lexer.c b/src/lexer/lexer.c index 01bbd41..e08a458 100644 --- a/src/lexer/lexer.c +++ b/src/lexer/lexer.c @@ -38,6 +38,7 @@ static enum TokenType identifier_type(const char* start, size_t length) { "struct", 6, TOKEN_STRUCT }, { "stack", 5, TOKEN_STACK }, { "if", 2, TOKEN_IF }, + { "else", 4, TOKEN_ELSE }, { "goto", 4, TOKEN_GOTO }, { "syscall", 7, TOKEN_SYSCALL }, { "byte", 4, TOKEN_BYTE }, @@ -265,6 +266,7 @@ const char* token_type_name(enum TokenType type) case TOKEN_STRUCT: return "struct"; case TOKEN_STACK: return "stack"; case TOKEN_IF: return "if"; + case TOKEN_ELSE: return "else"; case TOKEN_GOTO: return "goto"; case TOKEN_SYSCALL: return "syscall"; case TOKEN_BYTE: return "byte"; diff --git a/src/lexer/lexer.h b/src/lexer/lexer.h index 45e9b3d..2bc43e0 100644 --- a/src/lexer/lexer.h +++ b/src/lexer/lexer.h @@ -19,6 +19,7 @@ enum TokenType TOKEN_STRUCT, TOKEN_STACK, TOKEN_IF, + TOKEN_ELSE, TOKEN_GOTO, TOKEN_SYSCALL, TOKEN_BYTE, diff --git a/src/parser/ast.c b/src/parser/ast.c index a387533..9c78d1f 100644 --- a/src/parser/ast.c +++ b/src/parser/ast.c @@ -26,7 +26,7 @@ void free_expr(struct Expr* expr) free(expr); } -static void free_statement(struct Statement* statement) +void free_statement(struct Statement* statement) { switch (statement->kind) { @@ -36,8 +36,12 @@ static void free_statement(struct Statement* statement) case STATEMENT_IF: free_expr(statement->branch.left); free_expr(statement->branch.right); - free_statement(statement->branch.body); + for (size_t i = 0; i < statement->branch.body_count; i += 1) + free_statement(&statement->branch.body[i]); free(statement->branch.body); + for (size_t i = 0; i < statement->branch.else_count; i += 1) + free_statement(&statement->branch.else_body[i]); + free(statement->branch.else_body); break; case STATEMENT_CALL: for (size_t i = 0; i < statement->call.arg_count; i += 1) diff --git a/src/parser/ast.h b/src/parser/ast.h index c439747..75d0bdd 100644 --- a/src/parser/ast.h +++ b/src/parser/ast.h @@ -135,6 +135,9 @@ struct IfStatement struct Token comparison; struct Expr* right; struct Statement* body; + size_t body_count; + struct Statement* else_body; + size_t else_count; }; struct CallStatement @@ -230,3 +233,4 @@ void add_statement(struct ProcDecl* proc, struct Statement statement); void add_proc(struct Program* program, struct ProcDecl decl); void free_expr(struct Expr* expr); +void free_statement(struct Statement* statement); diff --git a/src/parser/parser.c b/src/parser/parser.c index a83c3c4..ebf07c7 100644 --- a/src/parser/parser.c +++ b/src/parser/parser.c @@ -215,7 +215,7 @@ static bool is_assign_op(enum TokenType type) static struct Expr* alloc_expr(enum ExprKind kind) { - struct Expr* expr = malloc(sizeof(struct Expr)); + struct Expr* expr = malloc(sizeof(*expr)); if (expr != NULL) expr->kind = kind; return expr; @@ -390,6 +390,66 @@ error: static bool parse_statement(struct Parser* parser, struct Statement* out); +// A branch body is either a braced block or a single bare statement, always +// returned as a list so codegen and freeing treat both the same way. +static bool parse_block(struct Parser* parser, struct Statement** out_body, size_t* out_count) +{ + if (!match_token(parser, TOKEN_LEFT_BRACE)) + { + struct Statement* body = malloc(sizeof(*body)); + if (body == NULL) + return false; + + if (!parse_statement(parser, body)) + { + free(body); + return false; + } + + *out_body = body; + *out_count = 1; + return true; + } + + struct Statement* body = NULL; + size_t count = 0; + size_t capacity = 0; + + while (!check(parser, TOKEN_RIGHT_BRACE)) + { + if (check(parser, TOKEN_EOF)) + { + error_at(parser, parser->current, "unterminated block"); + goto error; + } + + if (count == capacity) + { + size_t grown_capacity = capacity == 0 ? 4 : capacity * 2; + struct Statement* grown = realloc(body, grown_capacity * sizeof(*grown)); + if (grown == NULL) + goto error; + body = grown; + capacity = grown_capacity; + } + + if (!parse_statement(parser, &body[count])) + goto error; + count += 1; + } + advance_parser(parser); + + *out_body = body; + *out_count = count; + return true; + +error: + for (size_t i = 0; i < count; i += 1) + free_statement(&body[i]); + free(body); + return false; +} + static bool parse_if(struct Parser* parser, struct Statement* out) { struct Expr* left = parse_expression(parser); @@ -412,9 +472,21 @@ static bool parse_if(struct Parser* parser, struct Statement* out) return false; } - struct Statement* body = malloc(sizeof(struct Statement)); - if (!parse_statement(parser, body)) + struct Statement* body; + size_t body_count; + if (!parse_block(parser, &body, &body_count)) + { + free_expr(left); + free_expr(right); + return false; + } + + struct Statement* else_body = NULL; + size_t else_count = 0; + if (match_token(parser, TOKEN_ELSE) && !parse_block(parser, &else_body, &else_count)) { + for (size_t i = 0; i < body_count; i += 1) + free_statement(&body[i]); free(body); free_expr(left); free_expr(right); @@ -426,6 +498,9 @@ static bool parse_if(struct Parser* parser, struct Statement* out) out->branch.comparison = comparison; out->branch.right = right; out->branch.body = body; + out->branch.body_count = body_count; + out->branch.else_body = else_body; + out->branch.else_count = else_count; return true; } diff --git a/src/sema/sema.c b/src/sema/sema.c index 4b4bc26..6505077 100644 --- a/src/sema/sema.c +++ b/src/sema/sema.c @@ -32,7 +32,7 @@ static bool check_duplicate_names(struct Source source, struct Program* program) if (count == 0) return true; - struct Token* names = malloc(count * sizeof(struct Token)); + struct Token* names = malloc(count * sizeof(*names)); if (names == NULL) return true; size_t n = 0; @@ -449,7 +449,10 @@ static void check_statement(struct RefCheck* check, struct Statement* statement) case STATEMENT_IF: check_expr(check, statement->branch.left); check_expr(check, statement->branch.right); - check_statement(check, statement->branch.body); + for (size_t i = 0; i < statement->branch.body_count; i += 1) + check_statement(check, &statement->branch.body[i]); + for (size_t i = 0; i < statement->branch.else_count; i += 1) + check_statement(check, &statement->branch.else_body[i]); break; case STATEMENT_CALL: { |
