diff --git a/mypyc/irbuild/for_helpers.py b/mypyc/irbuild/for_helpers.py index 30101371bb3f..a97253174b97 100644 --- a/mypyc/irbuild/for_helpers.py +++ b/mypyc/irbuild/for_helpers.py @@ -70,7 +70,7 @@ ) from mypyc.irbuild.builder import IRBuilder from mypyc.irbuild.constant_fold import constant_fold_expr -from mypyc.irbuild.targets import AssignmentTarget, AssignmentTargetTuple +from mypyc.irbuild.targets import AssignmentTargetTuple from mypyc.irbuild.vec import vec_append, vec_create, vec_get_item_unsafe, vec_init_item_unsafe from mypyc.primitives.dict_ops import ( dict_check_size_op, @@ -654,6 +654,7 @@ def __init__( self.index = index self.body_block = body_block self.line = line + self.nested = nested # Some for loops need a cleanup block that we execute at exit. We # create a cleanup block if needed. However, if we are generating a for # loop for a nested iterator, such as "e" in "enumerate(e)", the @@ -946,19 +947,34 @@ def gen_condition(self) -> None: # (unless input is immutable type). len_reg = builder.read(self.length_reg, line) comparison = builder.binary_op(builder.read(self.index_target, line), len_reg, "<", line) - builder.add_bool_branch(comparison, self.body_block, self.loop_exit) + if self.nested: + # In zip() or enumerate(), read the item right after the length check, as + # the sequence's iterator would. Other iterators may be advanced and other + # loop targets assigned before our begin_body(), and either could shrink + # the sequence. + read_block = BasicBlock() + builder.add_bool_branch(comparison, read_block, self.loop_exit) + builder.activate_block(read_block) + self.next_reg = self.read_item() + builder.goto(self.body_block) + else: + builder.add_bool_branch(comparison, self.body_block, self.loop_exit) - def begin_body(self) -> None: + def read_item(self) -> Value: builder = self.builder line = self.line - # Read the next list item. - value_box = unsafe_index( + return unsafe_index( builder, builder.read(self.expr_target, line), builder.read(self.index_target, line), line, ) - assert value_box + + def begin_body(self) -> None: + builder = self.builder + line = self.line + # Read the next list item, unless gen_condition() already did. + value_box = self.next_reg if self.nested else self.read_item() # We coerce to the type of list elements here so that # iterating with tuple unpacking generates a tuple based # unpack instead of an iterator based one. @@ -1139,9 +1155,6 @@ def init(self, start_reg: Value, end_reg: Value, step: int) -> None: index_reg = Register(index_type, line=self.line) builder.assign(index_reg, start_reg, self.line) self.index_reg = index_reg - # Initialize loop index to 0. Assert that the index target is assignable. - self.index_target: Register | AssignmentTarget = builder.get_assignment_target(self.index) - builder.assign(self.index_target, builder.read(self.index_reg, self.line), self.line) def convert_arg(self, value: Value) -> Value: """Convert a range() argument to an int using __index__, like range() does.""" @@ -1175,9 +1188,13 @@ def gen_condition(self) -> None: def begin_body(self) -> None: # Update the user-visible loop variable at the start of the body, # after the condition check passes. This ensures the variable isn't - # "overshot" when the loop exits (matching CPython semantics). + # assigned if the range is empty, and isn't "overshot" when the loop + # exits (matching CPython semantics). builder = self.builder - builder.assign(self.index_target, builder.read(self.index_reg, self.line), self.line) + line = self.line + builder.assign( + builder.get_assignment_target(self.index), builder.read(self.index_reg, line), line + ) def gen_step(self) -> None: builder = self.builder @@ -1209,10 +1226,9 @@ class ForInfiniteCounter(ForGenerator): def init(self) -> None: builder = self.builder # Create a register to store the state of the loop index and - # initialize this register along with the loop index to 0. + # initialize this register to 0. zero = Integer(0) self.index_reg = builder.ensure_register(zero) - self.index_target: Register | AssignmentTarget = builder.get_assignment_target(self.index) def gen_step(self) -> None: builder = self.builder @@ -1226,8 +1242,10 @@ def gen_step(self) -> None: builder.assign(self.index_reg, new_val, line) def begin_body(self) -> None: - self.builder.assign( - self.index_target, self.builder.read(self.index_reg, self.line), self.line + builder = self.builder + line = self.line + builder.assign( + builder.get_assignment_target(self.index), builder.read(self.index_reg, line), line ) diff --git a/mypyc/test-data/irbuild-basic.test b/mypyc/test-data/irbuild-basic.test index 76663b96f77d..f22f9afff3d3 100644 --- a/mypyc/test-data/irbuild-basic.test +++ b/mypyc/test-data/irbuild-basic.test @@ -3512,14 +3512,12 @@ L5: def range_in_loop(): sum :: int r0 :: short_int - i :: int r1 :: bit - r2 :: int + i, r2 :: int r3 :: short_int L0: sum = 0 r0 = 8 - i = r0 L1: r1 = int_lt r0, 24 if r1 goto L2 else goto L4 :: bool diff --git a/mypyc/test-data/irbuild-i64.test b/mypyc/test-data/irbuild-i64.test index 26f8a2f4c495..b861406b43e7 100644 --- a/mypyc/test-data/irbuild-i64.test +++ b/mypyc/test-data/irbuild-i64.test @@ -536,13 +536,13 @@ def g(a): L0: return 1 def f(x): - x, r0, n :: i64 + x, r0 :: i64 r1 :: bit + n :: i64 r2 :: None r3 :: i64 L0: r0 = 0 - n = r0 L1: r1 = r0 < x :: signed if r1 goto L2 else goto L4 :: bool @@ -1661,12 +1661,11 @@ def f() -> None: y = x [out] def f(): - r0, x :: i64 + r0 :: i64 r1 :: bit - y, r2 :: i64 + x, y, r2 :: i64 L0: r0 = 0 - x = r0 L1: r1 = r0 < 4 :: signed if r1 goto L2 else goto L4 :: bool @@ -1688,12 +1687,11 @@ def f() -> None: y = x [out] def f(): - r0, x :: i64 + r0 :: i64 r1 :: bit - y, r2 :: i64 + x, y, r2 :: i64 L0: r0 = 0 - x = r0 L1: r1 = r0 < 4 :: signed if r1 goto L2 else goto L4 :: bool diff --git a/mypyc/test-data/irbuild-lists.test b/mypyc/test-data/irbuild-lists.test index e70d81c96bf7..ed89db8f702f 100644 --- a/mypyc/test-data/irbuild-lists.test +++ b/mypyc/test-data/irbuild-lists.test @@ -290,8 +290,8 @@ def increment(l): l :: list r0 :: native_int r1, r2 :: short_int - i :: int r3 :: bit + i :: int r4, r5, r6 :: object r7 :: bit r8 :: short_int @@ -299,7 +299,6 @@ L0: r0 = var_object_size l r1 = r0 << 1 r2 = 0 - i = r2 L1: r3 = int_lt r2, r1 if r3 goto L2 else goto L4 :: bool diff --git a/mypyc/test-data/irbuild-set.test b/mypyc/test-data/irbuild-set.test index ac01c5840111..2f5047f4bd05 100644 --- a/mypyc/test-data/irbuild-set.test +++ b/mypyc/test-data/irbuild-set.test @@ -213,9 +213,8 @@ L5: def test4(): r0 :: set r1 :: short_int - x :: int r2 :: bit - r3 :: int + x, r3 :: int r4 :: object r5 :: i32 r6 :: bit @@ -224,7 +223,6 @@ def test4(): L0: r0 = PySet_New(0) r1 = 2 - x = r1 L1: r2 = int_lt r1, 12 if r2 goto L2 else goto L4 :: bool @@ -244,9 +242,8 @@ L4: def test5(): r0 :: set r1 :: short_int - x :: int r2 :: bit - r3 :: int + x, r3 :: int r4 :: object r5 :: i32 r6 :: bit @@ -255,7 +252,6 @@ def test5(): L0: r0 = PySet_New(0) r1 = 2 - x = r1 L1: r2 = int_lt r1, 12 if r2 goto L2 else goto L4 :: bool diff --git a/mypyc/test-data/irbuild-statements.test b/mypyc/test-data/irbuild-statements.test index 7aed23a3b530..68e0e2d79f9c 100644 --- a/mypyc/test-data/irbuild-statements.test +++ b/mypyc/test-data/irbuild-statements.test @@ -7,14 +7,12 @@ def f() -> None: def f(): x :: int r0 :: short_int - i :: int r1 :: bit - r2 :: int + i, r2 :: int r3 :: short_int L0: x = 0 r0 = 0 - i = r0 L1: r1 = int_lt r0, 10 if r1 goto L2 else goto L4 :: bool @@ -35,12 +33,11 @@ def f(a: int) -> None: pass [out] def f(a): - a, r0, i :: int + a, r0 :: int r1 :: bit - r2 :: int + i, r2 :: int L0: r0 = 0 - i = r0 L1: r1 = int_lt r0, a if r1 goto L2 else goto L4 :: bool @@ -60,12 +57,11 @@ def f() -> None: [out] def f(): r0 :: short_int - i :: int r1 :: bit + i :: int r2 :: short_int L0: r0 = 20 - i = r0 L1: r1 = int_gt r0, 0 if r1 goto L2 else goto L4 :: bool @@ -89,16 +85,15 @@ def f(a, b): a, b, r0 :: object r1 :: int r2 :: object - r3, r4, i :: int + r3, r4 :: int r5 :: bit - r6 :: int + i, r6 :: int L0: r0 = PyNumber_Index(a) r1 = unbox(int, r0) r2 = PyNumber_Index(b) r3 = unbox(int, r2) r4 = r1 - i = r4 L1: r5 = int_lt r4, r3 if r5 goto L2 else goto L4 :: bool @@ -126,13 +121,12 @@ L0: return 4 def f(c): c :: __main__.C - r0, r1, i :: int + r0, r1 :: int r2 :: bit - r3 :: int + i, r3 :: int L0: r0 = c.__index__() r1 = 0 - i = r1 L1: r2 = int_lt r1, r0 if r2 goto L2 else goto L4 :: bool @@ -170,12 +164,11 @@ def f() -> None: [out] def f(): r0 :: short_int - n :: int r1 :: bit + n :: int r2 :: short_int L0: r0 = 0 - n = r0 L1: r1 = int_lt r0, 10 if r1 goto L2 else goto L4 :: bool @@ -240,12 +233,11 @@ def f() -> None: [out] def f(): r0 :: short_int - n :: int r1 :: bit + n :: int r2 :: short_int L0: r0 = 0 - n = r0 L1: r1 = int_lt r0, 10 if r1 goto L2 else goto L4 :: bool @@ -1018,9 +1010,8 @@ def f(a): r0 :: short_int r1, r2 :: native_int r3 :: bit - i :: int r4 :: object - r5, x, r6 :: int + i, r5, x, r6 :: int r7 :: short_int r8 :: native_int L0: @@ -1029,21 +1020,22 @@ L0: L1: r2 = var_object_size a r3 = r1 < r2 :: signed - if r3 goto L2 else goto L4 :: bool + if r3 goto L2 else goto L5 :: bool L2: - i = r0 r4 = list_get_item_unsafe a, r1 +L3: + i = r0 r5 = unbox(int, r4) x = r5 r6 = CPyTagged_Add(i, x) -L3: +L4: r7 = r0 + 2 r0 = r7 r8 = r1 + 1 r1 = r8 goto L1 -L4: L5: +L6: return 1 def g(x): x :: object @@ -1104,30 +1096,31 @@ L0: L1: r2 = var_object_size a r3 = r0 < r2 :: signed - if r3 goto L2 else goto L7 :: bool + if r3 goto L2 else goto L8 :: bool L2: - r4 = PyIter_Next(r1) - if is_error(r4) goto L7 else goto L3 + r4 = list_get_item_unsafe a, r0 L3: - r5 = list_get_item_unsafe a, r0 - r6 = unbox(int, r5) + r5 = PyIter_Next(r1) + if is_error(r5) goto L8 else goto L4 +L4: + r6 = unbox(int, r4) x = r6 - r7 = unbox(bool, r4) + r7 = unbox(bool, r5) y = r7 r8 = PyObject_IsTrue(b) r9 = r8 >= 0 :: signed r10 = truncate r8: i32 to builtins.bool - if r10 goto L4 else goto L5 :: bool -L4: - x = 2 + if r10 goto L5 else goto L6 :: bool L5: + x = 2 L6: +L7: r11 = r0 + 1 r0 = r11 goto L1 -L7: - r12 = CPy_NoErrOccurred() L8: + r12 = CPy_NoErrOccurred() +L9: return 1 def g(a, b): a :: object @@ -1135,13 +1128,13 @@ def g(a, b): r0 :: object r1 :: native_int r2 :: short_int - z :: int r3 :: object r4 :: native_int - r5, r6 :: bit - r7, x :: bool - r8 :: object - r9, y :: int + r5 :: bit + r6 :: object + r7 :: bit + r8, x :: bool + r9, y, z :: int r10 :: native_int r11 :: short_int r12 :: bit @@ -1149,34 +1142,34 @@ L0: r0 = PyObject_GetIter(a) r1 = 0 r2 = 0 - z = r2 L1: r3 = PyIter_Next(r0) - if is_error(r3) goto L6 else goto L2 + if is_error(r3) goto L7 else goto L2 L2: r4 = var_object_size b r5 = r1 < r4 :: signed - if r5 goto L3 else goto L6 :: bool + if r5 goto L3 else goto L7 :: bool L3: - r6 = int_lt r2, 10 - if r6 goto L4 else goto L6 :: bool + r6 = list_get_item_unsafe b, r1 L4: - r7 = unbox(bool, r3) - x = r7 - r8 = list_get_item_unsafe b, r1 - r9 = unbox(int, r8) + r7 = int_lt r2, 10 + if r7 goto L5 else goto L7 :: bool +L5: + r8 = unbox(bool, r3) + x = r8 + r9 = unbox(int, r6) y = r9 z = r2 x = 0 -L5: +L6: r10 = r1 + 1 r1 = r10 r11 = r2 + 2 r2 = r11 goto L1 -L6: - r12 = CPy_NoErrOccurred() L7: + r12 = CPy_NoErrOccurred() +L8: return 1 [case testConditionalFunctionDefinition] diff --git a/mypyc/test-data/irbuild-tuple.test b/mypyc/test-data/irbuild-tuple.test index 808b5b9e0d4a..c3d81e638d2f 100644 --- a/mypyc/test-data/irbuild-tuple.test +++ b/mypyc/test-data/irbuild-tuple.test @@ -1190,8 +1190,8 @@ def f2() -> tuple[str, ...]: def f(): r0 :: list r1 :: short_int - x :: int r2 :: bit + x :: int r3 :: str r4 :: i32 r5 :: bit @@ -1200,7 +1200,6 @@ def f(): L0: r0 = PyList_New(0) r1 = 0 - x = r1 L1: r2 = int_lt r1, 10 if r2 goto L2 else goto L4 :: bool @@ -1219,8 +1218,8 @@ L4: def f2(): r0 :: list r1 :: short_int - x :: int r2 :: bit + x :: int r3 :: str r4 :: i32 r5 :: bit @@ -1229,7 +1228,6 @@ def f2(): L0: r0 = PyList_New(0) r1 = 0 - x = r1 L1: r2 = int_lt r1, 10 if r2 goto L2 else goto L4 :: bool diff --git a/mypyc/test-data/irbuild-vec-i64.test b/mypyc/test-data/irbuild-vec-i64.test index 176e2616a485..718eabb89455 100644 --- a/mypyc/test-data/irbuild-vec-i64.test +++ b/mypyc/test-data/irbuild-vec-i64.test @@ -309,16 +309,15 @@ def f(n: i64) -> vec[i64]: def f(n): n :: i64 r0, r1 :: vec[i64] - r2, x :: i64 + r2 :: i64 r3 :: bit - r4 :: i64 + x, r4 :: i64 r5 :: vec[i64] r6 :: i64 L0: r0 = VecI64Api.alloc(0, 0) r1 = r0 r2 = 0 - x = r2 L1: r3 = r2 < 5 :: signed if r3 goto L2 else goto L4 :: bool @@ -435,28 +434,25 @@ def f() -> vec[i64]: def f(): r0, r1 :: vec[i64] r2 :: short_int - r3, ___tmp_6 :: i64 - r4 :: bit - r5 :: i64 - r6 :: vec[i64] - r7 :: short_int + r3 :: bit + r4, ___tmp_6 :: i64 + r5 :: vec[i64] + r6 :: short_int L0: r0 = VecI64Api.alloc(0, 0) r1 = r0 r2 = 0 - r3 = r2 >> 1 - ___tmp_6 = r3 L1: - r4 = int_lt r2, 14 - if r4 goto L2 else goto L4 :: bool + r3 = int_lt r2, 14 + if r3 goto L2 else goto L4 :: bool L2: - r5 = r2 >> 1 - ___tmp_6 = r5 - r6 = VecI64Api.append(r1, ___tmp_6) - r1 = r6 + r4 = r2 >> 1 + ___tmp_6 = r4 + r5 = VecI64Api.append(r1, ___tmp_6) + r1 = r5 L3: - r7 = r2 + 2 - r2 = r7 + r6 = r2 + 2 + r2 = r6 goto L1 L4: return r1 diff --git a/mypyc/test-data/irbuild-vec-t.test b/mypyc/test-data/irbuild-vec-t.test index ce162bd0806c..e720c2d14651 100644 --- a/mypyc/test-data/irbuild-vec-t.test +++ b/mypyc/test-data/irbuild-vec-t.test @@ -300,8 +300,9 @@ def f(n): r0 :: object r1 :: ptr r2, r3 :: vec[str] - r4, x :: i64 + r4 :: i64 r5 :: bit + x :: i64 r6 :: str r7 :: object r8 :: ptr @@ -313,7 +314,6 @@ L0: r2 = VecTApi.alloc(0, 0, r1) r3 = r2 r4 = 0 - x = r4 L1: r5 = r4 < 5 :: signed if r5 goto L2 else goto L4 :: bool diff --git a/mypyc/test-data/run-loops.test b/mypyc/test-data/run-loops.test index 077acafdb814..5f2f0ae52101 100644 --- a/mypyc/test-data/run-loops.test +++ b/mypyc/test-data/run-loops.test @@ -302,6 +302,221 @@ assert m == 8, f"expected 8, got {m}" from native import * [out] +[case testRangeLoopVariableEmptyRange] +from typing import Iterator +from mypy_extensions import i64 +from testutil import assertRaises + +def last_range(start: int, end: int) -> int: + i = -1 + for i in range(start, end): + pass + return i + +def unbound_range(n: int) -> int: + for i in range(n): + pass + return i + +def last_range_negative_step() -> int: + i = -1 + for i in range(0, 5, -1): + pass + return i + +def unbound_range_negative_step() -> int: + for i in range(0, 5, -1): + pass + return i + +def last_range_i64(n: i64) -> i64: + i: i64 = -1 + for i in range(n): + pass + return i + +def unbound_range_i64(n: i64) -> i64: + for i in range(n): + pass + return i + +def test_empty_range_keeps_old_value() -> None: + assert last_range(5, 0) == -1 + assert last_range(0, 3) == 2 + assert last_range_negative_step() == -1 + assert last_range_i64(0) == -1 + assert last_range_i64(3) == 2 + +def test_empty_range_leaves_variable_unbound() -> None: + assert unbound_range(3) == 2 + with assertRaises(UnboundLocalError): + unbound_range(0) + with assertRaises(UnboundLocalError): + unbound_range_negative_step() + assert unbound_range_i64(3) == 2 + with assertRaises(UnboundLocalError): + unbound_range_i64(0) + +def test_empty_range_in_enumerate_and_zip() -> None: + i = x = -1 + for i, x in enumerate(range(0)): + pass + assert (i, x) == (-1, -1) + a = b = -1 + for a, b in zip(range(3), range(0)): + pass + assert (a, b) == (-1, -1) + +def gen_range(n: int) -> Iterator[int]: + i = -1 + for i in range(n): + yield i + yield i + +def test_empty_range_in_generator() -> None: + assert list(gen_range(0)) == [-1] + assert list(gen_range(2)) == [0, 1, 1] + +for module_var in range(0): + pass + +def test_empty_range_at_module_level() -> None: + assert "module_var" not in globals() + +class A: + def __init__(self, n: int) -> None: + for self.x in range(n): + pass + +def test_empty_range_with_attribute_target() -> None: + assert A(2).x == 1 + with assertRaises(AttributeError): + A(0).x + +[case testForLoopTargetEvaluatedEachIteration] +class C: + def __init__(self) -> None: + self.assigned: list[int] = [] + + @property + def x(self) -> int: + return -1 + + @x.setter + def x(self, value: int) -> None: + self.assigned.append(value) + +def test_attribute_target() -> None: + c = C() + for c.x in range(0): + pass + assert c.assigned == [] + for c.x in range(3): + pass + assert c.assigned == [0, 1, 2] + +def index(log: list[int]) -> int: + log.append(0) + return 0 + +def test_index_target_in_range_loop() -> None: + a = [-1] + log: list[int] = [] + for a[index(log)] in range(0): + pass + assert a == [-1] + assert log == [] + for a[index(log)] in range(3): + pass + assert a == [2] + assert len(log) == 3 + +def test_index_target_in_enumerate_loop() -> None: + a = [-1] + log: list[int] = [] + empty: list[str] = [] + for a[index(log)], s in enumerate(empty): + pass + assert a == [-1] + assert log == [] + for a[index(log)], s in enumerate("abc"): + pass + assert a == [2] + assert len(log) == 3 + +[case testForZipAndEnumerateListClearedBeforeBody] +def index(items: list[int], calls: list[int]) -> int: + calls.append(1) + if len(calls) == 2: + items.clear() + return 0 + +def test_enumerate_index_target() -> None: + items = [10, 20] + target = [-1] + calls: list[int] = [] + result: list[int] = [] + for target[index(items, calls)], value in enumerate(items): + result.append(value) + assert items == [] + assert result == [10, 20] + +def test_zip_range_target() -> None: + items = [10, 20] + target = [-1] + calls: list[int] = [] + result: list[int] = [] + for target[index(items, calls)], value in zip(range(2), items): + result.append(value) + assert items == [] + assert result == [10, 20] + +def test_zip_reversed() -> None: + items = [10, 20] + target = [-1] + calls: list[int] = [] + result: list[int] = [] + for target[index(items, calls)], value in zip([1, 2], reversed(items)): + result.append(value) + assert items == [] + assert result == [20, 10] + +def test_zip_iterator_clears_list() -> None: + items = [10, 20] + + def clear_items(n: int) -> int: + if n == 2: + items.clear() + return n + + result: list[tuple[int, int]] = [] + for value, n in zip(items, map(clear_items, [1, 2, 3])): + result.append((value, n)) + assert items == [] + assert result == [(10, 1), (20, 2)] + +class C: + def __init__(self, items: list[int]) -> None: + self.items = items + + @property + def x(self) -> int: + return -1 + + @x.setter + def x(self, value: int) -> None: + if value == 1: + self.items.clear() + +def test_enumerate_index_setter_clears_list() -> None: + items = [10, 20] + c = C(items) + result: list[int] = [] + for c.x, value in enumerate(items): + result.append(value) + assert items == [] + assert result == [10, 20] + [case testForIterable] from typing import Iterable, Dict, Any, Tuple, TypeVar