Skip to content
Merged
2 changes: 1 addition & 1 deletion spy/analyze/importing.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,7 +23,7 @@
MODULE = Union[ast.Module, "W_Module", None]

# Cache version: increment this when ast.Module or SymTable structure changes
SPYC_VERSION = 8
SPYC_VERSION = 9


@dataclass
Expand Down
19 changes: 19 additions & 0 deletions spy/analyze/scope.py
Original file line number Diff line number Diff line change
Expand Up @@ -386,6 +386,25 @@ def declare_AugAssign(self, augassign: ast.AugAssign) -> None:
self._promote_const_to_var_maybe(augassign.target)
self.declare(augassign.value)

def _declare_hidden_local(self, name: str, value: ast.Expr) -> None:
self.define_name(name, "var", "auto", value.loc, value.loc)

def declare_AugSetAttr(self, augsetattr: ast.AugSetAttr) -> None:
target_name = f"_$aug_target{augsetattr.seq}"
self._declare_hidden_local(target_name, augsetattr.target)
self.declare(augsetattr.target)
self.declare(augsetattr.value)

def declare_AugSetItem(self, augsetitem: ast.AugSetItem) -> None:
target_name = f"_$aug_target{augsetitem.seq}"
self._declare_hidden_local(target_name, augsetitem.target)
self.declare(augsetitem.target)
for i, arg in enumerate(augsetitem.args):
arg_name = f"_$aug_arg{augsetitem.seq}_{i}"
self._declare_hidden_local(arg_name, arg)
self.declare(arg)
self.declare(augsetitem.value)

def declare_AssignExpr(self, assignexpr: ast.AssignExpr) -> None:
self._declare_target_maybe(assignexpr.target, assignexpr.value)
self.declare(assignexpr.value)
Expand Down
18 changes: 18 additions & 0 deletions spy/ast.py
Original file line number Diff line number Diff line change
Expand Up @@ -949,13 +949,31 @@ class SetAttr(Stmt):
value: Expr


@astnode("parsed")
class AugSetAttr(Stmt):
seq: int # unique id within a funcdef
target: Expr
attr: StrLiteral
op: str
value: Expr


@astnode
class SetItem(Stmt):
target: Expr
args: list[Expr]
value: Expr


@astnode("parsed")
class AugSetItem(Stmt):
seq: int # unique id within a funcdef
target: Expr
args: list[Expr]
op: str
value: Expr


@astnode
class If(Stmt):
test: Expr
Expand Down
77 changes: 77 additions & 0 deletions spy/astcompile.py
Original file line number Diff line number Diff line change
Expand Up @@ -320,6 +320,83 @@ def compile_stmt_AugAssign(self, stmt: ast.AugAssign) -> list[ast.Stmt]:
)
return self.compile_stmt(desugared)

def compile_stmt_AugSetAttr(self, stmt: ast.AugSetAttr) -> list[ast.Stmt]:
target_loc = stmt.target.loc.replace(colorize=False)
value_loc = stmt.loc.replace(colorize=False)
target_name = f"_$aug_target{stmt.seq}"
target = ast.SingleTarget(
target_loc,
ast.StrLiteral(target_loc, target_name),
)
desugared: list[ast.Stmt] = [
ast.Assign(loc=target_loc, target=target, value=stmt.target),
ast.SetAttr(
loc=stmt.loc,
target=ast.Name(loc=target_loc, id=target_name),
attr=stmt.attr,
value=ast.BinOp(
loc=value_loc,
op=stmt.op,
left=ast.GetAttr(
loc=target_loc,
value=ast.Name(loc=target_loc, id=target_name),
attr=stmt.attr,
),
right=stmt.value,
),
),
]
return self.compile_body(desugared)

def compile_stmt_AugSetItem(self, stmt: ast.AugSetItem) -> list[ast.Stmt]:
target_loc = stmt.target.loc.replace(colorize=False)
value_loc = stmt.loc.replace(colorize=False)
target_name = f"_$aug_target{stmt.seq}"
target = ast.SingleTarget(
target_loc,
ast.StrLiteral(target_loc, target_name),
)
desugared: list[ast.Stmt] = [
ast.Assign(loc=target_loc, target=target, value=stmt.target)
]
arg_names = []
for i, arg in enumerate(stmt.args):
arg_loc = arg.loc.replace(colorize=False)
arg_name = f"_$aug_arg{stmt.seq}_{i}"
arg_names.append((arg_name, arg_loc))
desugared.append(
ast.Assign(
loc=arg_loc,
target=ast.SingleTarget(
arg_loc,
ast.StrLiteral(arg_loc, arg_name),
),
value=arg,
)
)

def make_args() -> list[ast.Expr]:
return [ast.Name(loc=loc, id=name) for name, loc in arg_names]

desugared.append(
ast.SetItem(
loc=stmt.loc,
target=ast.Name(loc=target_loc, id=target_name),
args=make_args(),
value=ast.BinOp(
loc=value_loc,
op=stmt.op,
left=ast.GetItem(
loc=target_loc,
value=ast.Name(loc=target_loc, id=target_name),
args=make_args(),
),
right=stmt.value,
),
)
)
return self.compile_body(desugared)

def compile_stmt_SetItem(self, stmt: ast.SetItem) -> list[ast.Stmt]:
return [
stmt.replace(
Expand Down
13 changes: 13 additions & 0 deletions spy/backend/spy.py
Original file line number Diff line number Diff line change
Expand Up @@ -355,13 +355,26 @@ def emit_stmt_SetAttr(self, node: ast.SetAttr) -> None:
v = self.fmt_expr(node.value)
self.wl(f"{t}.{a} = {v}")

def emit_stmt_AugSetAttr(self, node: ast.AugSetAttr) -> None:
t = self.fmt_expr(node.target)
a = node.attr.value
v = self.fmt_expr(node.value)
self.wl(f"{t}.{a} {node.op}= {v}")

def emit_stmt_SetItem(self, node: ast.SetItem) -> None:
t = self.fmt_expr(node.target)
arglist = [self.fmt_expr(arg) for arg in node.args]
args = ", ".join(arglist)
v = self.fmt_expr(node.value)
self.wl(f"{t}[{args}] = {v}")

def emit_stmt_AugSetItem(self, node: ast.AugSetItem) -> None:
t = self.fmt_expr(node.target)
arglist = [self.fmt_expr(arg) for arg in node.args]
args = ", ".join(arglist)
v = self.fmt_expr(node.value)
self.wl(f"{t}[{args}] {node.op}= {v}")

def emit_stmt_VarDef(self, vardef: ast.VarDef) -> None:
varname = vardef.name.value
is_auto = isinstance(vardef.type, ast.Auto)
Expand Down
48 changes: 43 additions & 5 deletions spy/parser.py
Original file line number Diff line number Diff line change
Expand Up @@ -78,11 +78,13 @@ class Parser:
src: str
filename: str
for_loop_seq: int # counter for for loops within the current function
augassign_seq: int # counter for complex augassigns within the current function

def __init__(self, src: str, filename: str) -> None:
self.src = src
self.filename = filename
self.for_loop_seq = 0
self.augassign_seq = 0

@classmethod
def from_filename(cls, filename: str) -> "Parser":
Expand Down Expand Up @@ -314,9 +316,11 @@ def _parse_py_funcdef(
# by doing this "saved_seq" dance, we ensure that nested functions "continue"
# the numbering from the their parent, but sibling functions reset the
# numbering. See test_scope::test_for_loop_nested_funcs
saved_seq = self.for_loop_seq
saved_for_loop_seq = self.for_loop_seq
saved_augassign_seq = self.augassign_seq
body = self.from_py_body(py_body)
self.for_loop_seq = saved_seq
self.for_loop_seq = saved_for_loop_seq
self.augassign_seq = saved_augassign_seq

return spy.ast.FuncDef(
loc=py_funcdef.loc,
Expand Down Expand Up @@ -625,17 +629,51 @@ def from_py_stmt_Assign(self, py_node: py_ast.Assign) -> spy.ast.Stmt:
else:
self.unsupported(py_target, "assign to complex expressions")

def from_py_stmt_AugAssign(self, py_node: py_ast.AugAssign) -> spy.ast.AugAssign:
def from_py_stmt_AugAssign(self, py_node: py_ast.AugAssign) -> spy.ast.Stmt:
py_target = py_node.target
opname = type(py_node.op).__name__
op = self._binops[opname]

if isinstance(py_target, py_ast.Name):
opname = type(py_node.op).__name__
op = self._binops[opname]
# Simple case: x += 1
return spy.ast.AugAssign(
loc=py_node.loc,
op=op,
target=spy.ast.StrLiteral(py_target.loc, py_target.id),
value=self.from_py_expr(py_node.value),
)
elif isinstance(py_target, py_ast.Attribute):
# Attribute access: a.b += 1
seq = self.augassign_seq
self.augassign_seq += 1
return spy.ast.AugSetAttr(
loc=py_node.loc,
seq=seq,
op=op,
target=self.from_py_expr(py_target.value),
attr=spy.ast.StrLiteral(py_target.loc, py_target.attr),
value=self.from_py_expr(py_node.value),
)
elif isinstance(py_target, py_ast.Subscript):
# Subscript access: arr[i] += 1
seq = self.augassign_seq
self.augassign_seq += 1
target = self.from_py_expr(py_target.value)
index = self.from_py_expr(py_target.slice)

if isinstance(index, spy.ast.Tuple):
args = index.items
else:
args = [index]

return spy.ast.AugSetItem(
loc=py_node.loc,
seq=seq,
op=op,
target=target,
args=args,
value=self.from_py_expr(py_node.value),
)
else:
self.unsupported(py_target, "assign to complex expressions")

Expand Down
80 changes: 80 additions & 0 deletions spy/tests/compiler/test_basic.py
Original file line number Diff line number Diff line change
Expand Up @@ -891,6 +891,86 @@ def foo(x: i32) -> i32:
""")
assert mod.foo(10) == ((10 + 1) * 2) - 3

@no_C
def test_aug_assign_subscript(self):
mod = self.compile("""
from __spy__ import interp_list

var count: i32 = 0

def get_count() -> i32:
return count

def idx() -> i32:
count = count + 1
return 0

def test_add() -> i32:
arr = interp_list[i32](10)
arr[idx()] += 5
return arr[0]

def test_mul() -> i32:
arr = interp_list[i32](10)
arr[idx()] *= 3
return arr[0]

def test_sub() -> i32:
arr = interp_list[i32](10)
arr[idx()] -= 3
return arr[0]
""")
assert mod.test_add() == 15
assert mod.get_count() == 1
assert mod.test_mul() == 30
assert mod.get_count() == 2
assert mod.test_sub() == 7
assert mod.get_count() == 3

def test_aug_assign_attribute(self):
mod = self.compile("""
from unsafe import raw_alloc, raw_ptr

@struct
class Box:
value: i32

def test_add() -> i32:
b = raw_alloc[Box](1)
setattr(b, 'value', 10)
b.value += 5
return b.value

def test_sub() -> i32:
b = raw_alloc[Box](1)
setattr(b, 'value', 20)
b.value -= 3
return b.value

def test_mul() -> i32:
b = raw_alloc[Box](1)
setattr(b, 'value', 10)
b.value *= 2
return b.value

def test_div() -> i32:
b = raw_alloc[Box](1)
setattr(b, 'value', 20)
b.value //= 2
return b.value

def test_mod() -> i32:
b = raw_alloc[Box](1)
setattr(b, 'value', 10)
b.value %= 3
return b.value
""")
assert mod.test_add() == 15
assert mod.test_sub() == 17
assert mod.test_mul() == 20
assert mod.test_div() == 10
assert mod.test_mod() == 1

def test_resolve_name(self):
mod = self.compile("""
from builtins import i32 as my_int
Expand Down
38 changes: 38 additions & 0 deletions spy/tests/test_astcompile.py
Original file line number Diff line number Diff line change
Expand Up @@ -177,3 +177,41 @@ def foo() -> i32:
return x
"""
self.assert_dump_decls(expected)

def test_augassign(self):
self.compile_src("""
def foo(x: i32) -> i32:
x += 1
return x
""")
expected = """
def foo(x: i32) -> i32:
x = x + 1
return x
"""
self.assert_dump(expected)

def test_augsetitem(self):
self.compile_src("""
def foo(items: dynamic, idx: i32, val: i32) -> None:
items[idx] += val
""")
expected = """
def foo(items: dynamic, idx: i32, val: i32) -> None:
_$aug_target0 = items
_$aug_arg0_0 = idx
_$aug_target0[_$aug_arg0_0] = _$aug_target0[_$aug_arg0_0] + val
"""
self.assert_dump(expected)

def test_augsetattr(self):
self.compile_src("""
def foo(obj: dynamic, val: i32) -> None:
obj.x += val
""")
expected = """
def foo(obj: dynamic, val: i32) -> None:
_$aug_target0 = obj
_$aug_target0.x = _$aug_target0.x + val
"""
self.assert_dump(expected)
Loading
Loading