aboutsummaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
Diffstat (limited to 'src')
-rw-r--r--src/ast_compiler.py321
1 files changed, 265 insertions, 56 deletions
diff --git a/src/ast_compiler.py b/src/ast_compiler.py
index 2208e43..bcd06bf 100644
--- a/src/ast_compiler.py
+++ b/src/ast_compiler.py
@@ -1,7 +1,6 @@
#!/usr/bin/env python3
# Generate .atk16 assembly from a subset of Python
-from _ast import AnnAssign, Expr, Module
import sys
import ast
@@ -51,6 +50,7 @@ class Compiler(ast.NodeVisitor):
self.const_bindings: dict[str, Label] = {}
self.call_depth: int = 0
self.unique_name_counter = 0
+ self.latest_break_target: Label | None = None
self.reserved_regs: OrderedDict[RegChar, None] = OrderedDict()
@@ -94,38 +94,22 @@ class Compiler(ast.NodeVisitor):
])
def emit_builtin_call(self, name: str, args: list[ast.expr]):
- self.emit(f"; builtin call {name} {args}")
+ self.emit(f"; Builtin call {name} {args}")
match name:
case "asm":
match args:
- case [ast.Constant(value)]:
- if type(value) != str:
- raise Exception(f"asm: invalid arg type: {type(value)}")
-
+ case [ast.Constant(str(value))]:
self.emit(value)
- case _: raise Exception(f"asm: invalid args: {args}")
- case "set_graphics_mode":
- match args:
- case [ast.Constant(value)]:
- if type(value) != int:
- raise Exception(f"set_graphics_mode: invalid arg type: {type(value)}")
-
- reg1 = self.alloc_reg()
- reg2 = self.alloc_reg()
- self.emit(f"ldi vt_gr_mode_addr {reg1}")
- self.emit(f"ldr {reg1} {reg1}")
- self.emit(f"ldi {value} {reg2}")
- self.emit(f"str {reg2} {reg1}")
- self.free_reg(reg1)
- self.free_reg(reg2)
- case _: raise Exception(f"set_graphics_mode: invalid args: {args}")
+ case other: raise Exception(f"asm: invalid args: {other}")
case "store":
if len(args) != 2:
raise Exception("Invalid number of arguments to store: " + str(len(args)))
- self.eval_expr_and_spu(args[0])
- self.eval_expr_and_spu(args[1])
+ # Evaluate args before call
+ for arg in args:
+ self.visit(arg)
+
arg1 = self.alloc_reg()
arg2 = self.alloc_reg()
self.emit(f"spo {arg2}")
@@ -147,43 +131,237 @@ class Compiler(ast.NodeVisitor):
self.emit(f"spu {reg}")
self.free_reg(reg)
- def eval_expr_and_spu(self, expr: ast.expr):
- match expr:
- case ast.Constant(int(value)):
+ def visit_Module(self, node: ast.Module):
+ for stmt in node.body:
+ self.visit(stmt)
+
+ def visit_Expr(self, expr: ast.Expr):
+ # if self.call_depth == 0 and type(expr.value) != ast.Call:
+ # self.emit("; NOP top-level expression")
+ # return
+
+ self.visit(expr.value)
+
+ def visit_UnaryOp(self, node: ast.UnaryOp):
+ self.emit(f"; {node}")
+ 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)
+ 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)
+ 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)
+ case other:
+ raise NotImplementedError(f"Unhandled UnaryOp: {other}")
+
+ def visit_BoolOp(self, node: ast.BoolOp):
+ self.emit(f"; {node}")
+
+ 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)
+
+ self.emit(f"spo {reg1}")
+ self.emit(f"addi {reg1} 0 {reg1}")
+ self.emit(f"bri zero {label_short_circuit}")
+
+ 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)
+
+ 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.free_reg(reg1)
+ self.free_reg(reg2)
+ case other:
+ raise NotImplementedError(f"Unhandled BoolOp: {other}")
+
+ def visit_BinOp(self, node: ast.BinOp):
+ self.emit(f"; {node}")
+
+ self.emit(f"; BinOp lhs {node}")
+ self.visit(node.left)
+ 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}")
+
+ self.emit(f"spu {reg1}")
+
+ self.free_reg(reg1)
+ self.free_reg(reg2)
+
+ def visit_Constant(self, node: ast.Constant):
+ self.emit(f"; {node}")
+ match node.value:
+ case bool(value):
+ int_value = 1 if value else 0
+ self.eval_int_constant_and_spu(int_value)
+ case int(value):
self.eval_int_constant_and_spu(value)
- case ast.Constant(str(value)):
+ case str(value):
if len(value) > 1:
raise Exception("Invalid string, only single char values allowed: " + value)
c = value[0]
int_value = ord(c)
self.eval_int_constant_and_spu(int_value)
+ case other:
+ raise NotImplementedError(f"Unhandled Constant: {other}")
+ def visit_If(self, node: ast.If):
+ self.emit(f"; {node}")
- def visit_Module(self, node: Module):
- for stmt in node.body:
- self.visit(stmt)
+ 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)
+
+ for true_branch_stmt in node.body:
+ self.visit(true_branch_stmt)
+
+ self.emit(f"jpi {label_end}")
+ self.emit(f"@label {label_false}")
+
+ for false_branch_stmt in node.orelse:
+ self.visit(false_branch_stmt)
+
+ self.emit(f"@label {label_end}")
+
+ 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")
- def visit_Expr(self, expr: Expr):
- if self.call_depth == 0 and type(expr.value) != ast.Call:
- self.emit("; NOP top-level expression")
- return
+ prev_break_target = self.latest_break_target
+ self.latest_break_target = label_end
- match expr.value:
- case ast.Call(func, args, keywords):
- match func:
- case ast.Attribute(ast.Name(id="atk16"), attr):
- print("FOUND CAPTURED ATK16 CALL: " + attr)
- self.emit_builtin_call(attr, args)
- case ast.Name(name):
- addr = self.const_bindings[name]
- self.emit(f"csi {addr}")
- case _:
- raise NotImplementedError("Unhandled Call: " + str(func))
- case _:
- raise NotImplementedError("Unhandled Expr.value: " + str(expr.value))
+ 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}")
- def visit_AnnAssign(self, node: AnnAssign):
+ for body_stmt in node.body:
+ self.visit(body_stmt)
+
+ self.emit(f"jpi {label_test}")
+
+ self.emit(f"@label {label_else}")
+
+ for else_stmt in node.orelse:
+ self.visit(else_stmt)
+
+ self.emit(f"@label {label_end}")
+
+ self.free_reg(reg1)
+ self.latest_break_target = prev_break_target
+
+ def visit_Break(self, node: ast.Break):
+ self.emit(f"; {node}")
+
+ if self.latest_break_target is None:
+ raise Exception("Invalid break: no break target defined, i.e. no place to break out to")
+
+ self.emit(f"jpi {self.latest_break_target}")
+
+ 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)
+ case ast.Name(name):
+ # Evaluate args before call
+ for arg in node.args:
+ self.visit(arg)
+
+ addr = self.const_bindings[name]
+ self.emit(f"csi {addr}")
+ case other:
+ raise NotImplementedError(f"Unhandled Call: {other}")
+
+ def visit_Assign(self, node: ast.Assign) -> Any:
+ raise NotImplementedError(f"TODO assign {node}")
+
+ def generic_visit(self, node: ast.AST) -> Any:
+ raise NotImplementedError(f"type {type(node)}, value: {node}")
+
+ def visit_Import(self, node: ast.Import) -> Any:
+ match node.names:
+ case [ast.alias(name="atk16")]:
+ return
+ case other:
+ raise NotImplementedError(f"Unsupported import {other}")
+
+ def visit_AnnAssign(self, node: ast.AnnAssign):
match node:
case ast.AnnAssign(
target=ast.Name(name),
@@ -209,9 +387,12 @@ class Optimizer:
asm = [row for row in asm if not len(row) == 0]
asm = [row.split() for row in asm]
asm = self.compact_spu_spo_pattern(asm)
- asm = self.compact_load_mov_pattern(asm)
- asm = self.compact_load_mov_pattern(asm)
+ asm = self.compact_target_mov_pattern(asm)
+ asm = self.compact_target_mov_pattern(asm)
+ asm = self.compact_mov_source_pattern(asm)
+ asm = self.compact_mov_source_pattern(asm)
asm = self.compact_spu_load_spo_pattern(asm)
+ # asm = self.convert_alr_to_ali(asm)
result = "\n".join([format_asm_row(" ".join(row)) for row in asm])
return result
@@ -252,7 +433,8 @@ class Optimizer:
return result
- def compact_load_mov_pattern(self, asm: list[list[str]]) -> list[list[str]]:
+ def compact_target_mov_pattern(self, asm: list[list[str]]) -> list[list[str]]:
+ ops_with_target = ["ldi", "ldr", "add", "sub", "addi", "subi", "and", "or", "xor", "sll", "slr", "sar", "slli", "slri", "sari", "inc", "dec", "mov", "ali", "alr", "lpc"]
i = 0
result: list[list[str]] = []
while i < len(asm):
@@ -263,12 +445,39 @@ class Optimizer:
result.append(current)
continue
- if (current[0] == "ldi" or current[0] == "ldr") and next[0] == "mov":
- load_op, load_from, load_target_reg = current[0], current[1], current[2]
+ if current[0] in ops_with_target and next[0] == "mov":
+
+ op_op, op_operands, op_target_reg = current[0], current[1:len(current) - 1], current[len(current) - 1]
mov_from_reg, mov_to_reg = next[1], next[2]
- if load_target_reg == mov_from_reg:
- ret = [load_op, load_from, mov_to_reg]
+ if op_target_reg == mov_from_reg:
+ ret = [op_op, *op_operands, mov_to_reg]
+ result.append(ret)
+ i += 1
+ continue
+
+ result.append(current)
+
+ return result
+
+ def compact_mov_source_pattern(self, asm: list[list[str]]) -> list[list[str]]:
+ ops_with_source = ["str", "ldr", "add", "sub", "addi", "subi", "and", "or", "xor", "sll", "slr", "sar", "slli", "slri", "sari", "mov", "ali", "alr"]
+ i = 0
+ result: list[list[str]] = []
+ while i < len(asm):
+ current = asm[i]
+ next = asm[i + 1] if i + 1 < len(asm) else None
+ i += 1
+ if current[0].startswith("@") or next is None:
+ result.append(current)
+ continue
+
+ if current[0] == "mov" and next[0] in ops_with_source:
+ mov_from_reg, mov_to_reg = current[1], current[2]
+ op_op, op_source_reg, op_operands = next[0], next[1], next[2:]
+
+ if op_source_reg == mov_to_reg:
+ ret = [op_op, op_source_reg, *op_operands]
result.append(ret)
i += 1
continue