aboutsummaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
Diffstat (limited to 'src')
-rw-r--r--src/asm_pass4.py4
-rw-r--r--src/ast_compiler.py356
-rw-r--r--src/tokenizer.py6
3 files changed, 250 insertions, 116 deletions
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: