From a7e1a0da1ebd08a5cd27f2e389d4d44eaa32aa01 Mon Sep 17 00:00:00 2001 From: SM Date: Sat, 6 Jun 2026 13:12:39 +0200 Subject: [PATCH] align return generation with parser AST --- codegen/codegen.lua | 64 ++++++++++++++++++++++++++++++++-------- tests/codegen_return.lua | 36 ++++++++++++++++++++++ 2 files changed, 88 insertions(+), 12 deletions(-) create mode 100644 tests/codegen_return.lua diff --git a/codegen/codegen.lua b/codegen/codegen.lua index 8ffe238..e22d644 100644 --- a/codegen/codegen.lua +++ b/codegen/codegen.lua @@ -26,18 +26,41 @@ function Codegen:new_label(prefix) end function Codegen:gen_expression(node) - if node.type == "number" then + if node.type == "literal" then + if type(node.value) ~= "number" then + error("Unsupported literal value: "..tostring(node.value)) + end local r = self:next_reg() self:emit(string.format("MOV %s, %s", r, node.value)) return r - elseif node.type == "variable" then - return self.env[node.name] -- variable already in some register + elseif node.type == "identifier" then + local r = self.env[node.name] + if not r then + error("Undefined identifier: "..tostring(node.name)) + end + return r elseif node.type == "binary" then + if node.op == "=" then + if node.left.type ~= "identifier" then + error("Assignment target must be an identifier") + end + local dest = self.env[node.left.name] + if not dest then + error("Undefined identifier: "..tostring(node.left.name)) + end + local right = self:gen_expression(node.right) + self:emit(string.format("MOV %s, %s", dest, right)) + return dest + end + local left = self:gen_expression(node.left) local right = self:gen_expression(node.right) -- Assume left is dest, operate on it with right local op_map = {["+"]="ADD", ["-"]="SUB", ["*"]="MUL", ["/"]="DIV"} - local op = op_map[node.operator] + local op = op_map[node.op] + if not op then + error("Unsupported binary operator: "..tostring(node.op)) + end self:emit(string.format("%s %s, %s", op, left, right)) return left else @@ -62,34 +85,49 @@ function Codegen:gen_if(node) local end_label = self:new_label("endif") self:emit(string.format("JZ %s, %s", cond_reg, else_label)) - self:gen_block(node.then_block) + local then_returns = self:gen_block(node.thenBranch) self:emit(string.format("JMP %s", end_label)) self:emit(else_label .. ":") - if node.else_block then - self:gen_block(node.else_block) + local else_returns = false + if node.elseBranch then + else_returns = self:gen_block(node.elseBranch) end self:emit(end_label .. ":") + return then_returns and else_returns end function Codegen:gen_block(block) + local returns = false for _, stmt in ipairs(block) do - self:gen_statement(stmt) + returns = self:gen_statement(stmt) or false end + return returns end function Codegen:gen_statement(node) - if node.type == "declaration" then + if not node.type then + local returns = false + for _, stmt in ipairs(node) do + returns = self:gen_statement(stmt) or false + end + return returns + elseif node.type == "decl" then self:gen_declaration(node) elseif node.type == "expression" then self:gen_expression(node.expr) elseif node.type == "if" then - self:gen_if(node) + return self:gen_if(node) elseif node.type == "block" then self:gen_block(node.statements) local ret_reg = self:gen_expression(node.value) self:emit(string.format("MOV r0, %s", ret_reg)) -- r0 = return register elseif node.type == "return" then + if node.value then + local ret_reg = self:gen_expression(node.value) + self:emit(string.format("MOV r0, %s", ret_reg)) + end self:emit("RET") + return true elseif node.type == "function" then self.env = {} -- clear env for new func self.regCount = 0 @@ -100,11 +138,13 @@ function Codegen:gen_statement(node) -- assume parameters are passed in registers r1..rN self:emit(string.format("MOV %s, arg_%s", r, param.name)) end - self:gen_block(node.body) - self:emit("RET") + if not self:gen_block(node.body) then + self:emit("RET") + end else error("Unknown statement type: "..tostring(node.type)) end + return false end function Codegen:generate(ast) diff --git a/tests/codegen_return.lua b/tests/codegen_return.lua new file mode 100644 index 0000000..5df846d --- /dev/null +++ b/tests/codegen_return.lua @@ -0,0 +1,36 @@ +local Parser = require("lang/parser") +local Codegen = require("codegen/codegen") + +local function assert_equal(actual, expected, label) + if actual ~= expected then + error(string.format("%s: expected %q, got %q", label, expected, actual)) + end +end + +local code = [[ + fn int test(int x) { + return x + 2 + } +]] + +local parser = Parser:new(code) +local ast = parser:parse() + +local codegen = Codegen:new() +local insns = codegen:generate(ast.body) + +local expected = { + "test:", + "MOV r1, arg_x", + "MOV r2, 2", + "ADD r1, r2", + "MOV r0, r1", + "RET", +} + +assert_equal(#insns, #expected, "instruction count") +for i, expected_insn in ipairs(expected) do + assert_equal(insns[i], expected_insn, "instruction " .. i) +end + +print("ok")