diff --git a/mypyc/irbuild/for_helpers.py b/mypyc/irbuild/for_helpers.py index 30101371bb3f..f5c5ed2be452 100644 --- a/mypyc/irbuild/for_helpers.py +++ b/mypyc/irbuild/for_helpers.py @@ -746,10 +746,6 @@ def gen_cleanup(self) -> None: class ForNativeGenerator(ForGenerator): """Generate IR for a for loop over a native generator.""" - def need_cleanup(self) -> bool: - # Create a new cleanup block for when the loop is finished. - return True - def init(self, expr_reg: Value, target_type: RType) -> None: # The generator expression is also the iterator. The generator spill transform will # promote it to the private generator frame if needed. @@ -780,7 +776,16 @@ def gen_condition(self) -> None: helper_call.error_kind = ERR_NEVER self.next_reg = builder.add(helper_call) - builder.add(Branch(self.next_reg, self.loop_exit, self.body_block, Branch.IS_ERROR)) + stop_block = BasicBlock() + builder.add(Branch(self.next_reg, stop_block, self.body_block, Branch.IS_ERROR)) + + # The generator has stopped. If it didn't set the return value, it raised an + # exception that we need to propagate. Check this here rather than when the loop + # exits, since in zip() another iterator can end the loop while the generator + # is suspended, and the return value is NULL then too. + builder.activate_block(stop_block) + builder.primitive_op(propagate_if_error_op, [self.return_value], line) + builder.goto(self.loop_exit) def begin_body(self) -> None: # Assign the value obtained from the generator helper method to the @@ -796,11 +801,6 @@ def gen_step(self) -> None: # Nothing to do here, since we get the next item as part of gen_condition(). pass - def gen_cleanup(self) -> None: - # If return value is NULL (it wasn't assigned to by the generator helper method), - # an exception was raised that we need to propagate. - self.builder.primitive_op(propagate_if_error_op, [self.return_value], self.line) - class ForAsyncIterable(ForGenerator): """Generate IR for an async for loop.""" diff --git a/mypyc/test-data/run-generators.test b/mypyc/test-data/run-generators.test index b881079ab2fd..c1dea3023eca 100644 --- a/mypyc/test-data/run-generators.test +++ b/mypyc/test-data/run-generators.test @@ -522,6 +522,88 @@ with assertRaises(StopIteration): 2 42 +[case testNativeGeneratorInZip] +from typing import Generator, Iterator + +from testutil import assertRaises + +def gen() -> Iterator[int]: + yield 1 + yield 2 + yield 3 + +class C: + def items(self) -> Iterator[int]: + yield 1 + yield 2 + yield 3 + +def test_other_iterator_ends_loop() -> None: + result = [] + for i, x in zip(range(2), gen()): + result.append((i, x)) + assert result == [(0, 1), (1, 2)] + result = [] + for x, i in zip(gen(), range(2)): + result.append((x, i)) + assert result == [(1, 0), (2, 1)] + empty: list[int] = [] + assert [(i, x) for i, x in zip(empty, gen())] == [] + assert [(i, x) for i, x in zip(range(2), C().items())] == [(0, 1), (1, 2)] + assert [(a, x, b) for a, x, b in zip(range(5), gen(), range(2))] == [(0, 1, 0), (1, 2, 1)] + assert [(i, x, j) for (i, x), j in zip(enumerate(gen()), range(2))] == [(0, 1, 0), (1, 2, 1)] + +def test_generator_ends_loop() -> None: + assert [(i, x) for i, x in zip(range(5), gen())] == [(0, 1), (1, 2), (2, 3)] + assert [(x, i) for x, i in zip(gen(), range(5))] == [(1, 0), (2, 1), (3, 2)] + assert [(i, x) for i, x in enumerate(gen())] == [(0, 1), (1, 2), (2, 3)] + +def test_for_else() -> None: + result = [] + for i, x in zip(range(2), gen()): + result.append(x) + else: + result.append(-1) + assert result == [1, 2, -1] + result = [] + for x in gen(): + result.append(x) + else: + result.append(-1) + assert result == [1, 2, 3, -1] + +def gen_return() -> Generator[int, None, str]: + yield 1 + return "done" + +def test_generator_with_return_value() -> None: + assert [x for x in gen_return()] == [1] + assert [(x, i) for x, i in zip(gen_return(), range(5))] == [(1, 0)] + assert [(i, x) for i, x in zip(range(1), gen_return())] == [(0, 1)] + +def gen_raise() -> Iterator[int]: + yield 1 + raise ValueError("boom") + +def test_exception_from_generator() -> None: + result = [] + with assertRaises(ValueError, "boom"): + for i, x in zip(range(5), gen_raise()): + result.append((i, x)) + assert result == [(0, 1)] + with assertRaises(ValueError, "boom"): + for x in gen_raise(): + pass + # zip() stops before resuming the generator, so it doesn't raise + assert [(i, x) for i, x in zip(range(1), gen_raise())] == [(0, 1)] + +def pairs() -> Iterator[tuple[int, int]]: + for i, x in zip(range(2), gen()): + yield i, x + +def test_loop_in_generator() -> None: + assert list(pairs()) == [(0, 1), (1, 2)] + [case testGeneratorSuper] from typing import Iterator, Callable, Any