aboutsummaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorhachem <im@hachem.wtf>2026-09-10 05:52:08 +0200
committerhachem <im@hachem.wtf>2026-09-10 05:52:08 +0200
commitd50af2e474281a1e88d191f783c2792824698974 (patch)
tree571afdff3ac141bbfcee37bc61dfdb6ea8fda163
parentbf157cf09b8b0888c54a165b887309a4e801f608 (diff)
feat: conditional select
-rw-r--r--docs/language.md6
-rw-r--r--src/ast.c6
-rw-r--r--src/ast.h14
-rw-r--r--src/codegen.c96
-rw-r--r--src/parser.c37
-rw-r--r--src/sema.c7
-rw-r--r--tests/codegen_test.c37
-rw-r--r--tests/parser_test.c20
8 files changed, 223 insertions, 0 deletions
diff --git a/docs/language.md b/docs/language.md
index a4743b3..ecdbf19 100644
--- a/docs/language.md
+++ b/docs/language.md
@@ -171,6 +171,12 @@ while .countdown rcx > 0 // emits `.countdown:` … `jmp .countdown` … `.c
rcx -= 1
```
+A **conditional select** picks one of two register values without a branch: `dst = a if <cond> else b`. It lowers to `csel` on AArch64 (one instruction) and `cmov` on x86 (a default move plus a conditional move, arranged so `dst` may safely alias either source). Both sources must be registers.
+
+```hdass
+r3 = r1 if r1 > r2 else r2 // r3 = max(r1, r2), branchless
+```
+
## 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 b775857..b613abe 100644
--- a/src/ast.c
+++ b/src/ast.c
@@ -63,6 +63,12 @@ void free_statement(struct Statement* statement)
free_expr(statement->instruction.operands[i]);
free(statement->instruction.operands);
break;
+ case STATEMENT_SELECT:
+ free_expr(statement->select.if_value);
+ free_expr(statement->select.left);
+ free_expr(statement->select.right);
+ free_expr(statement->select.else_value);
+ break;
case STATEMENT_STACK:
free_expr(statement->stack.size);
break;
diff --git a/src/ast.h b/src/ast.h
index 75d8b12..1db8b99 100644
--- a/src/ast.h
+++ b/src/ast.h
@@ -119,6 +119,7 @@ enum StatementKind
STATEMENT_CALL,
STATEMENT_STACK,
STATEMENT_INSTRUCTION,
+ STATEMENT_SELECT,
};
struct AssignStatement
@@ -187,6 +188,18 @@ struct InstructionStatement
size_t operand_capacity;
};
+// a branchless conditional move: `target = if_value if left cmp right else
+// else_value`. Lowers to csel on AArch64 and cmov on x86.
+struct SelectStatement
+{
+ struct Token target;
+ struct Expr* if_value;
+ struct Expr* left;
+ struct Token comparison;
+ struct Expr* right;
+ struct Expr* else_value;
+};
+
struct Statement
{
enum StatementKind kind;
@@ -200,6 +213,7 @@ struct Statement
struct CallStatement call;
struct StackStatement stack;
struct InstructionStatement instruction;
+ struct SelectStatement select;
};
};
diff --git a/src/codegen.c b/src/codegen.c
index 271c5ae..c0e67be 100644
--- a/src/codegen.c
+++ b/src/codegen.c
@@ -1162,6 +1162,55 @@ static void emit_instruction(struct Emitter* emitter, struct InstructionStatemen
fprintf(emitter->out, "\n");
}
+// the x86 condition-move suffix for a comparison, true or inverted
+static const char* cmov_cc(enum TokenType comparison, bool when_true)
+{
+ switch (comparison)
+ {
+ case TOKEN_EQUAL_EQUAL: return when_true ? "e" : "ne";
+ case TOKEN_BANG_EQUAL: return when_true ? "ne" : "e";
+ case TOKEN_LESS: return when_true ? "l" : "ge";
+ case TOKEN_LESS_EQUAL: return when_true ? "le" : "g";
+ case TOKEN_GREATER: return when_true ? "g" : "le";
+ case TOKEN_GREATER_EQUAL: return when_true ? "ge" : "l";
+ default: return NULL;
+ }
+}
+
+// target = if_value if L cmp R else else_value, branchlessly via cmov. The
+// destination defaults to one source and conditionally takes the other, choosing
+// which to default so the target can alias either source safely.
+static void emit_select(struct Emitter* emitter, struct SelectStatement* select)
+{
+ FILE* out = emitter->out;
+ if (select->if_value->kind != EXPR_PRIMARY || select->else_value->kind != EXPR_PRIMARY
+ || cmov_cc(select->comparison.type, true) == NULL)
+ {
+ fprintf(out, "\t; TODO: unsupported select\n");
+ return;
+ }
+
+ struct Token dst = resolve_register(emitter, select->target);
+ struct Token a = resolve_register(emitter, select->if_value->primary.token);
+ struct Token b = resolve_register(emitter, select->else_value->primary.token);
+ bool dst_is_a = tokens_equal(dst, a);
+
+ if (!dst_is_a && !tokens_equal(dst, b))
+ fprintf(out, "\tmov %.*s, %.*s\n", (int)dst.length, dst.start, (int)b.length, b.start);
+
+ fprintf(out, "\tcmp ");
+ emit_operand(emitter, select->left);
+ fprintf(out, ", ");
+ emit_operand(emitter, select->right);
+ fprintf(out, "\n");
+
+ // if the destination already holds if_value, keep it when the condition is
+ // false and move else_value in otherwise; otherwise the normal true form
+ struct Token source = dst_is_a ? b : a;
+ fprintf(out, "\tcmov%s %.*s, %.*s\n", cmov_cc(select->comparison.type, !dst_is_a),
+ (int)dst.length, dst.start, (int)source.length, source.start);
+}
+
static void emit_statement(struct Emitter* emitter, struct Statement* statement)
{
FILE* out = emitter->out;
@@ -1170,6 +1219,9 @@ static void emit_statement(struct Emitter* emitter, struct Statement* statement)
case STATEMENT_ASSIGN:
emit_assign(emitter, &statement->assign);
break;
+ case STATEMENT_SELECT:
+ emit_select(emitter, &statement->select);
+ break;
case STATEMENT_LABEL:
fprintf(out, "%.*s:\n", (int)statement->label.name.length, statement->label.name.start);
break;
@@ -1641,6 +1693,47 @@ static void emit_a64_instruction(struct Emitter* emitter, struct InstructionStat
fprintf(emitter->out, "\n");
}
+// the AArch64 condition code for a true comparison (used by csel)
+static const char* a64_cond(enum TokenType comparison)
+{
+ switch (comparison)
+ {
+ case TOKEN_EQUAL_EQUAL: return "eq";
+ case TOKEN_BANG_EQUAL: return "ne";
+ case TOKEN_LESS: return "lt";
+ case TOKEN_LESS_EQUAL: return "le";
+ case TOKEN_GREATER: return "gt";
+ case TOKEN_GREATER_EQUAL: return "ge";
+ default: return NULL;
+ }
+}
+
+// target = if_value if L cmp R else else_value -> cmp; csel (one instruction,
+// no aliasing hazard)
+static void emit_a64_select(struct Emitter* emitter, struct SelectStatement* select)
+{
+ FILE* out = emitter->out;
+ const char* cond = a64_cond(select->comparison.type);
+ if (cond == NULL || select->if_value->kind != EXPR_PRIMARY
+ || select->else_value->kind != EXPR_PRIMARY)
+ {
+ fprintf(out, "\t; TODO: unsupported select\n");
+ return;
+ }
+
+ fprintf(out, "\tcmp ");
+ emit_a64_operand(emitter, select->left);
+ fprintf(out, ", ");
+ emit_a64_operand(emitter, select->right);
+ fprintf(out, "\n\tcsel ");
+ emit_a64_reg(emitter, select->target);
+ fprintf(out, ", ");
+ emit_a64_reg(emitter, select->if_value->primary.token);
+ fprintf(out, ", ");
+ emit_a64_reg(emitter, select->else_value->primary.token);
+ fprintf(out, ", %s\n", cond);
+}
+
static void emit_a64_statement(struct Emitter* emitter, struct Statement* statement)
{
FILE* out = emitter->out;
@@ -1649,6 +1742,9 @@ static void emit_a64_statement(struct Emitter* emitter, struct Statement* statem
case STATEMENT_ASSIGN:
emit_a64_assign(emitter, &statement->assign);
break;
+ case STATEMENT_SELECT:
+ emit_a64_select(emitter, &statement->select);
+ break;
case STATEMENT_LABEL:
fprintf(out, "%.*s:\n", (int)statement->label.name.length, statement->label.name.start);
break;
diff --git a/src/parser.c b/src/parser.c
index 5379114..b0e049c 100644
--- a/src/parser.c
+++ b/src/parser.c
@@ -721,6 +721,43 @@ static bool parse_statement(struct Parser* parser, struct Statement* out)
if (value == NULL)
return false;
+ // conditional select: target = if_value if <cond> else else_value. The `if`
+ // must share the assignment's line, so a plain assignment followed by an
+ // `if` statement on the next line stays two statements.
+ if (op.type == TOKEN_EQUAL && !deref
+ && check(parser, TOKEN_IF) && parser->current.line == name.line)
+ {
+ advance_parser(parser);
+
+ struct Expr* left;
+ struct Token comparison;
+ struct Expr* right;
+ if (!parse_condition(parser, &left, &comparison, &right))
+ {
+ free_expr(value);
+ return false;
+ }
+
+ struct Expr* else_value = NULL;
+ if (!consume(parser, TOKEN_ELSE, "expected 'else' in conditional select")
+ || (else_value = parse_expression(parser)) == NULL)
+ {
+ free_expr(value);
+ free_expr(left);
+ free_expr(right);
+ return false;
+ }
+
+ out->kind = STATEMENT_SELECT;
+ out->select.target = name;
+ out->select.if_value = value;
+ out->select.left = left;
+ out->select.comparison = comparison;
+ out->select.right = right;
+ out->select.else_value = else_value;
+ return true;
+ }
+
out->kind = STATEMENT_ASSIGN;
out->assign.target_deref = deref;
out->assign.store_size = store_size;
diff --git a/src/sema.c b/src/sema.c
index 4b9af4a..01fe644 100644
--- a/src/sema.c
+++ b/src/sema.c
@@ -504,6 +504,13 @@ static void check_statement(struct RefCheck* check, struct Statement* statement)
// symbols or registers hdass does not model, so leave them to the
// assembler rather than flagging them as undefined
break;
+ case STATEMENT_SELECT:
+ check_target(check, statement->select.target);
+ check_expr(check, statement->select.if_value);
+ check_expr(check, statement->select.left);
+ check_expr(check, statement->select.right);
+ check_expr(check, statement->select.else_value);
+ break;
}
}
diff --git a/tests/codegen_test.c b/tests/codegen_test.c
index 2b5fdb7..0a4a841 100644
--- a/tests/codegen_test.c
+++ b/tests/codegen_test.c
@@ -104,6 +104,41 @@ static void test_generate_fasm(struct TestContext* context)
free_program(&program);
}
+static void test_generate_select(struct TestContext* context)
+{
+ // destination distinct from both sources: default-move else, then cmov
+ struct Lexer lexer = create_lexer("proc main\n{\nrax = rbx if rcx < rdx else rsi\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, "mov rax, rsi") != NULL);
+ check(context, strstr(buffer, "cmp rcx, rdx") != NULL);
+ check(context, strstr(buffer, "cmovl rax, rbx") != NULL);
+ check(context, strstr(buffer, "; TODO") == NULL);
+
+ free_program(&program);
+}
+
+static void test_generate_select_aarch64(struct TestContext* context)
+{
+ struct Lexer lexer = create_lexer(
+ "[enable: logical_registers]\nproc main\n{\nr1 = r2 if r3 < r4 else r5\n}\n");
+ struct Program program;
+ check(context, parse_program(&lexer, &program));
+
+ char buffer[1024];
+ generate_aarch64_to_buffer(&program, buffer, sizeof(buffer));
+
+ check(context, strstr(buffer, "cmp x2, x3") != NULL);
+ check(context, strstr(buffer, "csel x0, x1, x4, lt") != NULL);
+ check(context, strstr(buffer, "; TODO") == NULL);
+
+ free_program(&program);
+}
+
static void test_generate_aarch64(struct TestContext* context)
{
// logical registers map r1 -> x0, r2 -> x1; three-operand arithmetic and svc
@@ -633,6 +668,8 @@ void run_codegen_tests(struct TestContext* context)
test_generate_consts_and_data(context);
test_generate_fasm(context);
test_generate_aarch64(context);
+ test_generate_select(context);
+ test_generate_select_aarch64(context);
test_generate_text(context);
test_generate_if(context);
test_generate_if_else(context);
diff --git a/tests/parser_test.c b/tests/parser_test.c
index 9db85ec..eecc9b9 100644
--- a/tests/parser_test.c
+++ b/tests/parser_test.c
@@ -246,6 +246,25 @@ static void test_parse_while(struct TestContext* context)
free_program(&program);
}
+static void test_parse_select(struct TestContext* context)
+{
+ struct Lexer lexer = create_lexer("proc main\n{\nr1 = r2 if r3 < r4 else r5\n}\n");
+ struct Program program;
+
+ check(context, parse_program(&lexer, &program));
+
+ struct Statement s = program.procs[0].body[0];
+ check(context, s.kind == STATEMENT_SELECT);
+ check(context, text_is(s.select.target, "r1"));
+ check(context, primary_is(s.select.if_value, "r2"));
+ check(context, primary_is(s.select.left, "r3"));
+ check(context, text_is(s.select.comparison, "<"));
+ check(context, primary_is(s.select.right, "r4"));
+ check(context, primary_is(s.select.else_value, "r5"));
+
+ free_program(&program);
+}
+
static void test_parse_while_named(struct TestContext* context)
{
struct Lexer lexer = create_lexer("proc main\n{\nwhile .drain rcx > 0\n{\nrcx -= 1\n}\n}\n");
@@ -435,6 +454,7 @@ void run_parser_tests(struct TestContext* context)
test_parse_else_if(context);
test_parse_while(context);
test_parse_while_named(context);
+ test_parse_select(context);
test_parse_call(context);
test_parse_stack(context);
test_parse_sized_deref(context);