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
66 changes: 44 additions & 22 deletions mypyc/irbuild/expression.py
Original file line number Diff line number Diff line change
Expand Up @@ -402,21 +402,37 @@ def transform_super_expr(builder: IRBuilder, o: SuperExpr) -> Value:
else:
assert o.info is not None
typ = builder.load_native_type_object(o.info.fullname)
ir = builder.mapper.type_to_ir[o.info]
iter_env = iter(builder.builder.args)
# Grab first argument
vself: Value = next(iter_env)
if builder.fn_info.is_generator:
# grab seventh argument (see comment in translate_super_method_call)
self_targ = list(builder.symtables[-1].values())[7]
vself = builder.read(self_targ, builder.fn_info.fitem.line)
elif not ir.is_ext_class:
vself = next(iter_env) # second argument is self if non_extension class
vself = builder.read(builder.lookup(implicit_super_arg(builder)), o.line)
args = [typ, vself]
res = builder.py_call(sup_val, args, o.line)
return builder.py_get_attr(res, o.name, o.line)


def implicit_super_arg(builder: IRBuilder) -> Var:
"""Return the variable that zero-argument super() uses as its second argument.

This is the first argument of the enclosing function. It isn't always the first
argument of the function being generated, since generators and nested functions
are compiled to methods of generated classes.
"""
# A comprehension can have a scope of its own, but it has no arguments, and
# super() uses the function that contains the comprehension.
fn_info = next(info for info in reversed(builder.fn_infos) if not info.is_comprehension_scope)
return fn_info.fitem.arguments[0].variable


def is_instance_method_self(builder: IRBuilder, var: Var) -> bool:
"""Is this the self argument of an instance method?"""
if not var.is_self:
return False
# The cls argument of __new__ is also marked as a self argument.
for fn_info in builder.fn_infos:
fitem = fn_info.fitem
if fitem.name == "__new__" and fitem.arguments and fitem.arguments[0].variable is var:
return False
return True


# Calls


Expand Down Expand Up @@ -601,6 +617,9 @@ def translate_super_method_call(builder: IRBuilder, expr: CallExpr, callee: Supe
or callee.info is not typ_arg.node
):
return translate_call(builder, expr, callee)
self_var = self_arg.node
else:
self_var = implicit_super_arg(builder)

ir = builder.mapper.type_to_ir[callee.info]
# Search for the method in the mro, skipping ourselves. We
Expand Down Expand Up @@ -631,22 +650,25 @@ def translate_super_method_call(builder: IRBuilder, expr: CallExpr, callee: Supe
# super().prop(...) calls the property's value, so get it through super()
return translate_call(builder, expr, callee)

needs_self = decl.kind != FUNC_STATICMETHOD and decl.name != "__new__"
if needs_self and not (
is_instance_method_self(builder, self_var)
or (self_var.is_cls and decl.kind == FUNC_CLASSMETHOD)
):
# We can only bind the method statically if super() is given the self argument of
# an instance method, or the cls argument of a class method when calling a class
# method. Otherwise it's an instance method looked up through the class, or the
# first argument of a static method (such as __new__) or of a nested function,
# which can be either an instance or a class.
return translate_call(builder, expr, callee)

arg_values = [builder.accept(arg) for arg in expr.args]
arg_kinds, arg_names = expr.arg_kinds.copy(), expr.arg_names.copy()

if decl.kind != FUNC_STATICMETHOD and decl.name != "__new__":
# Grab first argument
vself: Value = builder.self()
if decl.kind == FUNC_CLASSMETHOD:
if needs_self:
vself = builder.read(builder.lookup(self_var), expr.line)
if decl.kind == FUNC_CLASSMETHOD and not self_var.is_cls:
vself = builder.primitive_op(type_op, [vself], expr.line)
elif builder.fn_info.is_generator:
# For generator classes, the self target is the 7th value
# in the symbol table (which is an ordered dict). This is sort
# of ugly, but we can't search by name since the 'self' parameter
# could be named anything, and it doesn't get added to the
# environment indexes.
self_targ = list(builder.symtables[-1].values())[7]
vself = builder.read(self_targ, builder.fn_info.fitem.line)
arg_values.insert(0, vself)
arg_kinds.insert(0, ARG_POS)
arg_names.insert(0, None)
Expand Down
78 changes: 78 additions & 0 deletions mypyc/test-data/irbuild-classes.test
Original file line number Diff line number Diff line change
Expand Up @@ -741,6 +741,84 @@ L0:
r0 = T.foo(self)
return 1

[case testSuperClassMethod]
class A:
@classmethod
def f(cls, x: int) -> int:
return x

def g(self) -> int:
return 1

class B(A):
@classmethod
def f(cls, x: int) -> int:
return super().f(x)

def call_f(self) -> int:
return super().f(1)

@classmethod
def call_g(cls, b: B) -> int:
# An instance method looked up through the class isn't bound
return super().g(b)
[out]
def A.f(cls, x):
cls :: object
x :: int
L0:
return x
def A.g(self):
self :: __main__.A
L0:
return 2
def B.f(cls, x):
cls :: object
x, r0 :: int
L0:
r0 = A.f(cls, x)
return r0
def B.call_f(self):
self :: __main__.B
r0 :: object
r1 :: int
L0:
r0 = CPy_TYPE(self)
r1 = A.f(r0, 2)
return r1
def B.call_g(cls, b):
cls :: object
b :: __main__.B
r0 :: object
r1 :: str
r2, r3 :: object
r4 :: object[2]
r5 :: object_ptr
r6 :: object
r7 :: str
r8 :: object
r9 :: object[1]
r10 :: object_ptr
r11 :: object
r12 :: int
L0:
r0 = builtins :: module
r1 = 'super'
r2 = CPyObject_GetAttr(r0, r1)
r3 = __main__.B :: type
r4 = [r3, cls]
r5 = load_address r4
r6 = PyObject_Vectorcall(r2, r5, 2, 0)
keep_alive r3, cls
r7 = 'g'
r8 = CPyObject_GetAttr(r6, r7)
r9 = [b]
r10 = load_address r9
r11 = PyObject_Vectorcall(r8, r10, 1, 0)
keep_alive b
r12 = unbox(int, r11)
return r12

[case testSuperCallToObjectInitIsOmitted]
class C:
def __init__(self) -> None:
Expand Down
190 changes: 190 additions & 0 deletions mypyc/test-data/run-classes.test
Original file line number Diff line number Diff line change
Expand Up @@ -1050,6 +1050,196 @@ yo!
3
yo!

[case testSuperInClassMethod]
from typing import Any
from mypy_extensions import trait

class A:
@classmethod
def name(cls) -> str:
return cls.__name__

@classmethod
def tagged(cls, tag: str, n: int = 1) -> str:
return f"{tag}{n}:{cls.__name__}"

@staticmethod
def static(x: int) -> int:
return x + 1

def method(self) -> str:
return "A.method:" + type(self).__name__

class B(A):
@classmethod
def name(cls) -> str:
return "B." + super().name()

@classmethod
def tagged(cls, tag: str, n: int = 1) -> str:
return "B." + super().tagged(tag, n=n + 1)

@classmethod
def two_arg(cls) -> str:
return "B." + super(B, cls).name()

@classmethod
def call_static(cls, x: int) -> int:
return super().static(x) + 10

@classmethod
def call_method(cls, b: "B") -> str:
# An instance method looked up through the class isn't bound
return "B." + super().method(b)

def from_instance(self) -> str:
return "B." + super().name() + "+" + super(B, self).name()

class C(B):
@classmethod
def name(cls) -> str:
return "C." + super().name()

def test_super_in_classmethod() -> None:
assert B.name() == "B.B"
assert C.name() == "C.B.C"
assert C().name() == "C.B.C"
assert B.tagged("x") == "B.x2:B"
assert C.tagged("x", 5) == "B.x6:C"
assert B.two_arg() == "B.B"
assert C.two_arg() == "B.C"
assert C.call_static(1) == 12
assert C.call_method(C()) == "B.A.method:C"
assert B().from_instance() == "B.B+B"
assert C().from_instance() == "B.C+C"

@trait
class T:
@classmethod
def describe(cls) -> str:
return "T:" + cls.__name__

class D(T):
@classmethod
def describe(cls) -> str:
return "D." + super().describe()

class E(D):
pass

def test_super_in_classmethod_with_trait() -> None:
assert D.describe() == "D.T:D"
assert E.describe() == "D.T:E"

init_subclass_log: list[str] = []

class Hooked:
def __init_subclass__(cls, **kwargs: Any) -> None:
super().__init_subclass__(**kwargs)
init_subclass_log.append("Hooked:" + cls.__name__)

class HookedChild(Hooked):
def __init_subclass__(cls, **kwargs: Any) -> None:
super().__init_subclass__(**kwargs)
init_subclass_log.append("HookedChild:" + cls.__name__)

class HookedGrandchild(HookedChild):
pass

def test_super_in_init_subclass() -> None:
assert init_subclass_log == [
"Hooked:HookedChild",
"Hooked:HookedGrandchild",
"HookedChild:HookedGrandchild",
]

class New(A):
made_by: str

def __new__(cls) -> "New":
obj = super().__new__(cls)
obj.made_by = super().name()
return obj

class NewChild(New):
pass

def test_super_in_dunder_new() -> None:
assert New().made_by == "New"
assert NewChild().made_by == "NewChild"

[case testSuperInGeneratorAndNestedFunction]
import asyncio
from typing import Callable, Iterator

class A:
@classmethod
def name(cls) -> str:
return cls.__name__

@classmethod
async def async_name(cls) -> str:
return "A.async_name:" + cls.__name__

def method(self) -> str:
return "A.method:" + type(self).__name__

class B(A):
@classmethod
def gen_cls(cls) -> Iterator[str]:
yield super().name()

def gen(self) -> Iterator[str]:
yield super().name()
yield super(B, self).name()
yield super().method()

@classmethod
async def async_name(cls) -> str:
return "B." + await super().async_name() + "+" + super().name()

async def coro(self) -> str:
return await super().async_name() + "+" + super().name()

def nested(self) -> list[str]:
def two_arg() -> str:
return super(B, self).method() + "+" + super(B, self).name()

# Zero-argument super() uses the first argument of the nested function
def zero_arg(b: B) -> str:
return super().method() + "+" + super().name()

f: Callable[[], str] = lambda: super(B, self).method()
return [two_arg(), zero_arg(self), f()]

@classmethod
def nested_cls(cls) -> list[str]:
def zero_arg(c: type[B]) -> str:
return super().name()

f: Callable[[], str] = lambda: super(B, cls).name()
return [zero_arg(cls), f()]

class C(B):
pass

def test_super_in_generator() -> None:
assert list(B.gen_cls()) == ["B"]
assert list(C.gen_cls()) == ["C"]
assert list(C().gen()) == ["C", "C", "A.method:C"]

def test_super_in_coroutine() -> None:
assert asyncio.run(B.async_name()) == "B.A.async_name:B+B"
assert asyncio.run(C.async_name()) == "B.A.async_name:C+C"
assert asyncio.run(C().coro()) == "A.async_name:C+C"

def test_super_in_nested_function() -> None:
assert C().nested() == ["A.method:C+C", "A.method:C+C", "A.method:C"]
assert C.nested_cls() == ["C", "C"]

[file asyncio/__init__.pyi]
def run(x: object) -> object: ...

[case testSubclassException]
class Failure(Exception):
def __init__(self, x: int) -> None:
Expand Down
Loading