diff --git a/python/src/com/jetbrains/python/psi/types/PyTypeInferenceFromUsedAttributesUtil.java b/python/src/com/jetbrains/python/psi/types/PyTypeInferenceFromUsedAttributesUtil.java index 7fdb20346154..3e55067fa35a 100644 --- a/python/src/com/jetbrains/python/psi/types/PyTypeInferenceFromUsedAttributesUtil.java +++ b/python/src/com/jetbrains/python/psi/types/PyTypeInferenceFromUsedAttributesUtil.java @@ -6,11 +6,11 @@ import com.google.common.collect.Lists; import com.google.common.collect.Sets; import com.intellij.openapi.diagnostic.Logger; import com.intellij.openapi.roots.ProjectRootManager; -import com.intellij.openapi.util.Comparing; +import com.intellij.openapi.util.*; import com.intellij.psi.PsiFile; import com.intellij.psi.PsiNamedElement; import com.intellij.psi.PsiReference; -import com.intellij.psi.util.QualifiedName; +import com.intellij.psi.util.*; import com.intellij.util.Function; import com.intellij.util.containers.ContainerUtil; import com.jetbrains.python.PyNames; @@ -47,6 +47,8 @@ public class PyTypeInferenceFromUsedAttributesUtil { ); private static final Logger LOG = Logger.getInstance(PyTypeInferenceFromUsedAttributesUtil.class); + private static final CachedAncestorsFastProvider ourCachedAncestorsProvider = new CachedAncestorsFastProvider(); + private static final RecursionGuard ourRecursionGuard = RecursionManager.createGuard("py.types.from.attrs.rec.guard"); private PyTypeInferenceFromUsedAttributesUtil() { // empty @@ -103,13 +105,13 @@ public class PyTypeInferenceFromUsedAttributesUtil { if (PyUserSkeletonsUtil.isUnderUserSkeletonsDirectory(candidate.getContainingFile())) { continue; } - if (getAllInheritedAttributeNames(candidate, context).containsAll(seenAttrs)) { + if (getAllInheritedAttributeNames(candidate).containsAll(seenAttrs)) { suitableClasses.add(candidate); } } for (PyClass candidate : Lists.newArrayList(suitableClasses)) { - for (PyClass ancestor : candidate.getAncestorClasses()) { + for (PyClass ancestor : getAncestorClassesFast(candidate)) { if (suitableClasses.contains(ancestor)) { suitableClasses.remove(candidate); } @@ -126,9 +128,9 @@ public class PyTypeInferenceFromUsedAttributesUtil { } @NotNull - private static Set getAllInheritedAttributeNames(@NotNull PyClass candidate, @NotNull TypeEvalContext context) { + private static Set getAllInheritedAttributeNames(@NotNull PyClass candidate) { final Set availableAttrs = Sets.newHashSet(getAllDeclaredAttributeNames(candidate)); - for (PyClass parent : candidate.getAncestorClasses(context)) { + for (PyClass parent : getAncestorClassesFast(candidate)) { availableAttrs.addAll(getAllDeclaredAttributeNames(parent)); } return availableAttrs; @@ -278,6 +280,37 @@ public class PyTypeInferenceFromUsedAttributesUtil { } } + @NotNull + private static Set getAncestorClassesFast(@NotNull PyClass pyClass) { + final CachedValuesManager manager = CachedValuesManager.getManager(pyClass.getProject()); + return manager.getParameterizedCachedValue(pyClass, CachedAncestorsFastProvider.KEY, ourCachedAncestorsProvider, false, pyClass); + } + + private static class CachedAncestorsFastProvider implements ParameterizedCachedValueProvider, PyClass> { + static Key, PyClass>> KEY = Key.create("py.types.from.attrs.cached.ancestors.fast"); + + @NotNull + @Override + public CachedValueProvider.Result> compute(@NotNull PyClass pyClass) { + final HashSet result = Sets.newHashSet(); + for (final PyClass baseClass : pyClass.getSuperClasses()) { + final Computable> computable = new Computable>() { + @Override + public Set compute() { + return getAncestorClassesFast(baseClass); + } + }; + final Set baseClassAncestors = ourRecursionGuard.doPreventingRecursion(baseClass, false, computable); + result.add(baseClass); + if (baseClassAncestors != null) { + result.addAll(baseClassAncestors); + } + } + return CachedValueProvider.Result.create(Collections.unmodifiableSet(result), + PsiModificationTracker.OUT_OF_CODE_BLOCK_MODIFICATION_COUNT); + } + } + enum Priority { BUILTIN, SAME_FILE, diff --git a/python/testData/typesFromAttributes/cyclicInheritance/main.py b/python/testData/typesFromAttributes/cyclicInheritance/main.py new file mode 100644 index 000000000000..530c522ca408 --- /dev/null +++ b/python/testData/typesFromAttributes/cyclicInheritance/main.py @@ -0,0 +1,9 @@ +from module import A + +class B(A): + def unique(self): + pass + +x = undefined() +x.unique() +x \ No newline at end of file diff --git a/python/testData/typesFromAttributes/cyclicInheritance/module.py b/python/testData/typesFromAttributes/cyclicInheritance/module.py new file mode 100644 index 000000000000..82cd4f4fcfa8 --- /dev/null +++ b/python/testData/typesFromAttributes/cyclicInheritance/module.py @@ -0,0 +1,4 @@ +from main import B + +class A(B): + pass diff --git a/python/testSrc/com/jetbrains/python/PyTypeFromUsedAttributesTest.java b/python/testSrc/com/jetbrains/python/PyTypeFromUsedAttributesTest.java index d9498e1419ac..39bf35b666be 100644 --- a/python/testSrc/com/jetbrains/python/PyTypeFromUsedAttributesTest.java +++ b/python/testSrc/com/jetbrains/python/PyTypeFromUsedAttributesTest.java @@ -190,6 +190,17 @@ public class PyTypeFromUsedAttributesTest extends PyTestCase { "list | MySortable | OtherClassA | OtherClassB | unknown"); } + public void testCyclicInheritance() { + myFixture.copyDirectoryToProject(getTestName(true), ""); + myFixture.configureByFile("main.py"); + final PyReferenceExpression referenceExpression = findLastReferenceByText("x"); + assertNotNull(referenceExpression); + final TypeEvalContext context = TypeEvalContext.userInitiated(referenceExpression.getContainingFile()).withTracing(); + final PyType actual = context.getType(referenceExpression); + final String actualType = PythonDocumentationProvider.getTypeName(actual, context); + assertEquals("unknown", actualType); + } + private void doTestType(@NotNull String text, @NotNull String expectedType) { myFixture.configureByText(PythonFileType.INSTANCE, text); final PyReferenceExpression referenceExpression = findLastReferenceByText("x");