diff --git a/python/src/com/jetbrains/python/psi/impl/PyClassImpl.java b/python/src/com/jetbrains/python/psi/impl/PyClassImpl.java index 5033e503ac41..38df49a07383 100644 --- a/python/src/com/jetbrains/python/psi/impl/PyClassImpl.java +++ b/python/src/com/jetbrains/python/psi/impl/PyClassImpl.java @@ -51,6 +51,7 @@ import com.jetbrains.python.psi.stubs.PyFunctionStub; import com.jetbrains.python.psi.stubs.PyTargetExpressionStub; import com.jetbrains.python.psi.types.*; import com.jetbrains.python.toolbox.Maybe; +import one.util.streamex.StreamEx; import org.jetbrains.annotations.NotNull; import org.jetbrains.annotations.Nullable; @@ -235,24 +236,72 @@ public class PyClassImpl extends PyBaseElementImpl implements PyCla } @NotNull - public static PyExpression unfoldClass(@NotNull PyExpression expression) { - if (expression instanceof PyCallExpression) { - PyCallExpression call = (PyCallExpression)expression; - final PyExpression callee = call.getCallee(); - final PyExpression[] arguments = call.getArguments(); - if (callee != null && "with_metaclass".equals(callee.getName()) && arguments.length > 1) { - final PyExpression secondArgument = arguments[1]; - if (secondArgument != null) { - return secondArgument; - } + public static List getUnfoldedSuperClassExpressions(@NotNull PyClass pyClass) { + return StreamEx + .of(pyClass.getSuperClassExpressions()) + .filter(expression -> !PyKeywordArgument.class.isInstance(expression)) + .flatCollection(PyClassImpl::unfoldSuperClassExpression) + .toList(); + } + + @NotNull + private static List unfoldSuperClassExpression(@NotNull PyExpression expression) { + if (isSixWithMetaclassCall(expression)) { + final PyExpression[] arguments = ((PyCallExpression)expression).getArguments(); + if (arguments.length > 1) { + return ContainerUtil.newArrayList(arguments, 1, arguments.length); + } + else { + return Collections.emptyList(); } } // Heuristic: unfold Foo[Bar] to Foo for subscription expressions for superclasses else if (expression instanceof PySubscriptionExpression) { final PySubscriptionExpression subscriptionExpr = (PySubscriptionExpression)expression; - return subscriptionExpr.getOperand(); + return Collections.singletonList(subscriptionExpr.getOperand()); } - return expression; + + return Collections.singletonList(expression); + } + + private static boolean isSixWithMetaclassCall(@NotNull PyExpression expression) { + if (expression instanceof PyCallExpression){ + final PyCallExpression call = (PyCallExpression)expression; + final PyExpression callee = call.getCallee(); + if (callee != null && "with_metaclass".equals(callee.getName())) { + // SUPPORTED CASES: + + // import six + // six.with_metaclass(...) + + // from six import metaclass + // with_metaclass(...) + return true; + } + + if (callee instanceof PyReferenceExpression) { + // SUPPORTED CASES: + + // from six import with_metaclass as w_m + // w_m(...) + + final boolean importedWithMetaclass = StreamEx + .of(PyResolveUtil.resolveLocally((PyReferenceExpression)callee)) + .select(PyImportElement.class) + .map(PyImportElement::getImportedQName) + .nonNull() + .map(QualifiedName::getLastComponent) + .nonNull() + .findAny("with_metaclass"::equals) + .isPresent(); + + if (importedWithMetaclass) { + return true; + } + } + } + + return false; } @NotNull @@ -1294,12 +1343,7 @@ public class PyClassImpl extends PyBaseElementImpl implements PyCla } private void fillSuperClassesSwitchingToAst(@NotNull TypeEvalContext context, List result) { - for (PyExpression expression : getSuperClassExpressions()) { - context.getType(expression); - expression = unfoldClass(expression); - if (expression instanceof PyKeywordArgument) { - continue; - } + for (PyExpression expression : getUnfoldedSuperClassExpressions(this)) { final PyType type = context.getType(expression); PyClassLikeType classLikeType = null; if (type instanceof PyClassLikeType) { @@ -1401,6 +1445,16 @@ public class PyClassImpl extends PyBaseElementImpl implements PyCla return attribute.findAssignedValue(); } } + + for (PyExpression expression : getSuperClassExpressions()) { + if (isSixWithMetaclassCall(expression)) { + final PyExpression[] arguments = ((PyCallExpression)expression).getArguments(); + if (arguments.length != 0) { + return arguments[0]; + } + } + } + return null; } diff --git a/python/src/com/jetbrains/python/psi/impl/stubs/PyClassElementType.java b/python/src/com/jetbrains/python/psi/impl/stubs/PyClassElementType.java index fcf5a6d083b9..15427b4386e8 100644 --- a/python/src/com/jetbrains/python/psi/impl/stubs/PyClassElementType.java +++ b/python/src/com/jetbrains/python/psi/impl/stubs/PyClassElementType.java @@ -70,11 +70,12 @@ public class PyClassElementType extends PyStubElementType public static Map getSuperClassQNames(@NotNull final PyClass pyClass) { final Map result = new LinkedHashMap<>(); - Arrays - .stream(pyClass.getSuperClassExpressions()) - .filter(expression -> !PyKeywordArgument.class.isInstance(expression)) - .map(PyClassImpl::unfoldClass) - .forEach(expression -> result.put(PyPsiUtils.asQualifiedName(expression), resolveOriginalSuperClassQName(expression))); + for (PyExpression expression : PyClassImpl.getUnfoldedSuperClassExpressions(pyClass)) { + final QualifiedName importedQName = PyPsiUtils.asQualifiedName(expression); + final QualifiedName originalQName = resolveOriginalSuperClassQName(expression); + + result.put(importedQName, originalQName); + } return result; } diff --git a/python/testData/codeInsight/classMRO/SixWithMetaclass.py b/python/testData/codeInsight/classMRO/SixWithMetaclass.py index 317453bd9370..a16f3bb08405 100644 --- a/python/testData/codeInsight/classMRO/SixWithMetaclass.py +++ b/python/testData/codeInsight/classMRO/SixWithMetaclass.py @@ -1,3 +1,5 @@ +import six + class M(type): pass @@ -6,5 +8,9 @@ class B(object): pass -class C(six.with_metaclass(M, B)): +class D(object): + pass + + +class C(six.with_metaclass(M, B, D)): pass \ No newline at end of file diff --git a/python/testData/codeInsight/classMRO/SixWithMetaclassWithAs.py b/python/testData/codeInsight/classMRO/SixWithMetaclassWithAs.py new file mode 100644 index 000000000000..6da901f95da1 --- /dev/null +++ b/python/testData/codeInsight/classMRO/SixWithMetaclassWithAs.py @@ -0,0 +1,16 @@ +from six import with_metaclass as w_m + +class M(type): + pass + + +class B(object): + pass + + +class D(object): + pass + + +class C(w_m(M, B, D)): + pass \ No newline at end of file diff --git a/python/testData/inspections/PyUnresolvedReferencesInspection/sixWithMetaclass.py b/python/testData/inspections/PyUnresolvedReferencesInspection/sixWithMetaclass.py new file mode 100644 index 000000000000..d61eea7d408a --- /dev/null +++ b/python/testData/inspections/PyUnresolvedReferencesInspection/sixWithMetaclass.py @@ -0,0 +1,30 @@ +import six + + +class M(type): + def baz(self): + pass + + +class B1(object): + pass + + +class B2(object): + def bar(self): + pass + + +class C(six.with_metaclass(M, B1, B2)): + def foo(self): + self.bar() + C.baz() + + +from six import with_metaclass as w_m + + +class D(w_m(M, B1, B2)): + def foo(self): + self.bar() + D.baz() \ No newline at end of file diff --git a/python/testSrc/com/jetbrains/python/codeInsight/PyClassMROTest.java b/python/testSrc/com/jetbrains/python/codeInsight/PyClassMROTest.java index 7ced73929b59..e19391baed79 100644 --- a/python/testSrc/com/jetbrains/python/codeInsight/PyClassMROTest.java +++ b/python/testSrc/com/jetbrains/python/codeInsight/PyClassMROTest.java @@ -65,7 +65,11 @@ public class PyClassMROTest extends PyTestCase { } public void testSixWithMetaclass() { - assertMRO(getClass("C"), "B", "object"); + assertMRO(getClass("C"), "B", "D", "object"); + } + + public void testSixWithMetaclassWithAs() { + assertMRO(getClass("C"), "B", "D", "object"); } // PY-4183 diff --git a/python/testSrc/com/jetbrains/python/inspections/PyUnresolvedReferencesInspectionTest.java b/python/testSrc/com/jetbrains/python/inspections/PyUnresolvedReferencesInspectionTest.java index bf26bca98ab2..0c1276f5c0e1 100644 --- a/python/testSrc/com/jetbrains/python/inspections/PyUnresolvedReferencesInspectionTest.java +++ b/python/testSrc/com/jetbrains/python/inspections/PyUnresolvedReferencesInspectionTest.java @@ -734,6 +734,11 @@ public class PyUnresolvedReferencesInspectionTest extends PyInspectionTestCase { doMultiFileTest(); } + // PY-21224 + public void testSixWithMetaclass() { + doTest(); + } + @NotNull @Override protected Class getInspectionClass() {