Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
20 changes: 10 additions & 10 deletions mypyc/irbuild/for_helpers.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -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
Expand All @@ -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."""
Expand Down
82 changes: 82 additions & 0 deletions mypyc/test-data/run-generators.test
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
Loading