diff --git a/python/src/com/jetbrains/python/psi/impl/PyClassImpl.java b/python/src/com/jetbrains/python/psi/impl/PyClassImpl.java index 6db8dea58c6d..f79d10d5dd34 100644 --- a/python/src/com/jetbrains/python/psi/impl/PyClassImpl.java +++ b/python/src/com/jetbrains/python/psi/impl/PyClassImpl.java @@ -22,6 +22,7 @@ import com.intellij.openapi.util.Comparing; import com.intellij.openapi.util.NotNullLazyValue; import com.intellij.openapi.util.Ref; import com.intellij.psi.*; +import com.intellij.psi.scope.BaseScopeProcessor; import com.intellij.psi.scope.PsiScopeProcessor; import com.intellij.psi.search.LocalSearchScope; import com.intellij.psi.search.SearchScope; @@ -553,7 +554,8 @@ public class PyClassImpl extends PyBaseElementImpl implements PyCla @Override @NotNull public PyFunction[] getMethods() { - return getClassChildren(PythonDialectsTokenSetProvider.INSTANCE.getFunctionDeclarationTokens(), PyFunction.ARRAY_FACTORY); + final TokenSet functionDeclarationTokens = PythonDialectsTokenSetProvider.INSTANCE.getFunctionDeclarationTokens(); + return getClassChildren(functionDeclarationTokens, PyFunction.class, PyFunction.ARRAY_FACTORY); } @Override @@ -565,24 +567,24 @@ public class PyClassImpl extends PyBaseElementImpl implements PyCla @Override public PyClass[] getNestedClasses() { - return getClassChildren(TokenSet.create(PyElementTypes.CLASS_DECLARATION), PyClass.ARRAY_FACTORY); + return getClassChildren(TokenSet.create(PyElementTypes.CLASS_DECLARATION), PyClass.class, PyClass.ARRAY_FACTORY); } - protected T[] getClassChildren(TokenSet elementTypes, ArrayFactory factory) { - // TODO: gather all top-level functions, maybe within control statements - final PyClassStub classStub = getStub(); - if (classStub != null) { - return classStub.getChildrenByType(elementTypes, factory); - } - List result = new ArrayList<>(); - final PyStatementList statementList = getStatementList(); - for (PsiElement element : statementList.getChildren()) { - if (elementTypes.contains(element.getNode().getElementType())) { - //noinspection unchecked - result.add((T)element); + @NotNull + private >> T[] getClassChildren(@NotNull TokenSet elementTypes, + @NotNull Class childrenClass, + @NotNull ArrayFactory factory) { + final List result = new ArrayList<>(); + processClassLevelDeclarations(new BaseScopeProcessor() { + @Override + public boolean execute(@NotNull PsiElement element, @NotNull ResolveState state) { + if (childrenClass.isInstance(element) && elementTypes.contains(((StubBasedPsiElement)element).getElementType())) { + result.add(childrenClass.cast(element)); + } + return true; } - } - return result.toArray(factory.create(result.size())); + }); + return ContainerUtil.toArray(result, factory); } private static class NameFinder implements Processor { diff --git a/python/testSrc/com/jetbrains/python/PyOverrideTest.java b/python/testSrc/com/jetbrains/python/PyOverrideTest.java index ce4c9597ab02..525486c0fd51 100644 --- a/python/testSrc/com/jetbrains/python/PyOverrideTest.java +++ b/python/testSrc/com/jetbrains/python/PyOverrideTest.java @@ -1,5 +1,5 @@ /* - * Copyright 2000-2013 JetBrains s.r.o. + * Copyright 2000-2017 JetBrains s.r.o. * * Licensed under the Apache License, Version 2.0 (the "License"); * you may not use this file except in compliance with the License. @@ -115,7 +115,7 @@ public class PyOverrideTest extends PyTestCase { public void testImplement() { myFixture.configureByFile("override/" + getTestName(true) + ".py"); - PyFunction toImplement = getTopLevelClass(0).getMethods()[1]; + PyFunction toImplement = getTopLevelClass(0).getMethods()[0]; PyOverrideImplementUtil.overrideMethods(myFixture.getEditor(), getTopLevelClass(1), Collections.singletonList(new PyMethodMember(toImplement)), true); myFixture.checkResultByFile("override/" + getTestName(true) + "_after.py", true);