aboutsummaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
authorJan Tuomi <jans.tuomi@gmail.com>2023-11-11 14:40:25 +0200
committerJan Tuomi <jans.tuomi@gmail.com>2023-11-11 14:40:25 +0200
commitf7dd9705d44b9cee5103b267797c0449d42319df (patch)
treea63c68f282cf64e250384a2e1e8525d2333b9dde
parentd3abfd22d4f5fff02adaebe28707aea7bd804ba3 (diff)
WIP functions
-rw-r--r--asm/ast_compiler_bootstrap.atk162
-rw-r--r--asm/bootstrap.atk162
-rw-r--r--asm/ext_std.py28
-rw-r--r--asm/test_py_src.atk1674
-rw-r--r--asm/test_py_src.atk16_optimized42
-rw-r--r--asm/test_py_src.py14
-rw-r--r--src/asm_pass4.py4
-rw-r--r--src/ast_compiler.py356
-rw-r--r--src/tokenizer.py6
9 files changed, 335 insertions, 193 deletions
diff --git a/asm/ast_compiler_bootstrap.atk16 b/asm/ast_compiler_bootstrap.atk16
index ad401f1..7297c90 100644
--- a/asm/ast_compiler_bootstrap.atk16
+++ b/asm/ast_compiler_bootstrap.atk16
@@ -31,7 +31,7 @@
@address 0x0
ldi vt_stack_addr RA
- ldr RA __STACK_POINTER
+ ldr RA SP
jpi program_segment
@address vector_table
diff --git a/asm/bootstrap.atk16 b/asm/bootstrap.atk16
index 6363d71..4e5cd24 100644
--- a/asm/bootstrap.atk16
+++ b/asm/bootstrap.atk16
@@ -61,7 +61,7 @@
@address 0x0
; set up stack pointer to point to beginning of stack segment
ldi vt_stack_addr RA
- ldr RA __STACK_POINTER
+ ldr RA SP
jpi program_segment
@address program_segment
diff --git a/asm/ext_std.py b/asm/ext_std.py
index fd18953..f4751a3 100644
--- a/asm/ext_std.py
+++ b/asm/ext_std.py
@@ -67,29 +67,35 @@ def expand_nop() -> ExpandResult:
return [["ali", "al_plus", "RA", "0", "RA"]]
def expand_spu(reg: str) -> ExpandResult:
- return [["str", reg, "__STACK_POINTER"]] + expand_inc("__STACK_POINTER")
+ return [["str", reg, "SP"]] + expand_inc("SP")
-def expand_spo(reg: str):
- return expand_dec("__STACK_POINTER") + [["ldr", "__STACK_POINTER", reg]]
+def expand_spo(reg: str) -> ExpandResult:
+ return expand_dec("SP") + [["ldr", "SP", reg]]
+
+def expand_sinc(imm: str) -> ExpandResult:
+ return expand_addi("SP", imm, "SP")
+
+def expand_sdec(imm: str) -> ExpandResult:
+ return expand_subi("SP", imm, "SP")
def expand_csr(addr_reg: str):
return [
- ["lpc", "__CSR_SCRATCH"],
- *expand_addi("__CSR_SCRATCH", "4", "__CSR_SCRATCH"),
- *expand_spu("__CSR_SCRATCH"),
+ ["lpc", "CSR_SCRATCH"],
+ *expand_addi("CSR_SCRATCH", "4", "CSR_SCRATCH"),
+ *expand_spu("CSR_SCRATCH"),
["jpr", addr_reg]
]
def expand_csi(addr_imm: str):
return [
- ["lpc", "__CSR_SCRATCH"],
- *expand_addi("__CSR_SCRATCH", "4", "__CSR_SCRATCH"),
- *expand_spu("__CSR_SCRATCH"),
+ ["lpc", "CSR_SCRATCH"],
+ *expand_addi("CSR_SCRATCH", "4", "CSR_SCRATCH"),
+ *expand_spu("CSR_SCRATCH"),
["jpi", addr_imm]
]
def expand_rsr():
- return expand_spo("__CSR_SCRATCH") + [["jpr", "__CSR_SCRATCH"]]
+ return expand_spo("CSR_SCRATCH") + [["jpr", "CSR_SCRATCH"]]
def stack_stash(*rs: str):
@@ -140,6 +146,8 @@ expansions: OpExpansionDict = {
"nop": expand_nop,
"spu": expand_spu,
"spo": expand_spo,
+ "sinc": expand_sinc,
+ "sdec": expand_sdec,
"csr": expand_csr,
"csi": expand_csi,
"rsr": expand_rsr,
diff --git a/asm/test_py_src.atk16 b/asm/test_py_src.atk16
index c035b72..ca44267 100644
--- a/asm/test_py_src.atk16
+++ b/asm/test_py_src.atk16
@@ -31,7 +31,7 @@
@address 0x0
ldi vt_stack_addr RA
- ldr RA __STACK_POINTER
+ ldr RA SP
jpi program_segment
@address vector_table
@@ -75,42 +75,48 @@
-@label main
-; <ast.If object at 0x1013c1e40>
-; <ast.Constant object at 0x1013c20b0>
- ldi 1 RA
+; <ast.FunctionDef object at 0x100c7b310>
+@label func
+; stack frame: ['arg1', 'arg2', 'a']
+ addi SP 1 SP
+; <ast.Name object at 0x100c7aef0> (arg1)
+ ldi ${SP - 3} RA
+ ldr RA RA
+ spu RA
+ spo RA
+ ldi ${SP - 1} RB
+ str RA RB
+; <ast.Return object at 0x100c7b0a0>
+; <ast.BinOp object at 0x100c7b370>
+; BinOp lhs <ast.BinOp object at 0x100c7b370>
+; <ast.Name object at 0x100c7b3a0> (a)
+ ldi ${SP - 1} RA
+ ldr RA RA
spu RA
+; BinOp rhs <ast.BinOp object at 0x100c7b370>
+; <ast.Name object at 0x100c7b3d0> (a)
+ ldi ${SP - 1} RA
+ ldr RA RA
+ spu RA
+ spo RB
spo RA
- addi RA 0 RA
- bri zero If_false_branch_0
-; <ast.Constant object at 0x1013c2110>
- ldi 2 RA
+ add RA RB RA
spu RA
- jpi If_end_branch_1
-@label If_false_branch_0
-; <ast.Constant object at 0x1013c2170>
+ subi SP 1 SP
+ rsr
+
+@label main
+; stack frame: ['b']
+ addi SP 1 SP
+; <ast.Call object at 0x100c7b520>
+; <ast.Constant object at 0x100c7b580>
+ ldi 1 RA
+ spu RA
+; <ast.Constant object at 0x100c7b5b0>
ldi 3 RA
spu RA
-@label If_end_branch_1
-; <ast.While object at 0x1013c2230>
-@label While_test_2
-; <ast.Constant object at 0x1013c2260>
- ldi 3 RB
- spu RB
+ csi func
spo RA
- addi RA 0 RA
- bri zero While_else_3
-; <ast.Constant object at 0x1013c22c0>
- ldi 4 RB
- spu RB
-; <ast.Break object at 0x1013c22f0>
- jpi While_end_4
-; <ast.Constant object at 0x1013c2350>
- ldi 5 RB
- spu RB
- jpi While_test_2
-@label While_else_3
-; <ast.Constant object at 0x1013c2470>
- ldi 6 RB
- spu RB
-@label While_end_4 \ No newline at end of file
+ ldi ${SP - 1} RB
+ str RA RB
+ subi SP 1 SP
diff --git a/asm/test_py_src.atk16_optimized b/asm/test_py_src.atk16_optimized
index 673800b..4a65ec4 100644
--- a/asm/test_py_src.atk16_optimized
+++ b/asm/test_py_src.atk16_optimized
@@ -24,7 +24,7 @@
@let gr_sprite_mode 0b10
@address 0x0
ldi vt_stack_addr RA
- ldr RA __STACK_POINTER
+ ldr RA SP
jpi program_segment
@address vector_table
keyboard_isr
@@ -58,28 +58,28 @@
ldi gr_disabled_mode RB
str RB RA
jpi main
+@label func
+ addi SP 1 SP
+ ldi ${SP - 3} RA
+ ldr RA RA
+ ldi ${SP - 1} RB
+ str RA RB
+ ldi ${SP - 1} RA
+ ldr RA RA
+ ldi ${SP - 1} RB
+ ldr RB RB
+ add RA RB RA
+ spu RA
+ subi SP 1 SP
+ rsr
@label main
+ addi SP 1 SP
ldi 1 RA
- addi RA 0 RA
- bri zero If_false_branch_0
- ldi 2 RA
spu RA
- jpi If_end_branch_1
-@label If_false_branch_0
ldi 3 RA
spu RA
-@label If_end_branch_1
-@label While_test_2
- ldi 3 RA
- addi RA 0 RA
- bri zero While_else_3
- ldi 4 RB
- spu RB
- jpi While_end_4
- ldi 5 RB
- spu RB
- jpi While_test_2
-@label While_else_3
- ldi 6 RB
- spu RB
-@label While_end_4 \ No newline at end of file
+ csi func
+ spo RA
+ ldi ${SP - 1} RB
+ str RA RB
+ subi SP 1 SP \ No newline at end of file
diff --git a/asm/test_py_src.py b/asm/test_py_src.py
index 48a5cac..9af944e 100644
--- a/asm/test_py_src.py
+++ b/asm/test_py_src.py
@@ -1,16 +1,10 @@
import atk16
-if True:
- 2
-else:
- 3
+def func(arg1: int, arg2: int) -> int:
+ a = arg1
+ return a + a
-while 3:
- 4
- break
- 5
-else:
- 6
+b = func(1, 3)
# TEXT_MODE: atk16.ConstInt = 1
# GRAPHICS_MODE_ADDR: atk16.ConstInt = 0x17
diff --git a/src/asm_pass4.py b/src/asm_pass4.py
index 462707a..bb30eb7 100644
--- a/src/asm_pass4.py
+++ b/src/asm_pass4.py
@@ -21,8 +21,8 @@ class Result4:
def translate_opt(arg: str, options: Options) -> str:
match arg:
- case "__STACK_POINTER": return options.stack_pointer
- case "__CSR_SCRATCH": return options.csr_scratch
+ case "SP": return options.stack_pointer
+ case "CSR_SCRATCH": return options.csr_scratch
case _: return arg
def pass_4(result3: Result3) -> Result4:
diff --git a/src/ast_compiler.py b/src/ast_compiler.py
index bcd06bf..929b979 100644
--- a/src/ast_compiler.py
+++ b/src/ast_compiler.py
@@ -7,6 +7,7 @@ import ast
from dataclasses import dataclass
from typing import Literal, Set, cast, Any, TypeAlias
from collections import OrderedDict
+from tokenizer import tokenize
if len(sys.argv) != 3:
print("usage: ast_compiler.py <infile.py> <outfile.atk16>")
@@ -20,7 +21,7 @@ Label = str
RegChar = Literal["A", "B", "C", "D", "E", "F", "G", "H"]
@dataclass
-class Reg:
+class Reg():
reg: RegChar
def __str__(self):
@@ -47,7 +48,10 @@ class Compiler(ast.NodeVisitor):
self.program_asm: list[str] = [
"@label main"
]
- self.const_bindings: dict[str, Label] = {}
+ self.function_def_asms: list[str] = []
+ self.currently_emitting_asm_list = self.program_asm
+ self.stack_pointer_offset: int | None = None
+ self.bindings: dict[str, Label] = {}
self.call_depth: int = 0
self.unique_name_counter = 0
self.latest_break_target: Label | None = None
@@ -62,13 +66,13 @@ class Compiler(ast.NodeVisitor):
def assign_const(self, name: str, value: int):
self.const_asm.append(f"@label {name}")
self.const_asm.append(f" {value}")
- self.const_bindings[name] = name
+ self.bindings[name] = name
def emit(self, asm: str):
asm = asm.strip()
asm = format_asm_row(asm)
- self.program_asm.append(asm)
+ self.currently_emitting_asm_list.append(asm)
def alloc_reg(self) -> Reg:
for reg in GENERIC_REGS:
@@ -81,6 +85,20 @@ class Compiler(ast.NodeVisitor):
def free_reg(self, reg: Reg):
self.reserved_regs.pop(reg.reg)
+ class RegContextManager:
+ def __init__(self, compiler):
+ self.compiler = compiler
+
+ def __enter__(self):
+ self.reg = self.compiler.alloc_reg()
+ return self.reg
+
+ def __exit__(self, exc_type, exc_value, exc_tb):
+ self.compiler.free_reg(self.reg)
+
+ def allocated_reg(self):
+ return Compiler.RegContextManager(self)
+
def compile(self, bootstrap_asm: str, source: str) -> str:
tree = ast.parse(source)
print(ast.dump(tree, indent=4))
@@ -90,9 +108,15 @@ class Compiler(ast.NodeVisitor):
"",
"\n".join(self.const_asm),
"",
- "\n".join(self.program_asm)
+ "\n".join(self.function_def_asms),
+ "",
+ "\n".join(self.program_asm),
+ ""
])
+ def generic_visit(self, node: ast.AST) -> Any:
+ raise NotImplementedError(f"type {type(node)}, value: {node}")
+
def emit_builtin_call(self, name: str, args: list[ast.expr]):
self.emit(f"; Builtin call {name} {args}")
match name:
@@ -110,31 +134,46 @@ class Compiler(ast.NodeVisitor):
for arg in args:
self.visit(arg)
- arg1 = self.alloc_reg()
- arg2 = self.alloc_reg()
- self.emit(f"spo {arg2}")
- self.emit(f"spo {arg1}")
- self.emit(f"str {arg2} {arg1}")
- self.free_reg(arg1)
- self.free_reg(arg2)
+ with self.allocated_reg() as arg1, self.allocated_reg() as arg2:
+ self.emit(f"spo {arg2}")
+ self.emit(f"spo {arg1}")
+ self.emit(f"str {arg2} {arg1}")
def eval_int_constant_and_spu(self, value: int):
- reg = self.alloc_reg()
- if value >= 0 and value < 8:
- self.emit(f"ldi {value} {reg}")
- else:
- name = self.get_unique_name("int")
- self.assign_const(name, value)
- self.emit(f"ldi {name} {reg}")
- self.emit(f"ldr {reg} {reg}")
+ with self.allocated_reg() as reg:
+ if value >= 0 and value < 8:
+ self.emit(f"ldi {value} {reg}")
+ else:
+ name = self.get_unique_name("int")
+ self.assign_const(name, value)
+ self.emit(f"ldi {name} {reg}")
+ self.emit(f"ldr {reg} {reg}")
- self.emit(f"spu {reg}")
- self.free_reg(reg)
+ self.emit(f"spu {reg}")
def visit_Module(self, node: ast.Module):
- for stmt in node.body:
+ stmts = node.body
+ stack_frame = self.collect_local_variables(stmts)
+
+ self.emit(f"; stack frame: {stack_frame}")
+ # Move stack pointer to accommodate local variables
+ self.stack_pointer_offset = len(stack_frame)
+
+ self.emit(f"addi SP {self.stack_pointer_offset} SP")
+
+ prev_bindings = self.bindings
+ self.bindings: dict[str, Label] = prev_bindings.copy()
+ for idx, name in enumerate(stack_frame):
+ offset = len(stack_frame) - idx
+ offset_label = f"${{SP - {offset}}}"
+ self.bindings[name] = offset_label
+
+ for stmt in stmts:
self.visit(stmt)
+ self.bindings = prev_bindings
+ self.emit(f"subi SP {self.stack_pointer_offset} SP")
+
def visit_Expr(self, expr: ast.Expr):
# if self.call_depth == 0 and type(expr.value) != ast.Call:
# self.emit("; NOP top-level expression")
@@ -147,71 +186,66 @@ class Compiler(ast.NodeVisitor):
self.visit(node.operand)
match node.op:
case ast.Not():
- reg1 = self.alloc_reg()
- reg2 = self.alloc_reg()
- self.emit(f"spo {reg1}")
- self.emit(f"ldi 1 {reg2}")
- self.emit(f"andi {reg1} 1 {reg1}")
- self.emit(f"xor {reg1} {reg2} {reg1}")
- self.emit(f"spu {reg1}")
- self.free_reg(reg1)
- self.free_reg(reg2)
+ with self.allocated_reg() as reg1, self.allocated_reg() as reg2:
+ self.emit(f"spo {reg1}")
+ self.emit(f"ldi 1 {reg2}")
+ self.emit(f"andi {reg1} 1 {reg1}")
+ self.emit(f"xor {reg1} {reg2} {reg1}")
+ self.emit(f"spu {reg1}")
+
case ast.Invert(): # aka bitwise not
- reg1 = self.alloc_reg()
- self.emit(f"spo {reg1}")
- self.emit(f"not {reg1}")
- self.emit(f"spu {reg1}")
- self.free_reg(reg1)
+ with self.allocated_reg() as reg1:
+ self.emit(f"spo {reg1}")
+ self.emit(f"not {reg1}")
+ self.emit(f"spu {reg1}")
+
case ast.UAdd(): # +a
pass # nop
+
case ast.USub(): # -a
- reg1 = self.alloc_reg()
- self.emit(f"spo {reg1}")
- self.emit(f"not {reg1}")
- self.emit(f"addi {reg1} 1 {reg1}")
- self.emit(f"spu {reg1}")
- self.free_reg(reg1)
+ with self.allocated_reg() as reg1:
+ self.emit(f"spo {reg1}")
+ self.emit(f"not {reg1}")
+ self.emit(f"addi {reg1} 1 {reg1}")
+ self.emit(f"spu {reg1}")
+
case other:
raise NotImplementedError(f"Unhandled UnaryOp: {other}")
def visit_BoolOp(self, node: ast.BoolOp):
- self.emit(f"; {node}")
+ self.emit(f"; {node.op}")
match node.op:
case ast.And():
- self.emit("; Boolean and")
- reg1 = self.alloc_reg()
label_short_circuit = self.get_unique_name("And_short_circuit")
- for arg in node.values:
- self.emit(f"; And operand {arg}")
- self.visit(arg)
+ with self.allocated_reg() as reg:
+ for arg in node.values:
+ self.emit(f"; And operand {arg}")
+ self.visit(arg)
- self.emit(f"spo {reg1}")
- self.emit(f"addi {reg1} 0 {reg1}")
- self.emit(f"bri zero {label_short_circuit}")
+ self.emit(f"spo {reg}")
+ self.emit(f"addi {reg} 0 {reg}")
+ self.emit(f"bri zero {label_short_circuit}")
+
+ self.emit(f"@label {label_short_circuit}")
+ self.emit(f"spu {reg}")
- self.emit(f"@label {label_short_circuit}")
- self.emit(f"spu {reg1}")
- self.free_reg(reg1)
case ast.Or():
- self.emit("; Boolean or")
- reg1 = self.alloc_reg()
- reg2 = self.alloc_reg()
label_short_circuit = self.get_unique_name("Or_short_circuit")
- for arg in node.values:
- self.emit(f"; Or operand {arg}")
- self.visit(arg)
+ with self.allocated_reg() as reg1, self.allocated_reg() as reg2:
+ for arg in node.values:
+ self.emit(f"; Or operand {arg}")
+ self.visit(arg)
- self.emit(f"spo {reg1}")
- self.emit(f"subi {reg1} 1 {reg2}")
- self.emit(f"bri carry {label_short_circuit}")
+ self.emit(f"spo {reg1}")
+ self.emit(f"subi {reg1} 1 {reg2}")
+ self.emit(f"bri carry {label_short_circuit}")
+
+ self.emit(f"@label {label_short_circuit}")
+ self.emit(f"spu {reg1}")
- self.emit(f"@label {label_short_circuit}")
- self.emit(f"spu {reg1}")
- self.free_reg(reg1)
- self.free_reg(reg2)
case other:
raise NotImplementedError(f"Unhandled BoolOp: {other}")
@@ -223,34 +257,29 @@ class Compiler(ast.NodeVisitor):
self.emit(f"; BinOp rhs {node}")
self.visit(node.right)
- reg1 = self.alloc_reg()
- reg2 = self.alloc_reg()
-
- self.emit(f"spo {reg2}")
- self.emit(f"spo {reg1}")
-
- match node.op:
- case ast.Add():
- self.emit(f"add {reg1} {reg2} {reg1}")
- case ast.Sub():
- self.emit(f"sub {reg1} {reg2} {reg1}")
- case ast.BitAnd():
- self.emit(f"and {reg1} {reg2} {reg1}")
- case ast.BitOr():
- self.emit(f"or {reg1} {reg2} {reg1}")
- case ast.BitXor():
- self.emit(f"xor {reg1} {reg2} {reg1}")
- case ast.LShift():
- self.emit(f"sll {reg1} {reg2} {reg1}")
- case ast.RShift():
- self.emit(f"slr {reg1} {reg2} {reg1}")
- case other:
- raise NotImplementedError(f"Unhandled BinOp: {other}")
+ with self.allocated_reg() as reg1, self.allocated_reg() as reg2:
+ self.emit(f"spo {reg2}")
+ self.emit(f"spo {reg1}")
- self.emit(f"spu {reg1}")
+ match node.op:
+ case ast.Add():
+ self.emit(f"add {reg1} {reg2} {reg1}")
+ case ast.Sub():
+ self.emit(f"sub {reg1} {reg2} {reg1}")
+ case ast.BitAnd():
+ self.emit(f"and {reg1} {reg2} {reg1}")
+ case ast.BitOr():
+ self.emit(f"or {reg1} {reg2} {reg1}")
+ case ast.BitXor():
+ self.emit(f"xor {reg1} {reg2} {reg1}")
+ case ast.LShift():
+ self.emit(f"sll {reg1} {reg2} {reg1}")
+ case ast.RShift():
+ self.emit(f"slr {reg1} {reg2} {reg1}")
+ case other:
+ raise NotImplementedError(f"Unhandled BinOp: {other}")
- self.free_reg(reg1)
- self.free_reg(reg2)
+ self.emit(f"spu {reg1}")
def visit_Constant(self, node: ast.Constant):
self.emit(f"; {node}")
@@ -274,13 +303,13 @@ class Compiler(ast.NodeVisitor):
self.emit(f"; {node}")
self.visit(node.test)
- reg1 = self.alloc_reg()
label_false = self.get_unique_name("If_false_branch")
label_end = self.get_unique_name("If_end_branch")
- self.emit(f"spo {reg1}")
- self.emit(f"addi {reg1} 0 {reg1}")
- self.emit(f"bri zero {label_false}")
- self.free_reg(reg1)
+
+ with self.allocated_reg() as reg:
+ self.emit(f"spo {reg}")
+ self.emit(f"addi {reg} 0 {reg}")
+ self.emit(f"bri zero {label_false}")
for true_branch_stmt in node.body:
self.visit(true_branch_stmt)
@@ -296,7 +325,6 @@ class Compiler(ast.NodeVisitor):
def visit_While(self, node: ast.While):
self.emit(f"; {node}")
- reg1 = self.alloc_reg()
label_test = self.get_unique_name("While_test")
label_else = self.get_unique_name("While_else")
label_end = self.get_unique_name("While_end")
@@ -304,11 +332,12 @@ class Compiler(ast.NodeVisitor):
prev_break_target = self.latest_break_target
self.latest_break_target = label_end
- self.emit(f"@label {label_test}")
- self.visit(node.test)
- self.emit(f"spo {reg1}")
- self.emit(f"addi {reg1} 0 {reg1}")
- self.emit(f"bri zero {label_else}")
+ with self.allocated_reg() as reg:
+ self.emit(f"@label {label_test}")
+ self.visit(node.test)
+ self.emit(f"spo {reg}")
+ self.emit(f"addi {reg} 0 {reg}")
+ self.emit(f"bri zero {label_else}")
for body_stmt in node.body:
self.visit(body_stmt)
@@ -322,7 +351,6 @@ class Compiler(ast.NodeVisitor):
self.emit(f"@label {label_end}")
- self.free_reg(reg1)
self.latest_break_target = prev_break_target
def visit_Break(self, node: ast.Break):
@@ -333,8 +361,86 @@ class Compiler(ast.NodeVisitor):
self.emit(f"jpi {self.latest_break_target}")
+ def visit_For(self, node: ast.For):
+ raise Exception("For loops not supported. Consider using a while loop instead.")
+
+ def visit_Lambda(self, node: ast.Lambda):
+ raise Exception("Lambda functions not supported. Consider using a named function instead.")
+
+ def collect_local_variables(self, stmts: list[ast.stmt]):
+ symbols: list[str] = []
+ for stmt in stmts:
+ match stmt:
+ case ast.Assign(targets=[ast.Name(name)]):
+ symbols.append(name)
+ case ast.Assign(other):
+ raise Exception(f"Unsupported assignment in function definition body: {other}")
+ case ast.While(body=body):
+ symbols += self.collect_local_variables(body)
+ case ast.If(body=tb, orelse=fb):
+ symbols += self.collect_local_variables(tb)
+ symbols += self.collect_local_variables(fb)
+ case other: pass
+
+ # remove duplicates
+ return list(set(symbols))
+
+ def visit_FunctionDef(self, node: ast.FunctionDef):
+ self.currently_emitting_asm_list = self.function_def_asms
+ self.emit(f"; {node}")
+
+ fn_name = node.name
+ fn_params = [str(param.arg) for param in node.args.args]
+ fn_stmts = node.body
+
+ # Add a "return None" to the end to make sure there the function returns
+ if len(fn_stmts) == 0 or type(fn_stmts[len(fn_stmts) - 1]) != ast.Return:
+ fn_stmts.append(ast.Return(value=None))
+
+ if len(node.args.kwonlyargs) > 0 or len(node.args.posonlyargs) > 0 or len(node.args.kw_defaults) > 0 or len(node.args.defaults) > 0:
+ raise Exception("Only simple positional args are supported for now in function definitions.")
+
+ self.emit(f"@label {fn_name}")
+ stack_frame = fn_params + self.collect_local_variables(fn_stmts)
+ self.emit(f"; stack frame: {stack_frame}")
+ # Move stack pointer to accommodate local variables
+ self.stack_pointer_offset = len(stack_frame) - len(fn_params)
+
+ self.emit(f"addi SP {self.stack_pointer_offset} SP")
+
+ prev_bindings = self.bindings
+ self.bindings: dict[str, Label] = prev_bindings.copy()
+ for idx, name in enumerate(stack_frame):
+ offset = len(stack_frame) - idx
+ offset_label = f"${{SP - {offset}}}"
+ self.bindings[name] = offset_label
+
+ for stmt in fn_stmts:
+ self.visit(stmt)
+
+ self.bindings = prev_bindings
+ self.currently_emitting_asm_list = self.program_asm
+
+ def visit_Return(self, node: ast.Return):
+ self.emit(f"; {node}")
+ if self.stack_pointer_offset is None:
+ raise Exception("Encountered Return outside a function def context")
+
+ ret = node.value
+
+ if ret is not None:
+ self.visit(ret)
+ else:
+ with self.allocated_reg() as reg:
+ self.emit(f"ldi 0 {reg}")
+ self.emit(f"spu {reg}")
+
+ self.emit(f"subi SP {self.stack_pointer_offset} SP")
+ self.emit("rsr")
+
def visit_Call(self, node: ast.Call):
self.emit(f"; {node}")
+
match node.func:
case ast.Attribute(ast.Name(id="atk16"), attr):
self.emit_builtin_call(attr, node.args)
@@ -343,16 +449,40 @@ class Compiler(ast.NodeVisitor):
for arg in node.args:
self.visit(arg)
- addr = self.const_bindings[name]
- self.emit(f"csi {addr}")
+ self.emit(f"csi {name}")
case other:
raise NotImplementedError(f"Unhandled Call: {other}")
- def visit_Assign(self, node: ast.Assign) -> Any:
- raise NotImplementedError(f"TODO assign {node}")
+ def visit_Name(self, node: ast.Name):
+ self.emit(f"; {node} ({node.id})")
- def generic_visit(self, node: ast.AST) -> Any:
- raise NotImplementedError(f"type {type(node)}, value: {node}")
+ name = node.id
+ if name not in self.bindings:
+ raise Exception(f"{name} is unbound")
+
+ addr = self.bindings[name]
+ with self.allocated_reg() as reg:
+ self.emit(f"ldi {addr} {reg}")
+ self.emit(f"ldr {reg} {reg}")
+ self.emit(f"spu {reg}")
+
+ def visit_Assign(self, node: ast.Assign):
+ targets = node.targets
+ value = node.value
+ match targets:
+ case [ast.Name(name)]:
+ if name not in self.bindings:
+ raise Exception(f"{name} is not in bindings. Bindings:\n{self.bindings}")
+
+ self.visit(value)
+ addr = self.bindings[name]
+
+ with self.allocated_reg() as reg1, self.allocated_reg() as reg2:
+ self.emit(f"spo {reg1}")
+ self.emit(f"ldi {addr} {reg2}")
+ self.emit(f"str {reg1} {reg2}")
+ case other:
+ raise Exception(f"Unsupported assign targets: {other}")
def visit_Import(self, node: ast.Import) -> Any:
match node.names:
@@ -385,7 +515,7 @@ class Optimizer:
asm = [row.strip() for row in asm]
asm = [self.strip_comment(row) for row in asm]
asm = [row for row in asm if not len(row) == 0]
- asm = [row.split() for row in asm]
+ asm = [tokenize(row, retain_curlies=True) for row in asm]
asm = self.compact_spu_spo_pattern(asm)
asm = self.compact_target_mov_pattern(asm)
asm = self.compact_target_mov_pattern(asm)
diff --git a/src/tokenizer.py b/src/tokenizer.py
index d89309f..cf31b56 100644
--- a/src/tokenizer.py
+++ b/src/tokenizer.py
@@ -1,4 +1,4 @@
-def tokenize(line: str) -> list[str]:
+def tokenize(line: str, retain_curlies = False) -> list[str]:
cur: str = ""
result: list[str] = []
is_py_expr = False
@@ -8,6 +8,8 @@ def tokenize(line: str) -> list[str]:
c = line[idx]
if not is_py_expr and c == "$":
is_py_expr = True
+ if retain_curlies:
+ cur += "${"
idx += 1
elif is_py_expr and c == "$":
raise Exception("Unexpected start of python expr while already parsing a python expression:\n" + line)
@@ -15,6 +17,8 @@ def tokenize(line: str) -> list[str]:
raise Exception("Unexpected end of python expr while not parsing a python expression:\n" + line)
elif is_py_expr and c == "}":
is_py_expr = False
+ if retain_curlies:
+ cur += c
result.append(cur)
cur = ""
elif is_py_expr: