aboutsummaryrefslogtreecommitdiff
path: root/src
diff options
context:
space:
mode:
Diffstat (limited to 'src')
-rw-r--r--src/codegen/nasm.c56
-rw-r--r--src/lexer/lexer.c2
-rw-r--r--src/lexer/lexer.h1
-rw-r--r--src/parser/ast.c8
-rw-r--r--src/parser/ast.h4
-rw-r--r--src/parser/parser.c81
-rw-r--r--src/sema/sema.c7
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:
{