Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
64 changes: 52 additions & 12 deletions codegen/codegen.lua
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Expand All @@ -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)
Expand Down
36 changes: 36 additions & 0 deletions tests/codegen_return.lua
Original file line number Diff line number Diff line change
@@ -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")