Skip to content
Open
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
7 changes: 7 additions & 0 deletions mypyc/codegen/emitclass.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)}"


Expand Down
176 changes: 176 additions & 0 deletions mypyc/test-data/run-dunders.test
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
Loading