diff options
Diffstat (limited to 'src')
| -rw-r--r-- | src/ast_compiler.py | 321 |
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 |
