Update name finder in class to prefer implementation over overloads (PY-35531)

GitOrigin-RevId: bca5e67a07ff96c79af9a290fe255a2c95e7baf9
This commit is contained in:
Semyon Proshev
2019-05-21 19:11:41 +03:00
committed by intellij-monorepo-bot
parent b8b12106c3
commit 9879dc60da
5 changed files with 55 additions and 13 deletions
@@ -118,7 +118,7 @@ public interface PyClass extends PsiNameIdentifierOwner, PyStatement, PyDocStrin
* @param name what to look for
* @param inherited true: search in superclasses; false: only look for methods defined in this class
* @param context context to be used to resolve ancestors
* @return method with given name or null.
* @return method with given name or null, prefers implementation over same name overloads.
*/
@Nullable
PyFunction findMethodByName(@Nullable @NonNls final String name, boolean inherited, TypeEvalContext context);
@@ -142,7 +142,8 @@ public interface PyClass extends PsiNameIdentifierOwner, PyStatement, PyDocStrin
*
* @param inherited true: search in superclasses, too.
* @param context context to be used to resolve ancestors and check if this class is a new-style class
* @return a method that would be called first when an instance of this class is instantiated.
* @return a method that would be called first when an instance of this class is instantiated,
* prefers implementation over same name overloads.
*/
@Nullable
PyFunction findInitOrNew(boolean inherited, @Nullable TypeEvalContext context);
@@ -36,6 +36,7 @@ import com.jetbrains.python.psi.stubs.PyClassStub;
import com.jetbrains.python.psi.stubs.PyFunctionStub;
import com.jetbrains.python.psi.stubs.PyTargetExpressionStub;
import com.jetbrains.python.psi.types.*;
import com.jetbrains.python.pyi.PyiUtil;
import com.jetbrains.python.toolbox.Maybe;
import one.util.streamex.StreamEx;
import org.jetbrains.annotations.NotNull;
@@ -504,12 +505,15 @@ public class PyClassImpl extends PyBaseElementImpl<PyClassStub> implements PyCla
}
private static class NameFinder<T extends PyElement> implements Processor<T> {
@NotNull
private final TypeEvalContext myContext;
private T myResult;
private final String[] myNames;
private int myLastResultIndex = -1;
private PyClass myLastVisitedClass = null;
NameFinder(String... names) {
NameFinder(@NotNull TypeEvalContext context, String... names) {
myContext = context;
myNames = names;
myResult = null;
}
@@ -535,11 +539,16 @@ public class PyClassImpl extends PyBaseElementImpl<PyClassStub> implements PyCla
final int index = ArrayUtil.indexOf(myNames, target.getName());
// Do not depend on the order in which elements appear, always try to find the first one
if (index >= 0 && (myLastResultIndex == -1 || index < myLastResultIndex)) {
myLastResultIndex = index;
myResult = target;
if (index == 0) {
return false;
if (index >= 0) {
if (myLastResultIndex == -1 ||
index < myLastResultIndex ||
index == myLastResultIndex && PyiUtil.isOverload(myResult, myContext) && !PyiUtil.isOverload(target, myContext)) {
myLastResultIndex = index;
myResult = target;
if (index == 0 && !PyiUtil.isOverload(myResult, myContext)) {
return false;
}
}
}
return true;
@@ -584,7 +593,7 @@ public class PyClassImpl extends PyBaseElementImpl<PyClassStub> implements PyCla
@Override
public PyFunction findMethodByName(@Nullable final String name, boolean inherited, @Nullable TypeEvalContext context) {
if (name == null) return null;
NameFinder<PyFunction> proc = new NameFinder<>(name);
NameFinder<PyFunction> proc = new NameFinder<>(notNullizeContext(context), name);
visitMethods(proc, inherited, context);
return proc.getResult();
}
@@ -601,7 +610,7 @@ public class PyClassImpl extends PyBaseElementImpl<PyClassStub> implements PyCla
@Override
public PyClass findNestedClass(String name, boolean inherited) {
if (name == null) return null;
NameFinder<PyClass> proc = new NameFinder<>(name);
NameFinder<PyClass> proc = new NameFinder<>(TypeEvalContext.codeInsightFallback(getProject()), name);
visitNestedClasses(proc, inherited);
return proc.getResult();
}
@@ -611,7 +620,7 @@ public class PyClassImpl extends PyBaseElementImpl<PyClassStub> implements PyCla
public PyFunction findInitOrNew(boolean inherited, final @Nullable TypeEvalContext context) {
NameFinder<PyFunction> proc;
if (isNewStyleClass(context)) {
proc = new NameFinder<PyFunction>(PyNames.INIT, PyNames.NEW) {
proc = new NameFinder<PyFunction>(notNullizeContext(context), PyNames.INIT, PyNames.NEW) {
@Nullable
@Override
protected PyClass getContainingClass(@NotNull PyFunction element) {
@@ -620,7 +629,7 @@ public class PyClassImpl extends PyBaseElementImpl<PyClassStub> implements PyCla
};
}
else {
proc = new NameFinder<>(PyNames.INIT);
proc = new NameFinder<>(notNullizeContext(context), PyNames.INIT);
}
visitMethods(proc, inherited, context);
return proc.getResult();
@@ -1041,7 +1050,7 @@ public class PyClassImpl extends PyBaseElementImpl<PyClassStub> implements PyCla
@Override
public PyTargetExpression findClassAttribute(@NotNull String name, boolean inherited, TypeEvalContext context) {
final NameFinder<PyTargetExpression> processor = new NameFinder<>(name);
final NameFinder<PyTargetExpression> processor = new NameFinder<>(notNullizeContext(context), name);
visitClassAttributes(processor, inherited, context);
return processor.getResult();
}
@@ -0,0 +1,6 @@
from typing import overload
class A:
@overload
def __init__(self, **kwargs): ...
def __init__(self, *args, **kwargs):
pass
@@ -15,6 +15,7 @@ import com.jetbrains.python.psi.resolve.ImportedResolveResult;
import com.jetbrains.python.psi.resolve.PyResolveContext;
import com.jetbrains.python.psi.types.PyClassTypeImpl;
import com.jetbrains.python.psi.types.TypeEvalContext;
import com.jetbrains.python.pyi.PyiUtil;
public class PyResolveTest extends PyResolveTestCase {
@Override
@@ -1339,4 +1340,14 @@ public class PyResolveTest extends PyResolveTestCase {
final PsiElement element = doResolve();
assertEquals(PyBuiltinCache.getInstance(myFixture.getFile()).getBuiltinsFile(), element);
}
// PY-35531
public void testOverloadedDunderInit() {
final PyFile file = (PyFile)myFixture.configureByFile("resolve/" + getTestName(false) + ".py");
final TypeEvalContext context = TypeEvalContext.codeAnalysis(myFixture.getProject(), file);
final PyFunction function = file.findTopLevelClass("A").findInitOrNew(false, context);
assertNotNull(function);
assertFalse(PyiUtil.isOverload(function, context));
}
}
@@ -787,6 +787,21 @@ public class PyUnresolvedReferencesInspectionTest extends PyInspectionTestCase {
doMultiFileTest();
}
// PY-35531
public void testAttributeDefinedInOverloadedDunderInit() {
runWithLanguageLevel(
LanguageLevel.PYTHON35,
() -> doTestByText("from typing import overload\n" +
"class Example:\n" +
" @overload\n" +
" def __init__(self, **kwargs): ...\n" +
" def __init__(self, *args, **kwargs):\n" +
" self.__data = None\n" +
" def test(self):\n" +
" return self.__data")
);
}
@NotNull
@Override
protected Class<? extends PyInspection> getInspectionClass() {