aboutsummaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--docs/language.md15
-rw-r--r--src/ast.c7
-rw-r--r--src/ast.h11
-rw-r--r--src/lexer.c2
-rw-r--r--src/lexer.h1
-rw-r--r--src/nasm.c105
-rw-r--r--src/parser.c46
-rw-r--r--src/sema.c6
-rw-r--r--tests/codegen_test.c22
-rw-r--r--tests/parser_test.c21
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:
diff --git a/src/ast.c b/src/ast.c
index 28bb98e..da51b05 100644
--- a/src/ast.c
+++ b/src/ast.c
@@ -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]);
diff --git a/src/ast.h b/src/ast.h
index 074c9b2..f5eea17 100644
--- a/src/ast.h
+++ b/src/ast.h
@@ -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,
diff --git a/src/nasm.c b/src/nasm.c
index 15f5f01..7a7d1e6 100644
--- a/src/nasm.c
+++ b/src/nasm.c
@@ -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;
diff --git a/src/sema.c b/src/sema.c
index 12f2689..274bd53 100644
--- a/src/sema.c
+++ b/src/sema.c
@@ -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);