From 0de43685922115c56d1ebf61e5b74cbde02c1a56 Mon Sep 17 00:00:00 2001 From: Semyon Proshev Date: Thu, 8 Nov 2018 15:49:31 +0300 Subject: [PATCH] Return next after `Base` in `self`'s MRO for `super(Base, self)` (PY-32533) Previously first class in `Base`'s MRO were returned. --- .../psi/impl/PyCallExpressionHelper.java | 12 +++++++--- .../com/jetbrains/python/PyTypeTest.java | 24 +++++++++++++++++++ 2 files changed, 33 insertions(+), 3 deletions(-) diff --git a/python/src/com/jetbrains/python/psi/impl/PyCallExpressionHelper.java b/python/src/com/jetbrains/python/psi/impl/PyCallExpressionHelper.java index eeacb2ec0eab..356ff6321588 100644 --- a/python/src/com/jetbrains/python/psi/impl/PyCallExpressionHelper.java +++ b/python/src/com/jetbrains/python/psi/impl/PyCallExpressionHelper.java @@ -659,9 +659,15 @@ public class PyCallExpressionHelper { return getSuperClassUnionType(firstClass, context); } if (secondClass.isSubclass(firstClass, context)) { - final Iterator iterator = firstClass.getAncestorClasses(context).iterator(); - if (iterator.hasNext()) { - return new PyClassTypeImpl(iterator.next(), false); // super(Foo, self) has type of Foo, modulo __get__() + final PyClass nextAfterFirstInMro = StreamEx + .of(secondClass.getAncestorClasses(context)) + .dropWhile(it -> it != firstClass) + .skip(1) + .findFirst() + .orElse(null); + + if (nextAfterFirstInMro != null) { + return new PyClassTypeImpl(nextAfterFirstInMro, false); } } } diff --git a/python/testSrc/com/jetbrains/python/PyTypeTest.java b/python/testSrc/com/jetbrains/python/PyTypeTest.java index a4648f85a19d..489b1bf72881 100644 --- a/python/testSrc/com/jetbrains/python/PyTypeTest.java +++ b/python/testSrc/com/jetbrains/python/PyTypeTest.java @@ -3294,6 +3294,30 @@ public class PyTypeTest extends PyTestCase { " expr = a"); } + // PY-32533 + public void testSuperWithAnotherType() { + runWithLanguageLevel( + LanguageLevel.PYTHON34, + () -> doTest("A", + "class A:\n" + + " def f(self):\n" + + " return 'A'\n" + + "\n" + + "class B:\n" + + " def f(self):\n" + + " return 'B'\n" + + "\n" + + "class C(B):\n" + + " def f(self):\n" + + " return 'C'\n" + + "\n" + + "class D(C, A):\n" + + " def f(self):\n" + + " expr = super(B, self)\n" + + " return expr.f()") + ); + } + private static List getTypeEvalContexts(@NotNull PyExpression element) { return ImmutableList.of(TypeEvalContext.codeAnalysis(element.getProject(), element.getContainingFile()).withTracing(), TypeEvalContext.userInitiated(element.getProject(), element.getContainingFile()).withTracing());