From f7dd9705d44b9cee5103b267797c0449d42319df Mon Sep 17 00:00:00 2001 From: Jan Tuomi Date: Sat, 11 Nov 2023 14:40:25 +0200 Subject: WIP functions --- asm/ast_compiler_bootstrap.atk16 | 2 +- asm/bootstrap.atk16 | 2 +- asm/ext_std.py | 28 +-- asm/test_py_src.atk16 | 74 ++++---- asm/test_py_src.atk16_optimized | 42 ++--- asm/test_py_src.py | 14 +- src/asm_pass4.py | 4 +- src/ast_compiler.py | 360 ++++++++++++++++++++++++++------------- src/tokenizer.py | 6 +- 9 files changed, 337 insertions(+), 195 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 -; -; - ldi 1 RA +; +@label func +; stack frame: ['arg1', 'arg2', 'a'] + addi SP 1 SP +; (arg1) + ldi ${SP - 3} RA + ldr RA RA + spu RA + spo RA + ldi ${SP - 1} RB + str RA RB +; +; +; BinOp lhs +; (a) + ldi ${SP - 1} RA + ldr RA RA spu RA +; BinOp rhs +; (a) + ldi ${SP - 1} RA + ldr RA RA + spu RA + spo RB spo RA - addi RA 0 RA - bri zero If_false_branch_0 -; - ldi 2 RA + add RA RB RA spu RA - jpi If_end_branch_1 -@label If_false_branch_0 -; + subi SP 1 SP + rsr + +@label main +; stack frame: ['b'] + addi SP 1 SP +; +; + ldi 1 RA + spu RA +; ldi 3 RA spu RA -@label If_end_branch_1 -; -@label While_test_2 -; - ldi 3 RB - spu RB + csi func spo 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 + 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 ") @@ -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}") - - 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 {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}") 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: -- cgit v1.3