From e7f5877888523a1eb94804c94fa49c1b2777a671 Mon Sep 17 00:00:00 2001 From: Andrey Vlasovskikh Date: Wed, 15 Oct 2014 17:58:13 +0400 Subject: [PATCH] Fixed type inference for subscription operator of a union type (PY-12862) --- .../impl/PySubscriptionExpressionImpl.java | 46 +++++++++++-------- .../MapReturnElementType.py | 6 +-- .../com/jetbrains/python/PyTypeTest.java | 40 ++++++++++++++-- 3 files changed, 66 insertions(+), 26 deletions(-) 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 e4c53bf11451..33cba24bb512 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" + @@ -906,6 +919,27 @@ public class PyTypeTest extends PyTestCase { "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.getContainingFile()).withTracing(); }