diff options
| author | hachem <im@hachem.wtf> | 2026-08-31 23:14:26 +0200 |
|---|---|---|
| committer | hachem <im@hachem.wtf> | 2026-08-31 23:14:26 +0200 |
| commit | 9dc06d01c2224ee1dbeb11ebee7c8afecfac8b05 (patch) | |
| tree | 110fe30b5b1f228568087a8fe7b0f7e0ca43fbb1 | |
| parent | 5b06b0e4b8effaacd23308d6a6b26c525e7d3af4 (diff) | |
feat: add support for floats through xmm registers
| -rw-r--r-- | README.md | 5 | ||||
| -rw-r--r-- | docs/language.md | 36 | ||||
| -rw-r--r-- | examples/circle.hdass | 17 | ||||
| -rw-r--r-- | examples/fibonacci.hdass | 2 | ||||
| -rw-r--r-- | examples/hello_world.hdass | 2 | ||||
| -rw-r--r-- | examples/load.hdass | 6 | ||||
| -rw-r--r-- | examples/logical.hdass | 10 | ||||
| -rw-r--r-- | examples/loop_sum.hdass | 4 | ||||
| -rw-r--r-- | examples/mandelbrot.hdass | 144 | ||||
| -rw-r--r-- | examples/records.hdass | 2 | ||||
| -rwxr-xr-x | m | bin | 0 -> 9176 bytes | |||
| -rw-r--r-- | m.asm | 119 | ||||
| -rw-r--r-- | m.o | bin | 0 -> 1472 bytes | |||
| -rwxr-xr-x | scripts/run.sh | 28 | ||||
| -rwxr-xr-x | scripts/test_examples.sh | 1 | ||||
| -rw-r--r-- | src/codegen/nasm.c | 414 | ||||
| -rw-r--r-- | src/lexer/lexer.c | 15 | ||||
| -rw-r--r-- | src/lexer/lexer.h | 1 | ||||
| -rw-r--r-- | src/parser/ast.c | 3 | ||||
| -rw-r--r-- | src/parser/ast.h | 2 | ||||
| -rw-r--r-- | src/parser/parser.c | 10 | ||||
| -rw-r--r-- | src/sema/sema.c | 43 | ||||
| -rw-r--r-- | tests/codegen_test.c | 101 | ||||
| -rw-r--r-- | tests/lexer_test.c | 9 | ||||
| -rw-r--r-- | tests/parser_test.c | 2 | ||||
| -rw-r--r-- | tests/sema_test.c | 9 |
26 files changed, 950 insertions, 35 deletions
@@ -79,6 +79,11 @@ Or, from the host, run the whole suite (unit tests plus the example tests) in on ./scripts/docker_test.sh ``` +To build, assemble, link and run a single program in the container: +```bash +./scripts/run.sh examples/circle.hdass +``` + When you're finished, stop and remove the container: ```bash docker compose down diff --git a/docs/language.md b/docs/language.md index 5e75c54..d67cbf0 100644 --- a/docs/language.md +++ b/docs/language.md @@ -90,7 +90,7 @@ rsi += Point.y // add rsi, 8 rax = Point.size // mov rax, 17 ``` -A struct is layout only — it allocates nothing. Pair it with a `stack` buffer and pointer arithmetic (see [examples/records.hdass](../examples/records.hdass)). +A struct is layout only — it allocates nothing. Pair it with a `stack` buffer sized by `Name.size` and pointer arithmetic (see [examples/records.hdass](../examples/records.hdass)). ## Procedures @@ -122,7 +122,7 @@ if rcx != 0 // == != < <= > >= ; runs the next statement only goto loop syscall print_number(r12) // call; args go into the callee's parameter registers -stack buffer[32] // stack buffer; buffer is its base address +stack buf[Point.size] // stack buffer (size is any constant); buf is its base address ``` ## Dereference (`^`) @@ -188,6 +188,38 @@ r1.64 // rax Arch `r8`–`r15` share the `rN` spelling, so with the extension on a bare `r8` is the *logical* register (which is arch `r9`). Reach arch `r8`–`r15` through logical `r7`–`r14`. Architecture names like `rax` and `rsi` still work everywhere. +## Floating point + +Floating-point values live in the SSE registers `xmm0`–`xmm15` (double precision). Float literals like `3.14` are placed in `.data` and loaded for you. + +```hdass +xmm0 = 3.5 // movsd from a .data slot +xmm0 *= xmm1 // += -= *= /= -> addsd subsd mulsd divsd +``` + +An `=` between a float register and a general-purpose register converts: + +```hdass +xmm0 = rax // int -> float (cvtsi2sd) +rbx = xmm0 // float -> int, truncating (cvttsd2si) +``` + +`^` loads and stores floats too, so float state can live in memory (a `stack` buffer or struct): + +```hdass +^rsi = xmm0 // movsd [rsi], xmm0 +xmm1 = ^rsi // movsd xmm1, [rsi] +``` + +`if` compares floats too, when the left side is an `xmm` register (`ucomisd`): + +```hdass +if xmm0 > 4.0 + goto escaped +``` + +See [examples/mandelbrot.hdass](../examples/mandelbrot.hdass) for a float program. Not yet supported: mixing floats and ints in one expression, and printing floats. + ## Building a program ```bash diff --git a/examples/circle.hdass b/examples/circle.hdass new file mode 100644 index 0000000..8a4edaa --- /dev/null +++ b/examples/circle.hdass @@ -0,0 +1,17 @@ +[entry: main] + +// Floating-point math on the SSE registers. Computes the area of a circle +// (pi * r * r) with r = 3, then exits with the truncated result (28). +const SYS_EXIT = 60 + +proc main +{ + xmm0 = 3.0 // radius + xmm0 *= xmm0 // r * r = 9.0 + xmm1 = 3.14159 // pi + xmm0 *= xmm1 // area ~= 28.27 + + rdi = xmm0 // truncate to int -> 28 + rax = SYS_EXIT + syscall +} diff --git a/examples/fibonacci.hdass b/examples/fibonacci.hdass index 95a59c8..9436580 100644 --- a/examples/fibonacci.hdass +++ b/examples/fibonacci.hdass @@ -43,7 +43,7 @@ proc main { r12 = 0 r13 = 1 - r15 = 10 // syscall clobbers rcx/r11, so keep the counter in r15 + r15 = 10 // syscall clobbers rcx/r11, so keep the counter in r15 loop: print_number(r12) diff --git a/examples/hello_world.hdass b/examples/hello_world.hdass index da6b2b5..149c856 100644 --- a/examples/hello_world.hdass +++ b/examples/hello_world.hdass @@ -14,7 +14,7 @@ proc main rsi = message rdx = message.len syscall - + rax = SYS_EXIT rdi = 0 syscall diff --git a/examples/load.hdass b/examples/load.hdass index 6e68ffc..8d29358 100644 --- a/examples/load.hdass +++ b/examples/load.hdass @@ -8,9 +8,9 @@ proc main stack cell[8] rbx = 7 - rsi = cell // address of the cell - ^rsi = rbx // store - rdi = ^rsi // load it back + rsi = cell // address of the cell + ^rsi = rbx // store + rdi = ^rsi // load it back rax = SYS_EXIT syscall diff --git a/examples/logical.hdass b/examples/logical.hdass index 5e82b42..5fb27fc 100644 --- a/examples/logical.hdass +++ b/examples/logical.hdass @@ -8,12 +8,12 @@ const SYS_EXIT = 60 proc main { - r1 = 4 // rax - r2 = 3 // rbx - r1 += r2 // 7 - r1 *= r2 // 21 + r1 = 4 // rax + r2 = 3 // rbx + r1 += r2 // 7 + r1 *= r2 // 21 - r6 = r1 // rdi = 21 (the exit status) + r6 = r1 // rdi = 21 (the exit status) r1 = SYS_EXIT syscall } diff --git a/examples/loop_sum.hdass b/examples/loop_sum.hdass index 5976914..2427b17 100644 --- a/examples/loop_sum.hdass +++ b/examples/loop_sum.hdass @@ -5,8 +5,8 @@ const SYS_EXIT = 60 proc main { - rbx = 0 // running total - rcx = 5 // counter + rbx = 0 // running total + rcx = 5 // counter loop: rbx += rcx diff --git a/examples/mandelbrot.hdass b/examples/mandelbrot.hdass new file mode 100644 index 0000000..9054c15 --- /dev/null +++ b/examples/mandelbrot.hdass @@ -0,0 +1,144 @@ +[entry: main] +[enable: logical_registers] + +// Renders the Mandelbrot set as ASCII art. The evolving z value is kept in a +// Complex struct on the stack (accessed through its field offsets), while the +// per-pixel constant c stays in xmm2/xmm3. Uses double-precision floats and the +// logical_registers extension (r1..r14 name the general-purpose registers). + +const SYS_WRITE = 1 +const SYS_EXIT = 60 +const STDOUT = 1 + +const W = 80 +const H = 30 +const MAX_ITER = 32 + +struct Complex +{ + re: qword + im: qword +} + +data palette = " .:-=+*#%@" + +proc main +{ + stack z[Complex.size] + stack row[W + 1] + + r11 = 0 // py + +row_loop: + // ci = py * 0.08 - 1.2 + xmm3 = r11 + xmm7 = 0.08 + xmm3 *= xmm7 + xmm7 = 1.2 + xmm3 -= xmm7 + + r12 = 0 // px + +col_loop: + // cr = px * 0.04375 - 2.5 + xmm2 = r12 + xmm7 = 0.04375 + xmm2 *= xmm7 + xmm7 = 2.5 + xmm2 -= xmm7 + + // z = 0 + 0i + xmm0 = 0.0 + r5 = z + r5 += Complex.re + ^r5 = xmm0 // z.re = 0 + r5 = z + r5 += Complex.im + ^r5 = xmm0 // z.im = 0 + + r2 = 0 // iteration count + +iter_loop: + // load z into xmm0 (zr) and xmm1 (zi) + r5 = z + r5 += Complex.re + xmm0 = ^r5 + r5 = z + r5 += Complex.im + xmm1 = ^r5 + + // zr2 = zr*zr, zi2 = zi*zi + xmm4 = xmm0 + xmm4 *= xmm0 + xmm5 = xmm1 + xmm5 *= xmm1 + + // escape when zr2 + zi2 > 4.0 + xmm6 = xmm4 + xmm6 += xmm5 + if xmm6 > 4.0 + goto plot + + // new_zi = 2*zr*zi + ci (uses the old zr and zi) + xmm6 = xmm0 + xmm6 *= xmm1 + xmm6 += xmm6 + xmm6 += xmm3 + xmm7 = xmm6 // stash new_zi + + // new_zr = zr2 - zi2 + cr + xmm4 -= xmm5 + xmm4 += xmm2 + + // store the new z back into the struct + r5 = z + r5 += Complex.re + ^r5 = xmm4 // z.re = new_zr + r5 = z + r5 += Complex.im + ^r5 = xmm7 // z.im = new_zi + + r2 += 1 + if r2 < MAX_ITER + goto iter_loop + +plot: + // index = iter * 9 / MAX_ITER, then char = palette[index] + r1 = r2 + r1 *= 9 + r3 = MAX_ITER + r1 /= r3 + + r5 = palette + r5 += r1 + r4 = ^byte r5 + + r6 = row + r6 += r12 + ^byte r6 = r4 + + r12 += 1 + if r12 < W + goto col_loop + + // terminate the row with a newline and write it + r6 = row + r6 += W + r1 = 10 + ^byte r6 = r1 + + r1 = SYS_WRITE + r6 = STDOUT + r5 = row + r4 = W + r4 += 1 + syscall + + r11 += 1 + if r11 < H + goto row_loop + + r1 = SYS_EXIT + r6 = 0 + syscall +} diff --git a/examples/records.hdass b/examples/records.hdass index 9800839..f8553d6 100644 --- a/examples/records.hdass +++ b/examples/records.hdass @@ -19,7 +19,7 @@ const SYS_EXIT = 60 proc main { - stack pair[16] // Pair.size + stack pair[Pair.size] rsi = pair rbx = 40 Binary files differ@@ -0,0 +1,119 @@ +bits 64 + +%define SYS_WRITE (1) +%define SYS_EXIT (60) +%define STDOUT (1) +%define SCALE (1024) +%define FOUR (4096) +%define MAX_ITER (32) +%define W (80) +%define H (30) +%define RSPAN (3584) +%define RMIN (2560) +%define ISPAN (2458) +%define IMIN (1229) + +section .data +palette: db ` .:-=+*#%@` +.len equ $ - palette + +section .text +global main + +main: + push rbp + mov rbp, rsp + sub rsp, 96 + mov r12, 0 +row_loop: + mov rax, r12 + imul rax, ISPAN + mov rcx, H + cqo + idiv rcx + sub rax, IMIN + mov r9, rax + mov r13, 0 +col_loop: + mov rax, r13 + imul rax, RSPAN + mov rcx, W + cqo + idiv rcx + sub rax, RMIN + mov r8, rax + mov r14, 0 + mov r15, 0 + mov rbx, 0 +iter_loop: + mov rax, r14 + imul rax, r14 + mov rcx, SCALE + cqo + idiv rcx + mov r10, rax + mov rax, r15 + imul rax, r15 + mov rcx, SCALE + cqo + idiv rcx + mov r11, rax + mov rax, r10 + add rax, r11 + cmp rax, FOUR + jle .if_end_0 + jmp plot +.if_end_0: + mov rax, r14 + imul rax, r15 + imul rax, 2 + mov rcx, SCALE + cqo + idiv rcx + add rax, r9 + mov rsi, rax + mov rax, r10 + sub rax, r11 + add rax, r8 + mov r14, rax + mov r15, rsi + add rbx, 1 + cmp rbx, MAX_ITER + jge .if_end_1 + jmp iter_loop +.if_end_1: +plot: + mov rax, rbx + imul rax, 9 + mov rcx, MAX_ITER + cqo + idiv rcx + mov rsi, palette + add rsi, rax + movzx rdx, byte [rsi] + lea rdi, [rbp - 81] + add rdi, r13 + mov byte [rdi], dl + add r13, 1 + cmp r13, W + jge .if_end_2 + jmp col_loop +.if_end_2: + lea rdi, [rbp - 81] + add rdi, W + mov rax, 10 + mov byte [rdi], al + mov rax, SYS_WRITE + mov rdi, STDOUT + lea rsi, [rbp - 81] + mov rdx, W + add rdx, 1 + syscall + add r12, 1 + cmp r12, H + jge .if_end_3 + jmp row_loop +.if_end_3: + mov rax, SYS_EXIT + mov rdi, 0 + syscall Binary files differdiff --git a/scripts/run.sh b/scripts/run.sh new file mode 100755 index 0000000..6d71521 --- /dev/null +++ b/scripts/run.sh @@ -0,0 +1,28 @@ +#!/usr/bin/env bash +set -euo pipefail + +root="$(cd "$(dirname "$0")/.." && pwd)" +cd "$root" + +if [ $# -lt 1 ]; then + echo "usage: $(basename "$0") <file.hdass>" >&2 + exit 1 +fi + +file="$1" +name="$(basename "${file%.*}")" + +docker compose up -d >/dev/null + +docker compose exec -T -e SRC="$file" -e NAME="$name" hdass bash -c ' + set -e + cd /hdass + premake5 gmake >/dev/null + make config=debug >/dev/null + ./bin/debug-linux/hdass "$SRC" -o "/tmp/$NAME.asm" + nasm -f elf64 "/tmp/$NAME.asm" -o "/tmp/$NAME.o" + ld -e main "/tmp/$NAME.o" -o "/tmp/$NAME" + set +e + "/tmp/$NAME" + printf "\n[exit %s]\n" "$?" +' </dev/null diff --git a/scripts/test_examples.sh b/scripts/test_examples.sh index 0928f7c..5290f09 100755 --- a/scripts/test_examples.sh +++ b/scripts/test_examples.sh @@ -85,6 +85,7 @@ check mul_div "multiply and non-rax division" examples/mul_div.hdass check load "stores then loads through a pointer" examples/load.hdass 7 "" check constants "hex literals and constant folding" examples/constants.hdass 42 "" check records "enum values and struct field offsets" examples/records.hdass 42 "" +check circle "floating-point math on SSE registers" examples/circle.hdass 28 "" check fibonacci "prints the first ten Fibonacci numbers" examples/fibonacci.hdass 0 "0 1 1 diff --git a/src/codegen/nasm.c b/src/codegen/nasm.c index 9c1a178..31bcee0 100644 --- a/src/codegen/nasm.c +++ b/src/codegen/nasm.c @@ -1,6 +1,7 @@ #include <stdio.h> #include <stddef.h> #include <stdint.h> +#include <stdlib.h> #include <string.h> #include <stdbool.h> @@ -64,14 +65,44 @@ static const char* assign_mnemonic(enum TokenType op) } } +struct FloatTable +{ + struct Token* items; + size_t count; + size_t capacity; +}; + struct Emitter { struct Program* program; struct ProcDecl* proc; + struct FloatTable* floats; FILE* out; uint32_t label_id; }; +static bool is_float_register(struct Token token) +{ + if (token.length < 4 || memcmp(token.start, "xmm", 3) != 0) + return false; + + for (size_t i = 3; i < token.length; i += 1) + if (token.start[i] < '0' || token.start[i] > '9') + return false; + + return true; +} + +static size_t float_index(struct FloatTable* floats, struct Token literal) +{ + for (size_t i = 0; i < floats->count; i += 1) + if (floats->items[i].length == literal.length + && memcmp(floats->items[i].start, literal.start, literal.length) == 0) + return i; + + return floats->count; +} + static const char* sized_register(struct Token reg, enum StoreSize size); static struct Token resolve_token(struct Emitter* emitter, struct Token token) @@ -170,8 +201,11 @@ static uint64_t token_to_u64(struct Token token) return value; } -static bool buffer_offset(struct ProcDecl* proc, struct Token name, uint64_t* out_offset) +static bool fold_const(struct Program* program, struct Expr* expr, uint64_t* out); + +static bool buffer_offset(struct Emitter* emitter, struct Token name, uint64_t* out_offset) { + struct ProcDecl* proc = emitter->proc; uint64_t cumulative = 0; for (size_t i = 0; i < proc->body_count; i += 1) { @@ -179,7 +213,9 @@ static bool buffer_offset(struct ProcDecl* proc, struct Token name, uint64_t* ou if (statement->kind != STATEMENT_STACK) continue; - cumulative += token_to_u64(statement->stack.size); + uint64_t size = 0; + fold_const(emitter->program, statement->stack.size, &size); + cumulative += size; if (statement->stack.name.length == name.length && memcmp(statement->stack.name.start, name.start, name.length) == 0) { @@ -194,7 +230,7 @@ static bool buffer_offset(struct ProcDecl* proc, struct Token name, uint64_t* ou static bool is_buffer_name(struct Emitter* emitter, struct Token token) { uint64_t offset; - return emitter->proc != NULL && buffer_offset(emitter->proc, token, &offset); + return emitter->proc != NULL && buffer_offset(emitter, token, &offset); } static enum StoreSize size_from_int(struct Token token) @@ -249,6 +285,117 @@ static uint64_t store_size_bytes(enum StoreSize size) } } +static uint64_t char_literal_value(struct Token token) +{ + if (token.length >= 4 && token.start[1] == '\\') + { + switch (token.start[2]) + { + case 'n': return 10; + case 't': return 9; + case 'r': return 13; + case '0': return 0; + case '\\': return 92; + case '\'': return 39; + default: return (unsigned char)token.start[2]; + } + } + + return (unsigned char)token.start[1]; +} + +static bool fold_member(struct Program* program, struct Expr* object, struct Token member, uint64_t* out) +{ + if (object->kind != EXPR_PRIMARY) + return false; + struct Token name = object->primary.token; + + struct EnumDecl* enumeration = find_enum(program, name); + if (enumeration != NULL) + { + for (size_t i = 0; i < enumeration->member_count; i += 1) + if (tokens_equal(enumeration->members[i], member)) + { + *out = i; + return true; + } + return false; + } + + struct StructDecl* layout = find_struct(program, name); + if (layout != NULL) + { + uint64_t offset = 0; + for (size_t i = 0; i < layout->field_count; i += 1) + { + if (tokens_equal(layout->fields[i].name, member)) + { + *out = offset; + return true; + } + offset += store_size_bytes(layout->fields[i].size); + } + if (token_matches(member, "size")) + { + *out = offset; + return true; + } + } + + return false; +} + +// evaluates a compile-time constant expression: integer/char literals, other +// constants, enum values and struct offsets, and + - * / over them +static bool fold_const(struct Program* program, struct Expr* expr, uint64_t* out) +{ + switch (expr->kind) + { + case EXPR_PRIMARY: + { + struct Token token = expr->primary.token; + if (token.type == TOKEN_INTEGER) + { + *out = token_to_u64(token); + return true; + } + if (token.type == TOKEN_CHAR) + { + *out = char_literal_value(token); + return true; + } + if (token.type == TOKEN_IDENTIFIER) + for (size_t i = 0; i < program->const_count; i += 1) + if (tokens_equal(program->consts[i].name, token)) + return fold_const(program, program->consts[i].value, out); + return false; + } + case EXPR_BINARY: + { + uint64_t left; + uint64_t right; + if (!fold_const(program, expr->binary.left, &left) + || !fold_const(program, expr->binary.right, &right)) + return false; + + switch (expr->binary.op.type) + { + case TOKEN_PLUS: *out = left + right; return true; + case TOKEN_MINUS: *out = left - right; return true; + case TOKEN_STAR: *out = left * right; return true; + case TOKEN_SLASH: *out = right != 0 ? left / right : 0; return true; + default: return false; + } + } + case EXPR_MEMBER: + return fold_member(program, expr->member.object, expr->member.member, out); + case EXPR_DEREF: + return false; + } + + return false; +} + // an enum member folds to its 0-based index; a struct member folds to its byte // offset (or the total size for `.size`) static bool emit_named_member(struct Emitter* emitter, struct Token object, struct Token member) @@ -442,7 +589,7 @@ static void emit_expr_into(struct Emitter* emitter, const char* dst, struct Expr if (expr->kind == EXPR_PRIMARY && is_buffer_name(emitter, expr->primary.token)) { uint64_t offset; - buffer_offset(emitter->proc, expr->primary.token, &offset); + buffer_offset(emitter, expr->primary.token, &offset); fprintf(emitter->out, "\tlea %s, [rbp - %llu]\n", dst, (unsigned long long)offset); return; } @@ -517,8 +664,108 @@ static const char* sized_register(struct Token reg, enum StoreSize size) return NULL; } +static const char* float_mnemonic(enum TokenType op) +{ + switch (op) + { + case TOKEN_EQUAL: return "movsd"; + case TOKEN_PLUS_EQUAL: return "addsd"; + case TOKEN_MINUS_EQUAL: return "subsd"; + case TOKEN_STAR_EQUAL: return "mulsd"; + case TOKEN_SLASH_EQUAL: return "divsd"; + default: return NULL; + } +} + +static bool value_is_float(struct Emitter* emitter, struct Expr* expr) +{ + if (expr->kind != EXPR_PRIMARY) + return false; + if (expr->primary.token.type == TOKEN_FLOAT) + return true; + + return is_float_register(resolve_register(emitter, expr->primary.token)); +} + +// floating point: xmm moves and arithmetic, conversions to/from general-purpose +// registers, and float literals loaded from their .data slot +static bool emit_float_assign(struct Emitter* emitter, struct AssignStatement* assign, + struct Token target, bool target_float) +{ + struct Expr* value = assign->value; + + // float store: ^ptr = xmm -> movsd [ptr], xmm + if (assign->target_deref) + { + if (assign->op.type != TOKEN_EQUAL || value->kind != EXPR_PRIMARY) + return false; + struct Token source = resolve_register(emitter, value->primary.token); + if (!is_float_register(source)) + return false; + fprintf(emitter->out, "\tmovsd [%.*s], %.*s\n", + (int)target.length, target.start, (int)source.length, source.start); + return true; + } + + // float load: xmm = ^ptr -> movsd xmm, [ptr] + if (value->kind == EXPR_DEREF) + { + if (!target_float || assign->op.type != TOKEN_EQUAL) + return false; + fprintf(emitter->out, "\tmovsd %.*s, [", (int)target.length, target.start); + emit_operand(emitter, value->deref.address); + fprintf(emitter->out, "]\n"); + return true; + } + + if (value->kind == EXPR_PRIMARY && value->primary.token.type == TOKEN_FLOAT) + { + if (!target_float || assign->op.type != TOKEN_EQUAL) + return false; + size_t index = float_index(emitter->floats, value->primary.token); + fprintf(emitter->out, "\tmovsd %.*s, [__float%zu]\n", + (int)target.length, target.start, index); + return true; + } + + if (value->kind != EXPR_PRIMARY) + return false; + + struct Token source = resolve_register(emitter, value->primary.token); + bool source_float = is_float_register(source); + + if (target_float && source_float) + { + const char* mnemonic = float_mnemonic(assign->op.type); + if (mnemonic == NULL) + return false; + fprintf(emitter->out, "\t%s %.*s, %.*s\n", mnemonic, + (int)target.length, target.start, (int)source.length, source.start); + return true; + } + + if (assign->op.type != TOKEN_EQUAL) + return false; + + if (target_float) + fprintf(emitter->out, "\tcvtsi2sd %.*s, %.*s\n", + (int)target.length, target.start, (int)source.length, source.start); + else + fprintf(emitter->out, "\tcvttsd2si %.*s, %.*s\n", + (int)target.length, target.start, (int)source.length, source.start); + return true; +} + static void emit_assign(struct Emitter* emitter, struct AssignStatement* assign) { + struct Token float_target = resolve_register(emitter, assign->target); + if (is_float_register(float_target) || value_is_float(emitter, assign->value)) + { + if (!emit_float_assign(emitter, assign, float_target, is_float_register(float_target))) + fprintf(emitter->out, "\t; TODO: unsupported float assignment\n"); + return; + } + if (assign->op.type == TOKEN_SLASH_EQUAL) { emit_divide(emitter, assign); @@ -552,6 +799,13 @@ static void emit_assign(struct Emitter* emitter, struct AssignStatement* assign) return; } + // adding or subtracting a constant zero (e.g. a struct field at offset 0) is a no-op + uint64_t folded; + if (!assign->target_deref + && (assign->op.type == TOKEN_PLUS_EQUAL || assign->op.type == TOKEN_MINUS_EQUAL) + && fold_const(emitter->program, assign->value, &folded) && folded == 0) + return; + if (assign->target_deref) { fprintf(emitter->out, "\t%s %s[%.*s], ", mnemonic, @@ -632,8 +886,63 @@ static void emit_call(struct Emitter* emitter, struct CallStatement* call) static void emit_statement(struct Emitter* emitter, struct Statement* statement); +// 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) +{ + switch (comparison) + { + case TOKEN_EQUAL_EQUAL: return "jne"; + case TOKEN_BANG_EQUAL: return "je"; + case TOKEN_LESS: return "jae"; + case TOKEN_LESS_EQUAL: return "ja"; + case TOKEN_GREATER: return "jbe"; + case TOKEN_GREATER_EQUAL: return "jb"; + default: return NULL; + } +} + +static void emit_float_operand(struct Emitter* emitter, struct Expr* expr) +{ + if (expr->kind == EXPR_PRIMARY && expr->primary.token.type == TOKEN_FLOAT) + { + fprintf(emitter->out, "[__float%zu]", float_index(emitter->floats, expr->primary.token)); + return; + } + + struct Token token = resolve_register(emitter, expr->primary.token); + fprintf(emitter->out, "%.*s", (int)token.length, token.start); +} + 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); + + uint32_t id = emitter->label_id; + emitter->label_id += 1; + + 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)) + { + fprintf(emitter->out, "\t; TODO: unsupported if\n"); + return; + } + + fprintf(emitter->out, "\tucomisd "); + 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; + } + const char* jump = jump_if_false(branch->comparison.type); if (jump == NULL || branch->left->kind == EXPR_BINARY || branch->left->kind == EXPR_DEREF @@ -643,9 +952,6 @@ static void emit_if(struct Emitter* emitter, struct IfStatement* branch) return; } - uint32_t id = emitter->label_id; - emitter->label_id += 1; - fprintf(emitter->out, "\tcmp "); emit_operand(emitter, branch->left); fprintf(emitter->out, ", "); @@ -689,14 +995,18 @@ static void emit_statement(struct Emitter* emitter, struct Statement* statement) } } -static uint64_t proc_stack_size(struct ProcDecl* proc) +static uint64_t proc_stack_size(struct Program* program, struct ProcDecl* proc) { uint64_t total = 0; for (size_t i = 0; i < proc->body_count; i += 1) { struct Statement* statement = &proc->body[i]; if (statement->kind == STATEMENT_STACK) - total += token_to_u64(statement->stack.size); + { + uint64_t size = 0; + fold_const(program, statement->stack.size, &size); + total += size; + } } if (total % 16 != 0) @@ -705,11 +1015,86 @@ static uint64_t proc_stack_size(struct ProcDecl* proc) return total; } -static void emit_proc(struct Program* program, struct ProcDecl* proc, FILE* out) +static void collect_float(struct FloatTable* floats, struct Token token) +{ + if (token.type != TOKEN_FLOAT || float_index(floats, token) != floats->count) + return; + + if (floats->count == floats->capacity) + { + size_t capacity = floats->capacity < 8 ? 8 : floats->capacity * 2; + floats->items = realloc(floats->items, capacity * sizeof(struct Token)); + floats->capacity = capacity; + } + + floats->items[floats->count] = token; + floats->count += 1; +} + +static void collect_floats_expr(struct FloatTable* floats, struct Expr* expr) +{ + switch (expr->kind) + { + case EXPR_PRIMARY: + collect_float(floats, expr->primary.token); + break; + case EXPR_BINARY: + collect_floats_expr(floats, expr->binary.left); + collect_floats_expr(floats, expr->binary.right); + break; + case EXPR_MEMBER: + collect_floats_expr(floats, expr->member.object); + break; + case EXPR_DEREF: + collect_floats_expr(floats, expr->deref.address); + break; + } +} + +static void collect_floats_statement(struct FloatTable* floats, struct Statement* statement) +{ + switch (statement->kind) + { + case STATEMENT_ASSIGN: + collect_floats_expr(floats, statement->assign.value); + break; + case STATEMENT_IF: + collect_floats_expr(floats, statement->branch.left); + collect_floats_expr(floats, statement->branch.right); + collect_floats_statement(floats, statement->branch.body); + break; + case STATEMENT_CALL: + for (size_t i = 0; i < statement->call.arg_count; i += 1) + collect_floats_expr(floats, statement->call.args[i]); + break; + default: + break; + } +} + +static struct FloatTable collect_floats(struct Program* program) +{ + struct FloatTable floats = { NULL, 0, 0 }; + for (size_t i = 0; i < program->proc_count; i += 1) + for (size_t j = 0; j < program->procs[i].body_count; j += 1) + collect_floats_statement(&floats, &program->procs[i].body[j]); + + return floats; +} + +static void emit_float_data(struct FloatTable* floats, FILE* out) +{ + for (size_t i = 0; i < floats->count; i += 1) + fprintf(out, "__float%zu: dq %.*s\n", i, + (int)floats->items[i].length, floats->items[i].start); +} + +static void emit_proc(struct Program* program, struct FloatTable* floats, struct ProcDecl* proc, FILE* out) { struct Emitter emitter; emitter.program = program; emitter.proc = proc; + emitter.floats = floats; emitter.out = out; emitter.label_id = 0; @@ -720,7 +1105,7 @@ static void emit_proc(struct Program* program, struct ProcDecl* proc, FILE* out) fprintf(out, "%.*s:\n", (int)proc->name.length, proc->name.start); - uint64_t stack_size = proc_stack_size(proc); + uint64_t stack_size = proc_stack_size(program, proc); if (stack_size > 0) { fprintf(out, "\tpush rbp\n"); @@ -741,6 +1126,8 @@ static void emit_proc(struct Program* program, struct ProcDecl* proc, FILE* out) void generate_nasm(struct Program* program, FILE* out) { + struct FloatTable floats = collect_floats(program); + fprintf(out, "bits %u\n\n", program->config.bits); if (program->const_count > 0) @@ -750,6 +1137,7 @@ void generate_nasm(struct Program* program, FILE* out) } emit_data(program, out); + emit_float_data(&floats, out); fprintf(out, "\n"); fprintf(out, "section .text\n"); @@ -759,6 +1147,8 @@ void generate_nasm(struct Program* program, FILE* out) for (size_t i = 0; i < program->proc_count; i += 1) { fprintf(out, "\n"); - emit_proc(program, &program->procs[i], out); + emit_proc(program, &floats, &program->procs[i], out); } + + free(floats.items); } diff --git a/src/lexer/lexer.c b/src/lexer/lexer.c index 22a671c..b4fb7d7 100644 --- a/src/lexer/lexer.c +++ b/src/lexer/lexer.c @@ -191,18 +191,28 @@ struct Token scan_token(struct Lexer* lexer) advance(lexer); while (is_hex_digit(peek(lexer))) advance(lexer); + return make_token(lexer, TOKEN_INTEGER, start); } - else if (character == '0' && (peek(lexer) == 'b' || peek(lexer) == 'B')) + + if (character == '0' && (peek(lexer) == 'b' || peek(lexer) == 'B')) { advance(lexer); while (peek(lexer) == '0' || peek(lexer) == '1') advance(lexer); + return make_token(lexer, TOKEN_INTEGER, start); } - else + + while (is_digit(peek(lexer))) + advance(lexer); + + if (peek(lexer) == '.' && is_digit(lexer->current[1])) { + advance(lexer); while (is_digit(peek(lexer))) advance(lexer); + return make_token(lexer, TOKEN_FLOAT, start); } + return make_token(lexer, TOKEN_INTEGER, start); } @@ -244,6 +254,7 @@ const char* token_type_name(enum TokenType type) case TOKEN_EOF: return "eof"; case TOKEN_IDENTIFIER: return "identifier"; case TOKEN_INTEGER: return "integer"; + case TOKEN_FLOAT: return "float"; case TOKEN_STRING: return "string"; case TOKEN_CHAR: return "char"; case TOKEN_CONST: return "const"; diff --git a/src/lexer/lexer.h b/src/lexer/lexer.h index 3bc215e..5c1e868 100644 --- a/src/lexer/lexer.h +++ b/src/lexer/lexer.h @@ -8,6 +8,7 @@ enum TokenType TOKEN_EOF, TOKEN_IDENTIFIER, TOKEN_INTEGER, + TOKEN_FLOAT, TOKEN_STRING, TOKEN_CHAR, diff --git a/src/parser/ast.c b/src/parser/ast.c index f8bd9cc..788341b 100644 --- a/src/parser/ast.c +++ b/src/parser/ast.c @@ -44,6 +44,9 @@ static void free_statement(struct Statement* statement) free_expr(statement->call.args[i]); free(statement->call.args); break; + case STATEMENT_STACK: + free_expr(statement->stack.size); + break; default: break; } diff --git a/src/parser/ast.h b/src/parser/ast.h index b114fb1..c439747 100644 --- a/src/parser/ast.h +++ b/src/parser/ast.h @@ -148,7 +148,7 @@ struct CallStatement struct StackStatement { struct Token name; - struct Token size; + struct Expr* size; }; struct Statement diff --git a/src/parser/parser.c b/src/parser/parser.c index 682e193..3ed9fff 100644 --- a/src/parser/parser.c +++ b/src/parser/parser.c @@ -248,7 +248,8 @@ static struct Expr* parse_primary(struct Parser* parser) return deref; } - if (check(parser, TOKEN_IDENTIFIER) || check(parser, TOKEN_INTEGER) || check(parser, TOKEN_CHAR)) + if (check(parser, TOKEN_IDENTIFIER) || check(parser, TOKEN_INTEGER) + || check(parser, TOKEN_FLOAT) || check(parser, TOKEN_CHAR)) { advance_parser(parser); @@ -439,12 +440,15 @@ static bool parse_statement(struct Parser* parser, struct Statement* out) if (!consume(parser, TOKEN_LEFT_BRACKET, "expected '[' after buffer name")) return false; - if (!consume(parser, TOKEN_INTEGER, "expected buffer size")) + struct Expr* size = parse_expression(parser); + if (size == NULL) return false; - struct Token size = parser->previous; if (!consume(parser, TOKEN_RIGHT_BRACKET, "expected ']' after buffer size")) + { + free_expr(size); return false; + } out->kind = STATEMENT_STACK; out->stack.name = name; diff --git a/src/sema/sema.c b/src/sema/sema.c index d2b4668..12fd53f 100644 --- a/src/sema/sema.c +++ b/src/sema/sema.c @@ -191,6 +191,8 @@ static bool is_arch_register(struct Token token) "r14", "r14d", "r14w", "r14b", "r15", "r15d", "r15w", "r15b", "rip", + "xmm0", "xmm1", "xmm2", "xmm3", "xmm4", "xmm5", "xmm6", "xmm7", + "xmm8", "xmm9", "xmm10", "xmm11", "xmm12", "xmm13", "xmm14", "xmm15", }; for (size_t i = 0; i < sizeof(names) / sizeof(names[0]); i += 1) @@ -392,6 +394,43 @@ static void check_target(struct RefCheck* check, struct Token target) ref_error(check, target, "cannot assign to '%.*s': not a register", (int)target.length, target.start); } +// a stack size must be a compile-time constant: integer/char literals, other +// constants, enum values, struct sizes/offsets, and arithmetic over them +static void check_stack_size(struct RefCheck* check, struct Expr* expr) +{ + switch (expr->kind) + { + case EXPR_BINARY: + check_stack_size(check, expr->binary.left); + check_stack_size(check, expr->binary.right); + break; + case EXPR_MEMBER: + { + struct Expr* object = expr->member.object; + if (object->kind == EXPR_PRIMARY + && (find_enum(check, object->primary.token) != NULL + || find_struct(check, object->primary.token) != NULL)) + check_expr(check, expr); + else + ref_error(check, first_token(expr), "stack size must be a constant"); + break; + } + case EXPR_PRIMARY: + { + struct Token token = expr->primary.token; + if (token.type == TOKEN_INTEGER || token.type == TOKEN_CHAR) + break; + if (token.type == TOKEN_IDENTIFIER && is_const(check, token)) + break; + ref_error(check, token, "stack size must be a constant"); + break; + } + default: + ref_error(check, first_token(expr), "stack size must be a constant"); + break; + } +} + static void check_statement(struct RefCheck* check, struct Statement* statement) { switch (statement->kind) @@ -425,9 +464,11 @@ static void check_statement(struct RefCheck* check, struct Statement* statement) check_expr(check, call->args[i]); break; } + case STATEMENT_STACK: + check_stack_size(check, statement->stack.size); + break; case STATEMENT_LABEL: case STATEMENT_SYSCALL: - case STATEMENT_STACK: break; } } diff --git a/tests/codegen_test.c b/tests/codegen_test.c index 61e441f..377ddeb 100644 --- a/tests/codegen_test.c +++ b/tests/codegen_test.c @@ -426,8 +426,109 @@ static void test_generate_enum_struct(struct TestContext* context) free_program(&program); } +static void test_generate_floats(struct TestContext* context) +{ + struct Lexer lexer = create_lexer( + "proc main\n{\nxmm0 = 3.5\nxmm0 *= xmm1\nrax = 4\nxmm2 = rax\nrbx = xmm0\n}\n"); + struct Program program; + check(context, parse_program(&lexer, &program)); + + FILE* out = tmpfile(); + generate_nasm(&program, out); + fflush(out); + rewind(out); + + char buffer[1024]; + size_t read = fread(buffer, 1, sizeof(buffer) - 1, out); + buffer[read] = '\0'; + fclose(out); + + check(context, strstr(buffer, "__float0: dq 3.5") != NULL); + check(context, strstr(buffer, "movsd xmm0, [__float0]") != NULL); + check(context, strstr(buffer, "mulsd xmm0, xmm1") != NULL); // float arithmetic + check(context, strstr(buffer, "cvtsi2sd xmm2, rax") != NULL); // int -> float + check(context, strstr(buffer, "cvttsd2si rbx, xmm0") != NULL); // float -> int + check(context, strstr(buffer, "; TODO") == NULL); + + free_program(&program); +} + +static void test_generate_float_compare(struct TestContext* context) +{ + struct Lexer lexer = create_lexer("proc main\n{\nif xmm0 > 4.0\ngoto done\ndone:\nsyscall\n}\n"); + struct Program program; + check(context, parse_program(&lexer, &program)); + + FILE* out = tmpfile(); + generate_nasm(&program, out); + fflush(out); + rewind(out); + + char buffer[1024]; + size_t read = fread(buffer, 1, sizeof(buffer) - 1, out); + buffer[read] = '\0'; + fclose(out); + + check(context, strstr(buffer, "ucomisd xmm0, [__float0]") != NULL); + check(context, strstr(buffer, "jbe .if_end") != NULL); // '>' skips when <= + check(context, strstr(buffer, "; TODO") == NULL); + + free_program(&program); +} + +static void test_generate_add_zero_peephole(struct TestContext* context) +{ + // `+= 0` / `-= 0` (e.g. a struct field at offset 0) is dropped; `*= 0` is not + struct Lexer lexer = create_lexer("proc main\n{\nrax += 0\nrbx -= 0\nrcx *= 0\n}\n"); + struct Program program; + check(context, parse_program(&lexer, &program)); + + FILE* out = tmpfile(); + generate_nasm(&program, out); + fflush(out); + rewind(out); + + char buffer[1024]; + size_t read = fread(buffer, 1, sizeof(buffer) - 1, out); + buffer[read] = '\0'; + fclose(out); + + check(context, strstr(buffer, "add rax, 0") == NULL); + check(context, strstr(buffer, "sub rbx, 0") == NULL); + check(context, strstr(buffer, "imul rcx, 0") != NULL); + + free_program(&program); +} + +static void test_generate_float_memory(struct TestContext* context) +{ + struct Lexer lexer = create_lexer("proc main\n{\n^rsi = xmm0\nxmm1 = ^rsi\n}\n"); + struct Program program; + check(context, parse_program(&lexer, &program)); + + FILE* out = tmpfile(); + generate_nasm(&program, out); + fflush(out); + rewind(out); + + char buffer[1024]; + size_t read = fread(buffer, 1, sizeof(buffer) - 1, out); + buffer[read] = '\0'; + fclose(out); + + check(context, strstr(buffer, "movsd [rsi], xmm0") != NULL); // float store + check(context, strstr(buffer, "movsd xmm1, [rsi]") != NULL); // float load + check(context, strstr(buffer, "; TODO") == NULL); + + free_program(&program); +} + void run_codegen_tests(struct TestContext* context) { + test_generate_floats(context); + test_generate_float_compare(context); + test_generate_float_memory(context); + test_generate_add_zero_peephole(context); test_generate_enum_struct(context); test_generate_load(context); test_generate_logical_registers(context); diff --git a/tests/lexer_test.c b/tests/lexer_test.c index 38644d9..a44b1f0 100644 --- a/tests/lexer_test.c +++ b/tests/lexer_test.c @@ -28,6 +28,14 @@ static void test_number_bases(struct TestContext* context) check(context, token_matches(scan_token(&lexer), TOKEN_INTEGER, "0b1010")); } +static void test_float_literals(struct TestContext* context) +{ + struct Lexer lexer = create_lexer("3.14 1.0 42"); + check(context, token_matches(scan_token(&lexer), TOKEN_FLOAT, "3.14")); + check(context, token_matches(scan_token(&lexer), TOKEN_FLOAT, "1.0")); + check(context, scan_token(&lexer).type == TOKEN_INTEGER); +} + static void test_operators(struct TestContext* context) { struct Lexer lexer = create_lexer("= == += != /"); @@ -92,6 +100,7 @@ void run_lexer_tests(struct TestContext* context) { test_identifiers_and_integers(context); test_number_bases(context); + test_float_literals(context); test_operators(context); test_literals(context); test_keywords(context); diff --git a/tests/parser_test.c b/tests/parser_test.c index 71e1433..dbb250f 100644 --- a/tests/parser_test.c +++ b/tests/parser_test.c @@ -199,7 +199,7 @@ static void test_parse_stack(struct TestContext* context) struct Statement statement = program.procs[0].body[0]; check(context, statement.kind == STATEMENT_STACK); check(context, text_is(statement.stack.name, "buffer")); - check(context, text_is(statement.stack.size, "32")); + check(context, primary_is(statement.stack.size, "32")); free_program(&program); } diff --git a/tests/sema_test.c b/tests/sema_test.c index e89f608..dc01d47 100644 --- a/tests/sema_test.c +++ b/tests/sema_test.c @@ -97,6 +97,14 @@ static void test_deref_needs_register(struct TestContext* context) check(context, analyze_source("proc main\n{\nrax = ^rsi\n}\n")); } +static void test_stack_size_constant(struct TestContext* context) +{ + check(context, analyze_source( + "struct P\n{\nx\ny\n}\nconst N = 4\n" + "proc main\n{\nstack a[P.size]\nstack b[N * 2]\nsyscall\n}\n")); + check(context, !analyze_source("proc main\n{\nstack a[rax]\n}\n")); +} + static void test_enum_struct_members(struct TestContext* context) { check(context, analyze_source( @@ -144,6 +152,7 @@ void run_sema_tests(struct TestContext* context) test_const_expr_rejects_register(context); test_const_expr_rejects_data(context); test_deref_needs_register(context); + test_stack_size_constant(context); test_enum_struct_members(context); test_references_resolve(context); } |
