diff --git a/mypyc/codegen/emitclass.py b/mypyc/codegen/emitclass.py index 69c0b3e7c8c6..6d3ba761426e 100644 --- a/mypyc/codegen/emitclass.py +++ b/mypyc/codegen/emitclass.py @@ -53,6 +53,13 @@ def native_slot(cl: ClassIR, fn: FuncIR, emitter: Emitter) -> str: + """Use a native method directly as a slot that takes self and returns an object. + + If the method returns an unboxed value, such as an int from __next__, a wrapper + that boxes the value fills the slot instead. + """ + if fn.ret_type.is_unboxed: + return generate_dunder_wrapper(cl, fn, emitter) return f"{NATIVE_PREFIX}{fn.cname(emitter.names)}" diff --git a/mypyc/test-data/run-dunders.test b/mypyc/test-data/run-dunders.test index 370de875b556..1159a0af6d1b 100644 --- a/mypyc/test-data/run-dunders.test +++ b/mypyc/test-data/run-dunders.test @@ -409,6 +409,182 @@ class InterpOverride(C): def __index__(self) -> int: return 2 +[case testDundersNext] +from typing import Any + +from mypy_extensions import i64, mypyc_attr +from testutil import assertRaises + +@mypyc_attr(allow_interpreted_subclasses=True) +class Counter: + def __init__(self, n: int) -> None: + self.i = 0 + self.n = n + + def __iter__(self) -> "Counter": + return self + + def __next__(self) -> int: + if self.i == self.n: + raise StopIteration + self.i += 1 + return self.i + +class Inherit(Counter): + pass + +class Override(Counter): + def __next__(self) -> int: + return super().__next__() * 10 + +class Container: + def __iter__(self) -> Counter: + return Counter(3) + +class BoolIter: + def __init__(self, items: list[bool]) -> None: + self.items = items + + def __iter__(self) -> "BoolIter": + return self + + def __next__(self) -> bool: + if not self.items: + raise StopIteration + return self.items.pop(0) + +class FloatIter: + def __init__(self, items: list[float], fail: bool = False) -> None: + self.items = items + self.fail = fail + + def __iter__(self) -> "FloatIter": + return self + + def __next__(self) -> float: + if not self.items: + if self.fail: + raise ValueError("bad float") + raise StopIteration + return self.items.pop(0) + +class I64Iter: + def __init__(self, items: list[i64]) -> None: + self.items = items + + def __iter__(self) -> "I64Iter": + return self + + def __next__(self) -> i64: + if not self.items: + raise StopIteration + return self.items.pop(0) + +class TupleIter: + def __init__(self) -> None: + self.items = [(1, "a"), (2, "b")] + + def __iter__(self) -> "TupleIter": + return self + + def __next__(self) -> tuple[int, str]: + if not self.items: + raise StopIteration + return self.items.pop(0) + +class StrIter: + def __init__(self) -> None: + self.items = ["a", "b"] + + def __iter__(self) -> "StrIter": + return self + + def __next__(self) -> str: + if not self.items: + raise StopIteration + return self.items.pop(0) + +class Raises: + def __iter__(self) -> "Raises": + return self + + def __next__(self) -> int: + raise ValueError("bad next") + +def test_int() -> None: + assert list(Counter(3)) == [1, 2, 3] + c = Counter(2) + assert next(c) == 1 + assert next(c) == 2 + with assertRaises(StopIteration): + next(c) + assert next(c, -1) == -1 + result = [] + for x in Counter(3): + result.append(x) + assert result == [1, 2, 3] + assert [x * 2 for x in Counter(3)] == [2, 4, 6] + assert list(Container()) == [1, 2, 3] + +def test_int_generic() -> None: + a: Any = Counter(3) + assert next(a) == 1 + assert a.__next__() == 2 + assert list(a) == [3] + with assertRaises(StopIteration): + a.__next__() + +def test_subclasses() -> None: + assert list(Inherit(3)) == [1, 2, 3] + assert list(Override(3)) == [10, 20, 30] + +def test_other_unboxed_types() -> None: + assert list(BoolIter([False, True, False])) == [False, True, False] + assert list(BoolIter([True, False])) == [True, False] + # -113.0 and -113 are the error values of float and i64. + assert list(FloatIter([1.5, -113.0, 2.0])) == [1.5, -113.0, 2.0] + assert list(I64Iter([1, -113, 2])) == [1, -113, 2] + assert list(TupleIter()) == [(1, "a"), (2, "b")] + +def test_boxed_type() -> None: + assert list(StrIter()) == ["a", "b"] + +def test_errors() -> None: + with assertRaises(ValueError, "bad next"): + list(Raises()) + with assertRaises(ValueError, "bad next"): + next(Raises()) + with assertRaises(ValueError, "bad float"): + list(FloatIter([1.0], fail=True)) + +def test_interpreted() -> None: + from interp import InterpInherit, InterpOverride, call_next, consume + + assert consume(Counter(3)) == [1, 2, 3] + assert call_next(Counter(3)) == 1 + assert consume(BoolIter([False, True])) == [False, True] + assert consume(TupleIter()) == [(1, "a"), (2, "b")] + assert consume(InterpInherit(2)) == [1, 2] + assert consume(InterpOverride(2)) == [101, 102] + +[file interp.py] +from typing import Any + +from native import Counter + +class InterpInherit(Counter): + pass + +class InterpOverride(Counter): + def __next__(self) -> int: + return super().__next__() + 100 + +def consume(it: Any) -> list[Any]: + return [x for x in it] + +def call_next(it: Any) -> Any: + return it.__next__() + [case testDundersBinarySimple] from typing import Any