diff --git a/python/psi-api/src/com/jetbrains/python/PyNames.java b/python/psi-api/src/com/jetbrains/python/PyNames.java index b8d657b15bf3..11c723e9de7e 100644 --- a/python/psi-api/src/com/jetbrains/python/PyNames.java +++ b/python/psi-api/src/com/jetbrains/python/PyNames.java @@ -185,6 +185,7 @@ public class PyNames { public static final String TUPLE = "tuple"; public static final String SET = "set"; + public static final String SLICE = "slice"; public static final String KEYS = "keys"; public static final String EXTEND = "extend"; diff --git a/python/src/com/jetbrains/python/codeInsight/stdlib/PyStdlibTypeProvider.java b/python/src/com/jetbrains/python/codeInsight/stdlib/PyStdlibTypeProvider.java index b7283d7d3029..8f4a16b80da6 100644 --- a/python/src/com/jetbrains/python/codeInsight/stdlib/PyStdlibTypeProvider.java +++ b/python/src/com/jetbrains/python/codeInsight/stdlib/PyStdlibTypeProvider.java @@ -25,6 +25,7 @@ import com.jetbrains.python.codeInsight.controlflow.ScopeOwner; import com.jetbrains.python.codeInsight.dataflow.scope.ScopeUtil; import com.jetbrains.python.psi.*; import com.jetbrains.python.psi.impl.PyBuiltinCache; +import com.jetbrains.python.psi.impl.PyCallExpressionHelper; import com.jetbrains.python.psi.impl.PyPsiUtils; import com.jetbrains.python.psi.impl.PyTypeProvider; import com.jetbrains.python.psi.impl.stubs.PyNamedTupleStubImpl; @@ -38,6 +39,7 @@ import org.jetbrains.annotations.Nullable; import java.util.List; import java.util.Map; +import java.util.Optional; import java.util.Set; import static com.jetbrains.python.psi.PyUtil.as; @@ -148,7 +150,7 @@ public class PyStdlibTypeProvider extends PyTypeProviderBase { if (OPEN_FUNCTIONS.contains(qname) && callSite instanceof PyCallExpression) { final PyCallExpression callExpr = (PyCallExpression)callSite; final PyResolveContext resolveContext = PyResolveContext.noImplicits().withTypeEvalContext(context); - final PyCallExpression.PyArgumentsMapping mapping = callExpr.mapArguments( resolveContext); + final PyCallExpression.PyArgumentsMapping mapping = callExpr.mapArguments(resolveContext); if (mapping.getMarkedCallee() != null) { final PyType type = getOpenFunctionType(qname, mapping.getMappedParameters(), callSite); if (type != null) { @@ -162,10 +164,63 @@ public class PyStdlibTypeProvider extends PyTypeProviderBase { else if ("__builtin__.tuple.__mul__".equals(qname) && callSite instanceof PyBinaryExpression) { return getTupleMultiplicationResultType((PyBinaryExpression)callSite, context); } + else if (callSite != null && isListGetItem(function)) { + final PyExpression receiver = PyTypeChecker.getReceiver(callSite, function); + final List arguments = PyTypeChecker.getArguments(callSite, function); + final List parameters = PyUtil.getParameters(function, context); + final PyResolveContext resolveContext = PyResolveContext.noImplicits().withTypeEvalContext(context); + final List explicitParameters = PyTypeChecker.filterExplicitParameters(parameters, function, callSite, resolveContext); + final Map mapping = PyCallExpressionHelper.mapArguments(arguments, explicitParameters); + final Map substitutions = PyTypeChecker.unifyGenericCall(receiver, mapping, context); + if (substitutions != null) { + return analyzeListGetItemCallType(receiver, mapping, substitutions, context); + } + } } return null; } + private static boolean isListGetItem(@NotNull PyFunction function) { + return PyNames.GETITEM.equals(function.getName()) && + Optional + .ofNullable(PyBuiltinCache.getInstance(function).getListType()) + .map(PyClassType::getPyClass) + .map(cls -> cls.equals(function.getContainingClass())) + .orElse(false); + } + + @Nullable + private static PyType analyzeListGetItemCallType(@Nullable PyExpression receiver, + @NotNull Map parameters, + @NotNull Map substitutions, + @NotNull TypeEvalContext context) { + if (parameters.size() != 1 || substitutions.size() != 1) { + return null; + } + + final PyType firstArgumentType = Optional + .ofNullable(parameters.keySet().iterator().next()) + .map(context::getType) + .orElse(null); + + if (firstArgumentType == null) { + return null; + } + + if (PyABCUtil.isSubtype(firstArgumentType, PyNames.ABC_INTEGRAL, context)) { + return substitutions.values().iterator().next(); + } + + if (PyNames.SLICE.equals(firstArgumentType.getName()) && firstArgumentType.isBuiltin()) { + return Optional + .ofNullable(receiver) + .map(context::getType) + .orElseGet(() -> PyTypeChecker.substitute(PyBuiltinCache.getInstance(receiver).getListType(), substitutions, context)); + } + + return null; + } + @Nullable private static PyType getTupleMultiplicationResultType(@NotNull PyBinaryExpression multiplication, @NotNull TypeEvalContext context) { final PyTupleType leftTupleType = as(context.getType(multiplication.getLeftExpression()), PyTupleType.class); diff --git a/python/testData/inspections/PyTypeCheckerInspection/MapReturnElementType.py b/python/testData/inspections/PyTypeCheckerInspection/MapReturnElementType.py index b9aa981911b9..14563237d69e 100644 --- a/python/testData/inspections/PyTypeCheckerInspection/MapReturnElementType.py +++ b/python/testData/inspections/PyTypeCheckerInspection/MapReturnElementType.py @@ -2,4 +2,4 @@ def test(): xs = map(lambda x: x + 1, [1, 2, 3]) 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]) + print(1 + ys[0], 'bar' + ys[1]) diff --git a/python/testData/inspections/PyTypeCheckerInspection/Subscript.py b/python/testData/inspections/PyTypeCheckerInspection/Subscript.py index adbbfed0df52..0a3ffe35fe3d 100644 --- a/python/testData/inspections/PyTypeCheckerInspection/Subscript.py +++ b/python/testData/inspections/PyTypeCheckerInspection/Subscript.py @@ -11,7 +11,7 @@ def test(): """ xs = [1, 2, 3] x = xs[0] - f(x) + f(x) c = C() c_0 = c[0] f(c_0) diff --git a/python/testSrc/com/jetbrains/python/PyTypeTest.java b/python/testSrc/com/jetbrains/python/PyTypeTest.java index dfe68303ca19..8f6d2026ba22 100644 --- a/python/testSrc/com/jetbrains/python/PyTypeTest.java +++ b/python/testSrc/com/jetbrains/python/PyTypeTest.java @@ -171,7 +171,7 @@ public class PyTypeTest extends PyTestCase { } public void testSubscriptType() { - doTest("Union[int, List[int]]", + doTest("int", "l = [1, 2, 3]; expr = l[0]"); } @@ -621,7 +621,7 @@ public class PyTypeTest extends PyTestCase { } public void testFunctionTypeAsUnificationArgumentWithSubscription() { - doTest("Union[int, List[int], str, unicode]", + doTest("Union[int, str, unicode]", "def map2(f, xs):\n" + " '''\n" + " :type f: (T) -> V | None\n" + diff --git a/python/testSrc/com/jetbrains/python/PyTypingTest.java b/python/testSrc/com/jetbrains/python/PyTypingTest.java index 7280891215be..a919ffd0994f 100644 --- a/python/testSrc/com/jetbrains/python/PyTypingTest.java +++ b/python/testSrc/com/jetbrains/python/PyTypingTest.java @@ -1,5 +1,5 @@ /* - * Copyright 2000-2014 JetBrains s.r.o. + * Copyright 2000-2016 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. @@ -642,7 +642,53 @@ public class PyTypingTest extends PyTestCase { "Type2 = Union[int, Type1]\n" + "\n" + "expr = None # type: Type1"); + } + // PY-19858 + public void testGetListItemByIntegral() { + doTest("list", + "from typing import List\n" + + "\n" + + "def foo(x: List[List]):\n" + + " expr = x[0]\n"); + } + + // PY-19858 + public void testGetListItemByIndirectIntegral() { + doTest("list", + "from typing import List\n" + + "\n" + + "def foo(x: List[List]):\n" + + " y = 0\n" + + " expr = x[y]\n"); + } + + // PY-19858 + public void testGetSublistBySlice() { + doTest("List[list]", + "from typing import List\n" + + "\n" + + "def foo(x: List[List]):\n" + + " expr = x[1:3]\n"); + } + + // PY-19858 + public void testGetSublistByIndirectSlice() { + doTest("List[list]", + "from typing import List\n" + + "\n" + + "def foo(x: List[List]):\n" + + " y = slice(1, 3)\n" + + " expr = x[y]\n"); + } + + // PY-19858 + public void testGetListItemByUnknown() { + doTest("Union[list, List[list]]", + "from typing import List\n" + + "\n" + + "def foo(x: List[List]):\n" + + " expr = x[y]\n"); } private void doTestNoInjectedText(@NotNull String text) {