diff --git a/python/src/com/jetbrains/python/psi/PyUtil.java b/python/src/com/jetbrains/python/psi/PyUtil.java index 1518b3410f11..67f8af119e92 100644 --- a/python/src/com/jetbrains/python/psi/PyUtil.java +++ b/python/src/com/jetbrains/python/psi/PyUtil.java @@ -68,6 +68,7 @@ import com.jetbrains.python.magicLiteral.PyMagicLiteralTools; import com.jetbrains.python.psi.impl.PyBuiltinCache; import com.jetbrains.python.psi.impl.PyPsiUtils; import com.jetbrains.python.psi.impl.PythonLanguageLevelPusher; +import com.jetbrains.python.psi.resolve.RatedResolveResult; import com.jetbrains.python.psi.types.*; import com.jetbrains.python.refactoring.classes.PyDependenciesComparator; import com.jetbrains.python.refactoring.classes.extractSuperclass.PyExtractSuperclassHelper; @@ -750,6 +751,40 @@ public class PyUtil { return currentElement; } + @NotNull + public static List multiResolveTopPriority(@NotNull PsiPolyVariantReference reference) { + return filterTopPriorityResults(reference.multiResolve(false)); + } + + @NotNull + private static List filterTopPriorityResults(@NotNull ResolveResult[] resolveResults) { + if (resolveResults.length == 0) { + return Collections.emptyList(); + } + final List filtered = new ArrayList(); + final int maxRate = getMaxRate(resolveResults); + for (ResolveResult resolveResult : resolveResults) { + final int rate = resolveResult instanceof RatedResolveResult ? ((RatedResolveResult)resolveResult).getRate() : 0; + if (rate >= maxRate) { + filtered.add(resolveResult.getElement()); + } + } + return filtered; + } + + private static int getMaxRate(@NotNull ResolveResult[] resolveResults) { + int maxRate = Integer.MIN_VALUE; + for (ResolveResult resolveResult : resolveResults) { + if (resolveResult instanceof RatedResolveResult) { + final int rate = ((RatedResolveResult)resolveResult).getRate(); + if (rate > maxRate) { + maxRate = rate; + } + } + } + return maxRate; + } + /** * Gets class init method * diff --git a/python/src/com/jetbrains/python/psi/impl/PyCallExpressionHelper.java b/python/src/com/jetbrains/python/psi/impl/PyCallExpressionHelper.java index d44c848049ba..30c6d9e82078 100644 --- a/python/src/com/jetbrains/python/psi/impl/PyCallExpressionHelper.java +++ b/python/src/com/jetbrains/python/psi/impl/PyCallExpressionHelper.java @@ -17,9 +17,10 @@ package com.jetbrains.python.psi.impl; import com.intellij.codeInsight.completion.CompletionUtil; import com.intellij.openapi.util.Pair; +import com.intellij.openapi.util.Ref; import com.intellij.psi.PsiElement; +import com.intellij.psi.PsiPolyVariantReference; import com.intellij.psi.PsiReference; -import com.intellij.psi.ResolveResult; import com.intellij.psi.util.PsiTreeUtil; import com.jetbrains.python.FunctionParameter; import com.jetbrains.python.PyNames; @@ -429,57 +430,18 @@ public class PyCallExpressionHelper { } // normal cases final PyResolveContext resolveContext = PyResolveContext.noImplicits().withTypeEvalContext(context); - ResolveResult[] targets = ((PyReferenceExpression)callee).getReference(resolveContext).multiResolve(false); - if (targets.length > 0) { - PsiElement target = targets[0].getElement(); - if (target == null) { - return null; - } - PyClass cls = null; - PyFunction init = null; - if (target instanceof PyClass) { - cls = (PyClass)target; - init = cls.findInitOrNew(true); - } - else if (target instanceof PyFunction) { - final PyFunction f = (PyFunction)target; - if (PyNames.INIT.equals(f.getName())) { - init = f; - cls = f.getContainingClass(); + final PsiPolyVariantReference reference = ((PyReferenceExpression)callee).getReference(resolveContext); + final List members = new ArrayList(); + for (PsiElement target : PyUtil.multiResolveTopPriority(reference)) { + if (target != null) { + final Ref typeRef = getCallTargetReturnType(call, target, context); + if (typeRef != null) { + members.add(typeRef.get()); } } - if (init != null) { - final PyType t = init.getCallType(context, call); - if (cls != null) { - if (init.getContainingClass() != cls) { - if (t instanceof PyCollectionType) { - final PyType elementType = ((PyCollectionType)t).getElementType(context); - return new PyCollectionTypeImpl(cls, false, elementType); - } - return new PyClassTypeImpl(cls, false); - } - } - if (t != null && !(t instanceof PyNoneType)) { - return t; - } - if (cls != null && t == null) { - final PyFunction newMethod = cls.findMethodByName(PyNames.NEW, true); - if (newMethod != null && !PyBuiltinCache.getInstance(call).isBuiltin(newMethod)) { - return PyUnionType.createWeakType(new PyClassTypeImpl(cls, false)); - } - } - } - if (cls != null) { - return new PyClassTypeImpl(cls, false); - } - final PyType providedType = PyReferenceExpressionImpl.getReferenceTypeFromProviders(target, context, call); - if (providedType instanceof PyCallableType) { - return ((PyCallableType)providedType).getCallType(context, call); - } - if (target instanceof Callable) { - final Callable callable = (Callable)target; - return callable.getCallType(context, call); - } + } + if (!members.isEmpty()) { + return PyUnionType.union(members); } } if (callee == null) { @@ -499,6 +461,57 @@ public class PyCallExpressionHelper { } } + @Nullable + private static Ref getCallTargetReturnType(@NotNull PyCallExpression call, @NotNull PsiElement target, + @NotNull TypeEvalContext context) { + PyClass cls = null; + PyFunction init = null; + if (target instanceof PyClass) { + cls = (PyClass)target; + init = cls.findInitOrNew(true); + } + else if (target instanceof PyFunction) { + final PyFunction f = (PyFunction)target; + if (PyNames.INIT.equals(f.getName())) { + init = f; + cls = f.getContainingClass(); + } + } + if (init != null) { + final PyType t = init.getCallType(context, call); + if (cls != null) { + if (init.getContainingClass() != cls) { + if (t instanceof PyCollectionType) { + final PyType elementType = ((PyCollectionType)t).getElementType(context); + return Ref.create(new PyCollectionTypeImpl(cls, false, elementType)); + } + return Ref.create(new PyClassTypeImpl(cls, false)); + } + } + if (t != null && !(t instanceof PyNoneType)) { + return Ref.create(t); + } + if (cls != null && t == null) { + final PyFunction newMethod = cls.findMethodByName(PyNames.NEW, true); + if (newMethod != null && !PyBuiltinCache.getInstance(call).isBuiltin(newMethod)) { + return Ref.create(PyUnionType.createWeakType(new PyClassTypeImpl(cls, false))); + } + } + } + if (cls != null) { + return Ref.create(new PyClassTypeImpl(cls, false)); + } + final PyType providedType = PyReferenceExpressionImpl.getReferenceTypeFromProviders(target, context, call); + if (providedType instanceof PyCallableType) { + return Ref.create(((PyCallableType)providedType).getCallType(context, call)); + } + if (target instanceof Callable) { + final Callable callable = (Callable)target; + return Ref.create(callable.getCallType(context, call)); + } + return null; + } + @NotNull private static Maybe getSuperCallType(@NotNull PyCallExpression call, TypeEvalContext context) { final PyExpression callee = call.getCallee(); diff --git a/python/src/com/jetbrains/python/psi/impl/PyReferenceExpressionImpl.java b/python/src/com/jetbrains/python/psi/impl/PyReferenceExpressionImpl.java index 126ed26acf32..edd902da1a55 100644 --- a/python/src/com/jetbrains/python/psi/impl/PyReferenceExpressionImpl.java +++ b/python/src/com/jetbrains/python/psi/impl/PyReferenceExpressionImpl.java @@ -216,12 +216,14 @@ public class PyReferenceExpressionImpl extends PyElementImpl implements PyRefere return typeOfProperty.get(); } } - ResolveResult[] targets = getReference(PyResolveContext.noImplicits().withTypeEvalContext(context)).multiResolve(false); - if (targets.length == 0) { + final PsiPolyVariantReference reference = getReference(PyResolveContext.noImplicits().withTypeEvalContext(context)); + final List targets = PyUtil.multiResolveTopPriority(reference); + if (targets.isEmpty()) { return getQualifiedReferenceTypeByControlFlow(context); } - for (ResolveResult resolveResult : targets) { - PsiElement target = resolveResult.getElement(); + + final List members = new ArrayList(); + for (PsiElement target : targets) { if (target == this || target == null) { continue; } @@ -229,12 +231,10 @@ public class PyReferenceExpressionImpl extends PyElementImpl implements PyRefere LOG.error("Reference " + this + " resolved to invalid element " + target + " (text=" + target.getText() + ")"); continue; } - type = getTypeFromTarget(target, context, this); - if (type != null) { - return type; - } + members.add(getTypeFromTarget(target, context, this)); } - return null; + + return PyUnionType.union(members); } finally { TypeEvalStack.evaluated(this); diff --git a/python/src/com/jetbrains/python/psi/impl/PySubscriptionExpressionImpl.java b/python/src/com/jetbrains/python/psi/impl/PySubscriptionExpressionImpl.java index 9cceec4fd5d5..c27e9195a910 100644 --- a/python/src/com/jetbrains/python/psi/impl/PySubscriptionExpressionImpl.java +++ b/python/src/com/jetbrains/python/psi/impl/PySubscriptionExpressionImpl.java @@ -30,6 +30,9 @@ import com.jetbrains.python.psi.types.*; import org.jetbrains.annotations.NotNull; import org.jetbrains.annotations.Nullable; +import java.util.ArrayList; +import java.util.List; + /** * @author yole */ @@ -66,29 +69,32 @@ public class PySubscriptionExpressionImpl extends PyElementImpl implements PySub @Nullable @Override public PyType getType(@NotNull TypeEvalContext context, @NotNull TypeEvalContext.Key key) { - PyType res = null; - final PsiReference ref = getReference(PyResolveContext.noImplicits().withTypeEvalContext(context)); - final PsiElement resolved = ref.resolve(); - if (resolved instanceof Callable) { - res = ((Callable)resolved).getCallType(context, this); - } - if (PyTypeChecker.isUnknown(res) || res instanceof PyNoneType) { - final PyExpression indexExpression = getIndexExpression(); - if (indexExpression != null) { - final PyType type = context.getType(getOperand()); - final PyClass cls = (type instanceof PyClassType) ? ((PyClassType)type).getPyClass() : null; - if (cls != null && PyABCUtil.isSubclass(cls, PyNames.MAPPING)) { - return res; - } - if (type instanceof PySubscriptableType) { - res = ((PySubscriptableType)type).getElementType(indexExpression, context); - } - else if (type instanceof PyCollectionType) { - res = ((PyCollectionType)type).getElementType(context); + final PsiPolyVariantReference reference = getReference(PyResolveContext.noImplicits().withTypeEvalContext(context)); + final List members = new ArrayList(); + for (PsiElement resolved : PyUtil.multiResolveTopPriority(reference)) { + PyType res = null; + if (resolved instanceof Callable) { + res = ((Callable)resolved).getCallType(context, this); + } + if (PyTypeChecker.isUnknown(res) || res instanceof PyNoneType) { + final PyExpression indexExpression = getIndexExpression(); + if (indexExpression != null) { + final PyType type = context.getType(getOperand()); + final PyClass cls = (type instanceof PyClassType) ? ((PyClassType)type).getPyClass() : null; + if (cls != null && PyABCUtil.isSubclass(cls, PyNames.MAPPING)) { + return res; + } + if (type instanceof PySubscriptableType) { + res = ((PySubscriptableType)type).getElementType(indexExpression, context); + } + else if (type instanceof PyCollectionType) { + res = ((PyCollectionType)type).getElementType(context); + } } } + members.add(res); } - return res; + return PyUnionType.union(members); } @Override diff --git a/python/testData/inspections/PyTypeCheckerInspection/MapReturnElementType.py b/python/testData/inspections/PyTypeCheckerInspection/MapReturnElementType.py index 87628aa630d4..a27266914eed 100644 --- a/python/testData/inspections/PyTypeCheckerInspection/MapReturnElementType.py +++ b/python/testData/inspections/PyTypeCheckerInspection/MapReturnElementType.py @@ -1,5 +1,5 @@ def test(): xs = map(lambda x: x + 1, [1, 2, 3]) - print('foo' + xs[0]) - ys = map(str, iter([1, 2, 3])) - print(1 + ys[0], 'bar' + ys[1]) + print('foo' + xs[0]) # Can be a str since map returns list[V] | str | unicode + ys = map(tuple, iter([1, 2, 3])) + print(1 + ys[0], 'bar' + ys[1]) diff --git a/python/testSrc/com/jetbrains/python/PyTypeTest.java b/python/testSrc/com/jetbrains/python/PyTypeTest.java index 1596596ad2de..e392371e13a0 100644 --- a/python/testSrc/com/jetbrains/python/PyTypeTest.java +++ b/python/testSrc/com/jetbrains/python/PyTypeTest.java @@ -621,12 +621,25 @@ public class PyTypeTest extends PyTestCase { } public void testFunctionTypeAsUnificationArgument() { - doTest("int", + doTest("list[int] | str | unicode", "def map2(f, xs):\n" + " '''\n" + " :type f: (T) -> V | None\n" + - " :type xs: collections.Iterable[T] | bytes | unicode\n" + - " :rtype: list[V] | bytes | unicode\n" + + " :type xs: collections.Iterable[T] | str | unicode\n" + + " :rtype: list[V] | str | unicode\n" + + " '''\n" + + " pass\n" + + "\n" + + "expr = map2(lambda x: 10, ['1', '2', '3'])\n"); + } + + public void testFunctionTypeAsUnificationArgumentWithSubscription() { + doTest("int | str | unicode", + "def map2(f, xs):\n" + + " '''\n" + + " :type f: (T) -> V | None\n" + + " :type xs: collections.Iterable[T] | str | unicode\n" + + " :rtype: list[V] | str | unicode\n" + " '''\n" + " pass\n" + "\n" + @@ -873,6 +886,60 @@ public class PyTypeTest extends PyTestCase { "expr = np.ones(10) * 2\n"); } + public void testUnionTypeAttributeOfDifferentTypes() { + doTest("list | int", + "class Foo:\n" + + " x = []\n" + + "\n" + + "class Bar:\n" + + " x = 42\n" + + "\n" + + "def f(c):\n" + + " o = Foo() if c else Bar()\n" + + " expr = o.x\n"); + } + + // PY-11364 + public void testUnionTypeAttributeCallOfDifferentTypes() { + doTest("C1 | C2", + "class C1:\n" + + " def foo(self):\n" + + " return self\n" + + "\n" + + "class C2:\n" + + " def foo(self):\n" + + " return self\n" + + "\n" + + "def f():\n" + + " '''\n" + + " :rtype: C1 | C2\n" + + " '''\n" + + " pass\n" + + "\n" + + "expr = f().foo()\n"); + } + + // PY-12862 + public void testUnionTypeAttributeSubscriptionOfDifferentTypes() { + doTest("C1 | C2", + "class C1:\n" + + " def __getitem__(self, item):\n" + + " return self\n" + + "\n" + + "class C2:\n" + + " def __getitem__(self, item):\n" + + " return self\n" + + "\n" + + "def f():\n" + + " '''\n" + + " :rtype: C1 | C2\n" + + " '''\n" + + " pass\n" + + "\n" + + "expr = f()[0]\n" + + "print(expr)\n"); + } + private static TypeEvalContext getTypeEvalContext(@NotNull PyExpression element) { return TypeEvalContext.userInitiated(element.getProject(), element.getContainingFile()).withTracing(); }