diff options
| author | hachem <im@hachem.wtf> | 2026-09-08 12:16:51 +0200 |
|---|---|---|
| committer | hachem <im@hachem.wtf> | 2026-09-08 12:16:51 +0200 |
| commit | 60c0e084b159801a45570b56f152f0f5e9193d72 (patch) | |
| tree | ad6c72d6782ab5d9edff336e1e852ca7b995bcdc | |
| parent | 9d47486aac9cd1a4422b78ee273b7a27bf7b733a (diff) | |
feat: add while loops
| -rw-r--r-- | docs/language.md | 15 | ||||
| -rw-r--r-- | src/ast.c | 7 | ||||
| -rw-r--r-- | src/ast.h | 11 | ||||
| -rw-r--r-- | src/lexer.c | 2 | ||||
| -rw-r--r-- | src/lexer.h | 1 | ||||
| -rw-r--r-- | src/nasm.c | 105 | ||||
| -rw-r--r-- | src/parser.c | 46 | ||||
| -rw-r--r-- | src/sema.c | 6 | ||||
| -rw-r--r-- | tests/codegen_test.c | 22 | ||||
| -rw-r--r-- | tests/parser_test.c | 21 |
10 files changed, 202 insertions, 34 deletions
diff --git a/docs/language.md b/docs/language.md index 3ccdab8..8556c8b 100644 --- a/docs/language.md +++ b/docs/language.md @@ -120,12 +120,14 @@ loop: // label goto loop if rcx != 0 // == != < <= > >= ; guards the next statement or a { block } goto loop +while rcx != 0 // same condition; repeats the next statement or { block } + rcx -= 1 syscall print_number(r12) // call; args go into the callee's parameter registers stack buf[Point.size] // stack buffer (size is any constant); buf is its base address ``` -## Branching (`if` / `else`) +## Control flow (`if` / `else` / `while`) `if <expr> <cmp> <expr>` guards either the single next statement or a `{ }` block, and an optional `else` takes its own statement or block. `else if` chains because the `else` body is itself a statement. Comparisons are `==` `!=` `<` `<=` `>` `>=`; a float compare needs an `xmm` register on the left (see [Floating point](#floating-point)). @@ -141,6 +143,17 @@ else rdi = -1 ``` +`while <expr> <cmp> <expr>` runs its statement or `{ }` block for as long as the condition holds, testing it before each pass. It is the same condition as `if`, and desugars to a label, the test, the body, and a jump back — the `loop:`/`goto` you would write by hand. Use `goto` to break out early. + +```hdass +rbx = 0 +while rcx > 0 +{ + rbx += rcx + rcx -= 1 +} +``` + ## Dereference (`^`) `^reg` is the memory at the address in `reg` — NASM's `[reg]`. On the left of `=` it stores there. The store width comes from the value operand, so a sized sub-register picks the size: @@ -46,6 +46,13 @@ void free_statement(struct Statement* statement) free_statement(&statement->branch.else_body[i]); free(statement->branch.else_body); break; + case STATEMENT_WHILE: + free_expr(statement->loop.left); + free_expr(statement->loop.right); + for (size_t i = 0; i < statement->loop.body_count; i += 1) + free_statement(&statement->loop.body[i]); + free(statement->loop.body); + break; case STATEMENT_CALL: for (size_t i = 0; i < statement->call.arg_count; i += 1) free_expr(statement->call.args[i]); @@ -115,6 +115,7 @@ enum StatementKind STATEMENT_GOTO, STATEMENT_SYSCALL, STATEMENT_IF, + STATEMENT_WHILE, STATEMENT_CALL, STATEMENT_STACK, }; @@ -149,6 +150,15 @@ struct IfStatement size_t else_count; }; +struct WhileStatement +{ + struct Expr* left; + struct Token comparison; + struct Expr* right; + struct Statement* body; + size_t body_count; +}; + struct CallStatement { struct Token name; @@ -172,6 +182,7 @@ struct Statement struct LabelStatement label; struct GotoStatement jump; struct IfStatement branch; + struct WhileStatement loop; struct CallStatement call; struct StackStatement stack; }; diff --git a/src/lexer.c b/src/lexer.c index 0d26f9a..b7fec3c 100644 --- a/src/lexer.c +++ b/src/lexer.c @@ -39,6 +39,7 @@ static enum TokenType identifier_type(const char* start, size_t length) { "stack", 5, TOKEN_STACK }, { "if", 2, TOKEN_IF }, { "else", 4, TOKEN_ELSE }, + { "while", 5, TOKEN_WHILE }, { "goto", 4, TOKEN_GOTO }, { "syscall", 7, TOKEN_SYSCALL }, { "byte", 4, TOKEN_BYTE }, @@ -268,6 +269,7 @@ const char* token_type_name(enum TokenType type) case TOKEN_STACK: return "stack"; case TOKEN_IF: return "if"; case TOKEN_ELSE: return "else"; + case TOKEN_WHILE: return "while"; case TOKEN_GOTO: return "goto"; case TOKEN_SYSCALL: return "syscall"; case TOKEN_BYTE: return "byte"; diff --git a/src/lexer.h b/src/lexer.h index 73de02f..bdb00a8 100644 --- a/src/lexer.h +++ b/src/lexer.h @@ -20,6 +20,7 @@ enum TokenType TOKEN_STACK, TOKEN_IF, TOKEN_ELSE, + TOKEN_WHILE, TOKEN_GOTO, TOKEN_SYSCALL, TOKEN_BYTE, @@ -999,51 +999,64 @@ static void emit_float_operand(struct Emitter* emitter, const struct Expr* expr) fprintf(emitter->out, "%.*s", (int)token.length, token.start); } -static void emit_if(struct Emitter* emitter, struct IfStatement* branch) +// Emits the comparison for `left cmp right` and a jump to `target` taken when +// the condition is false, so the code that follows runs when it is true. Both +// if and while build on this. Returns false (after a TODO note) for a form that +// isn't supported yet. +static bool emit_branch_test(struct Emitter* emitter, struct Expr* left, + struct Token comparison, struct Expr* right, const char* target) { - 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_"; + bool is_float = value_is_float(emitter, left) || value_is_float(emitter, right); if (is_float) { - const char* jump = float_jump_if_false(branch->comparison.type); - bool left_reg = branch->left->kind == EXPR_PRIMARY - && is_float_register(resolve_register(emitter, branch->left->primary.token)); - if (jump == NULL || !left_reg || !value_is_float(emitter, branch->right)) + const char* jump = float_jump_if_false(comparison.type); + bool left_reg = left->kind == EXPR_PRIMARY + && is_float_register(resolve_register(emitter, left->primary.token)); + if (jump == NULL || !left_reg || !value_is_float(emitter, right)) { - fprintf(emitter->out, "\t; TODO: unsupported if\n"); - return; + fprintf(emitter->out, "\t; TODO: unsupported condition\n"); + return false; } fprintf(emitter->out, "\tucomisd "); - emit_float_operand(emitter, branch->left); + emit_float_operand(emitter, left); fprintf(emitter->out, ", "); - emit_float_operand(emitter, branch->right); - fprintf(emitter->out, "\n\t%s %s%u\n", jump, target, id); + emit_float_operand(emitter, right); + fprintf(emitter->out, "\n\t%s %s\n", jump, target); + return true; } - else - { - 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); + const char* jump = jump_if_false(comparison.type); + if (jump == NULL + || left->kind == EXPR_BINARY || left->kind == EXPR_DEREF + || right->kind == EXPR_BINARY || right->kind == EXPR_DEREF) + { + fprintf(emitter->out, "\t; TODO: unsupported condition\n"); + return false; } + fprintf(emitter->out, "\tcmp "); + emit_operand(emitter, left); + fprintf(emitter->out, ", "); + emit_operand(emitter, right); + fprintf(emitter->out, "\n\t%s %s\n", jump, target); + return true; +} + +static void emit_if(struct Emitter* emitter, struct IfStatement* branch) +{ + bool has_else = branch->else_count > 0; + + uint32_t id = emitter->label_id; + emitter->label_id += 1; + + char target[32]; + snprintf(target, sizeof(target), ".if_%s_%u", has_else ? "else" : "end", id); + + if (!emit_branch_test(emitter, branch->left, branch->comparison, branch->right, target)) + return; + emit_block(emitter, branch->body, branch->body_count); if (has_else) @@ -1056,6 +1069,25 @@ static void emit_if(struct Emitter* emitter, struct IfStatement* branch) fprintf(emitter->out, ".if_end_%u:\n", id); } +static void emit_while(struct Emitter* emitter, struct WhileStatement* loop) +{ + uint32_t id = emitter->label_id; + emitter->label_id += 1; + + char target[32]; + snprintf(target, sizeof(target), ".while_end_%u", id); + + fprintf(emitter->out, ".while_%u:\n", id); + + if (!emit_branch_test(emitter, loop->left, loop->comparison, loop->right, target)) + return; + + emit_block(emitter, loop->body, loop->body_count); + + fprintf(emitter->out, "\tjmp .while_%u\n", id); + fprintf(emitter->out, ".while_end_%u:\n", id); +} + static void emit_statement(struct Emitter* emitter, struct Statement* statement) { FILE* out = emitter->out; @@ -1076,6 +1108,9 @@ static void emit_statement(struct Emitter* emitter, struct Statement* statement) case STATEMENT_IF: emit_if(emitter, &statement->branch); break; + case STATEMENT_WHILE: + emit_while(emitter, &statement->loop); + break; case STATEMENT_CALL: emit_call(emitter, &statement->call); break; @@ -1161,6 +1196,12 @@ static void collect_floats_statement(struct FloatTable* floats, struct Statement for (size_t i = 0; i < statement->branch.else_count; i += 1) collect_floats_statement(floats, &statement->branch.else_body[i]); break; + case STATEMENT_WHILE: + collect_floats_expr(floats, statement->loop.left); + collect_floats_expr(floats, statement->loop.right); + for (size_t i = 0; i < statement->loop.body_count; i += 1) + collect_floats_statement(floats, &statement->loop.body[i]); + break; case STATEMENT_CALL: for (size_t i = 0; i < statement->call.arg_count; i += 1) collect_floats_expr(floats, statement->call.args[i]); diff --git a/src/parser.c b/src/parser.c index 4c01c86..ddaee72 100644 --- a/src/parser.c +++ b/src/parser.c @@ -481,7 +481,8 @@ error: return false; } -static bool parse_if(struct Parser* parser, struct Statement* out) +static bool parse_condition(struct Parser* parser, struct Expr** out_left, + struct Token* out_comparison, struct Expr** out_right) { struct Expr* left = parse_expression(parser); if (left == NULL) @@ -503,6 +504,20 @@ static bool parse_if(struct Parser* parser, struct Statement* out) return false; } + *out_left = left; + *out_comparison = comparison; + *out_right = right; + return true; +} + +static bool parse_if(struct Parser* parser, struct Statement* out) +{ + struct Expr* left; + struct Token comparison; + struct Expr* right; + if (!parse_condition(parser, &left, &comparison, &right)) + return false; + struct Statement* body; size_t body_count; if (!parse_block(parser, &body, &body_count)) @@ -535,11 +550,40 @@ static bool parse_if(struct Parser* parser, struct Statement* out) return true; } +static bool parse_while(struct Parser* parser, struct Statement* out) +{ + struct Expr* left; + struct Token comparison; + struct Expr* right; + if (!parse_condition(parser, &left, &comparison, &right)) + return false; + + struct Statement* body; + size_t body_count; + if (!parse_block(parser, &body, &body_count)) + { + free_expr(left); + free_expr(right); + return false; + } + + out->kind = STATEMENT_WHILE; + out->loop.left = left; + out->loop.comparison = comparison; + out->loop.right = right; + out->loop.body = body; + out->loop.body_count = body_count; + return true; +} + static bool parse_statement(struct Parser* parser, struct Statement* out) { if (match_token(parser, TOKEN_IF)) return parse_if(parser, out); + if (match_token(parser, TOKEN_WHILE)) + return parse_while(parser, out); + if (match_token(parser, TOKEN_SYSCALL)) { out->kind = STATEMENT_SYSCALL; @@ -470,6 +470,12 @@ static void check_statement(struct RefCheck* check, struct Statement* statement) for (size_t i = 0; i < statement->branch.else_count; i += 1) check_statement(check, &statement->branch.else_body[i]); break; + case STATEMENT_WHILE: + check_expr(check, statement->loop.left); + check_expr(check, statement->loop.right); + for (size_t i = 0; i < statement->loop.body_count; i += 1) + check_statement(check, &statement->loop.body[i]); + break; case STATEMENT_CALL: { struct CallStatement* call = &statement->call; diff --git a/tests/codegen_test.c b/tests/codegen_test.c index 1da711c..c92dfc2 100644 --- a/tests/codegen_test.c +++ b/tests/codegen_test.c @@ -100,6 +100,27 @@ static void test_generate_if_else(struct TestContext* context) free_program(&program); } +static void test_generate_while(struct TestContext* context) +{ + struct Lexer lexer = create_lexer( + "proc main\n{\nwhile rcx > 0\n{\nrbx += rcx\nrcx -= 1\n}\n}\n"); + struct Program program; + check(context, parse_program(&lexer, &program)); + + char buffer[1024]; + generate_to_buffer(&program, buffer, sizeof(buffer)); + + check(context, strstr(buffer, ".while_0:") != NULL); + check(context, strstr(buffer, "cmp rcx, 0") != NULL); + check(context, strstr(buffer, "jle .while_end_0") != NULL); + check(context, strstr(buffer, "add rbx, rcx") != NULL); + check(context, strstr(buffer, "jmp .while_0") != NULL); + check(context, strstr(buffer, ".while_end_0:") != NULL); + check(context, strstr(buffer, "; TODO") == NULL); + + free_program(&program); +} + static void test_generate_negative(struct TestContext* context) { struct Lexer lexer = create_lexer( @@ -485,6 +506,7 @@ void run_codegen_tests(struct TestContext* context) test_generate_text(context); test_generate_if(context); test_generate_if_else(context); + test_generate_while(context); test_generate_negative(context); test_generate_call(context); test_generate_param_substitution(context); diff --git a/tests/parser_test.c b/tests/parser_test.c index b51093e..88e32a1 100644 --- a/tests/parser_test.c +++ b/tests/parser_test.c @@ -225,6 +225,26 @@ static void test_parse_else_if(struct TestContext* context) free_program(&program); } +static void test_parse_while(struct TestContext* context) +{ + struct Lexer lexer = create_lexer( + "proc main\n{\nwhile rcx > 0\n{\nrbx += rcx\nrcx -= 1\n}\n}\n"); + struct Program program; + + check(context, parse_program(&lexer, &program)); + check(context, program.procs[0].body_count == 1); + + struct Statement loop = program.procs[0].body[0]; + check(context, loop.kind == STATEMENT_WHILE); + check(context, primary_is(loop.loop.left, "rcx")); + check(context, text_is(loop.loop.comparison, ">")); + check(context, primary_is(loop.loop.right, "0")); + check(context, loop.loop.body_count == 2); + check(context, text_is(loop.loop.body[1].assign.target, "rcx")); + + free_program(&program); +} + static void test_parse_call(struct TestContext* context) { struct Lexer lexer = create_lexer("proc main\n{\nprint_number(r12)\nf(a, b)\n}\n"); @@ -396,6 +416,7 @@ void run_parser_tests(struct TestContext* context) test_parse_if(context); test_parse_if_block(context); test_parse_else_if(context); + test_parse_while(context); test_parse_call(context); test_parse_stack(context); test_parse_sized_deref(context); |
