From 0c63a3d7caaec8885e98bbafd6861b7631239961 Mon Sep 17 00:00:00 2001 From: Semyon Proshev Date: Fri, 24 May 2019 19:37:08 +0300 Subject: [PATCH] Fix type inference for types with `Final` qualifier (PEP 591) (PY-34945) GitOrigin-RevId: 1b2b580273df4d6edc0411849183b142392745b8 --- .../third_party/2and3/typing_extensions.pyi | 37 ++++++++++++++ .../typing/PyTypingTypeProvider.java | 30 ++++++++++-- .../PyProtocolInspection/typing_extensions.py | 2 - .../typing_extensions.py | 2 - .../com/jetbrains/python/PyTypeTest.java | 48 +++++++++++++++---- .../inspections/PyProtocolInspectionTest.java | 9 ---- .../jetbrains/python/tools/PyTypeShedSync.kts | 3 +- 7 files changed, 104 insertions(+), 27 deletions(-) create mode 100644 python/helpers/typeshed/third_party/2and3/typing_extensions.pyi delete mode 100644 python/testData/inspections/PyProtocolInspection/typing_extensions.py delete mode 100644 python/testData/types/GenericTypingProtocolExt/typing_extensions.py diff --git a/python/helpers/typeshed/third_party/2and3/typing_extensions.pyi b/python/helpers/typeshed/third_party/2and3/typing_extensions.pyi new file mode 100644 index 000000000000..8ac31fc35462 --- /dev/null +++ b/python/helpers/typeshed/third_party/2and3/typing_extensions.pyi @@ -0,0 +1,37 @@ +import sys +from typing import Callable +from typing import ClassVar as ClassVar +from typing import ContextManager as ContextManager +from typing import Counter as Counter +from typing import DefaultDict as DefaultDict +from typing import Deque as Deque +from typing import NewType as NewType +from typing import NoReturn as NoReturn +from typing import overload as overload +from typing import Text as Text +from typing import Type as Type +from typing import TYPE_CHECKING as TYPE_CHECKING +from typing import TypeVar, Any + +_F = TypeVar('_F', bound=Callable[..., Any]) +_TC = TypeVar('_TC', bound=Type[object]) +class _SpecialForm: + def __getitem__(self, typeargs: Any) -> Any: ... +def runtime(cls: _TC) -> _TC: ... +Protocol: _SpecialForm = ... +Final: _SpecialForm = ... +def final(f: _F) -> _F: ... +Literal: _SpecialForm = ... + +if sys.version_info >= (3, 3): + from typing import ChainMap as ChainMap + +if sys.version_info >= (3, 5): + from typing import AsyncIterable as AsyncIterable + from typing import AsyncIterator as AsyncIterator + from typing import AsyncContextManager as AsyncContextManager + from typing import Awaitable as Awaitable + from typing import Coroutine as Coroutine + +if sys.version_info >= (3, 6): + from typing import AsyncGenerator as AsyncGenerator diff --git a/python/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.java b/python/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.java index 15f47074a0af..798f5e9b9afc 100644 --- a/python/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.java +++ b/python/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.java @@ -88,6 +88,8 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { public static final String UNION = "typing.Union"; public static final String OPTIONAL = "typing.Optional"; public static final String NO_RETURN = "typing.NoReturn"; + private static final String FINAL = "typing.Final"; + private static final String FINAL_EXT = "typing_extensions.Final"; private static final String PY2_FILE_TYPE = "typing.BinaryIO"; private static final String PY3_BINARY_FILE_TYPE = "typing.BinaryIO"; @@ -120,10 +122,10 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { public static final ImmutableSet GENERIC_CLASSES = ImmutableSet.builder() // special forms - .add(TUPLE, GENERIC, PROTOCOL, CALLABLE, TYPE, CLASS_VAR) + .add(TUPLE, GENERIC, PROTOCOL, CALLABLE, TYPE, CLASS_VAR, FINAL) // type aliases .add(UNION, OPTIONAL, LIST, DICT, DEFAULT_DICT, SET, FROZEN_SET, COUNTER, DEQUE, CHAIN_MAP) - .add(PROTOCOL_EXT) + .add(PROTOCOL_EXT, FINAL_EXT) .build(); /** @@ -147,12 +149,13 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { .add(DEFAULT_DICT) .add(SET) .add(FROZEN_SET) - .add(PROTOCOL) + .add(PROTOCOL, PROTOCOL_EXT) .add(CLASS_VAR) .add(COUNTER) .add(DEQUE) .add(CHAIN_MAP) .add(NO_RETURN) + .add(FINAL, FINAL_EXT) .build(); @Nullable @@ -806,6 +809,10 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { if (classObjType != null) { return Ref.create(addTypeVarAlias(classObjType.get(), alias)); } + final Ref finalType = getFinalType(resolved, context); + if (finalType != null) { + return finalType; + } final PyType parameterizedType = getParameterizedType(resolved, context); if (parameterizedType != null) { return Ref.create(parameterizedType); @@ -947,6 +954,23 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { return null; } + @Nullable + private static Ref getFinalType(@NotNull PsiElement resolved, @NotNull Context context) { + if (resolved instanceof PySubscriptionExpression) { + final PySubscriptionExpression subscriptionExpr = (PySubscriptionExpression)resolved; + + final Collection operandNames = resolveToQualifiedNames(subscriptionExpr.getOperand(), context.getTypeContext()); + if (ContainerUtil.exists(operandNames, name -> name.equals(FINAL) || name.equals(FINAL_EXT))) { + final PyExpression indexExpr = subscriptionExpr.getIndexExpression(); + if (indexExpr != null) { + return getType(indexExpr, context); + } + } + } + + return null; + } + @Nullable private static PyExpression getAnnotationValue(@NotNull PyAnnotationOwner owner, @NotNull TypeEvalContext context) { if (context.maySwitchToAST(owner)) { diff --git a/python/testData/inspections/PyProtocolInspection/typing_extensions.py b/python/testData/inspections/PyProtocolInspection/typing_extensions.py deleted file mode 100644 index 4b38e92012b3..000000000000 --- a/python/testData/inspections/PyProtocolInspection/typing_extensions.py +++ /dev/null @@ -1,2 +0,0 @@ -class Protocol: - pass \ No newline at end of file diff --git a/python/testData/types/GenericTypingProtocolExt/typing_extensions.py b/python/testData/types/GenericTypingProtocolExt/typing_extensions.py deleted file mode 100644 index 4b38e92012b3..000000000000 --- a/python/testData/types/GenericTypingProtocolExt/typing_extensions.py +++ /dev/null @@ -1,2 +0,0 @@ -class Protocol: - pass \ No newline at end of file diff --git a/python/testSrc/com/jetbrains/python/PyTypeTest.java b/python/testSrc/com/jetbrains/python/PyTypeTest.java index 23a6f2ada436..d1698bb7088e 100644 --- a/python/testSrc/com/jetbrains/python/PyTypeTest.java +++ b/python/testSrc/com/jetbrains/python/PyTypeTest.java @@ -3233,16 +3233,16 @@ public class PyTypeTest extends PyTestCase { public void testGenericTypingProtocolExt() { runWithLanguageLevel( LanguageLevel.PYTHON37, - () -> doMultiFileTest("int", - "from typing_extensions import Protocol\n" + - "from typing import TypeVar\n" + - "T = TypeVar(\"T\")\n" + - "class MyProto1(Protocol[T]):\n" + - " def func(self) -> T:\n" + - " pass\n" + - "class MyClass1(MyProto1[int]):\n" + - " pass\n" + - "expr = MyClass1().func()") + () -> doTest("int", + "from typing_extensions import Protocol\n" + + "from typing import TypeVar\n" + + "T = TypeVar(\"T\")\n" + + "class MyProto1(Protocol[T]):\n" + + " def func(self) -> T:\n" + + " pass\n" + + "class MyClass1(MyProto1[int]):\n" + + " pass\n" + + "expr = MyClass1().func()") ); } @@ -3389,6 +3389,34 @@ public class PyTypeTest extends PyTestCase { ); } + // PY-34945 + public void testFinal() { + runWithLanguageLevel( + LanguageLevel.PYTHON35, + () -> { + doTest("int", + "from typing_extensions import Final\n" + + "expr: Final[int] = undefined"); + + doTest("int", + "from typing_extensions import Final\n" + + "expr: Final = 5"); + + doTest("int", + "from typing_extensions import Final\n" + + "expr: Final[int]"); + } + ); + + doTest("int", + "from typing_extensions import Final\n" + + "expr = undefined # type: Final[int]"); + + doTest("int", + "from typing_extensions import Final\n" + + "expr = 5 # type: Final"); + } + private static List getTypeEvalContexts(@NotNull PyExpression element) { return ImmutableList.of(TypeEvalContext.codeAnalysis(element.getProject(), element.getContainingFile()).withTracing(), TypeEvalContext.userInitiated(element.getProject(), element.getContainingFile()).withTracing()); diff --git a/python/testSrc/com/jetbrains/python/inspections/PyProtocolInspectionTest.java b/python/testSrc/com/jetbrains/python/inspections/PyProtocolInspectionTest.java index 92f1a785f4d8..278d59aa0dd1 100644 --- a/python/testSrc/com/jetbrains/python/inspections/PyProtocolInspectionTest.java +++ b/python/testSrc/com/jetbrains/python/inspections/PyProtocolInspectionTest.java @@ -1,11 +1,9 @@ // Copyright 2000-2018 JetBrains s.r.o. Use of this source code is governed by the Apache 2.0 license that can be found in the LICENSE file. package com.jetbrains.python.inspections; -import com.intellij.testFramework.LightProjectDescriptor; import com.jetbrains.python.fixtures.PyInspectionTestCase; import com.jetbrains.python.psi.LanguageLevel; import org.jetbrains.annotations.NotNull; -import org.jetbrains.annotations.Nullable; public class PyProtocolInspectionTest extends PyInspectionTestCase { @@ -36,7 +34,6 @@ public class PyProtocolInspectionTest extends PyInspectionTestCase { // PY-26628 public void testProtocolExtBases() { - myFixture.copyFileToProject(getTestCaseDirectory() + "typing_extensions.py", "typing_extensions.py"); doTest(); } @@ -45,12 +42,6 @@ public class PyProtocolInspectionTest extends PyInspectionTestCase { runWithLanguageLevel(LanguageLevel.PYTHON37, () -> super.doTest()); } - @Nullable - @Override - protected LightProjectDescriptor getProjectDescriptor() { - return ourPy3Descriptor; - } - @NotNull @Override protected Class getInspectionClass() { diff --git a/python/tools/src/com/jetbrains/python/tools/PyTypeShedSync.kts b/python/tools/src/com/jetbrains/python/tools/PyTypeShedSync.kts index c49c4ee577ab..eb0c4d9bff45 100644 --- a/python/tools/src/com/jetbrains/python/tools/PyTypeShedSync.kts +++ b/python/tools/src/com/jetbrains/python/tools/PyTypeShedSync.kts @@ -22,7 +22,8 @@ sync(repo, bundled) val whiteList = setOf("typing", "six", "__builtin__", "builtins", "exceptions", "types", "datetime", "functools", "shutil", "re", "time", "argparse", "uuid", "threading", "signal", "collections", "subprocess", "math", "queue", "socket", "sqlite3", "attr", - "pathlib", "io", "_io", "itertools", "ssl", "multiprocessing", "asyncio", "mock", "unittest", "_importlib_modulespec") + "pathlib", "io", "_io", "itertools", "ssl", "multiprocessing", "asyncio", "mock", "unittest", "_importlib_modulespec", + "typing_extensions") clean(topLevelPackages(bundled), whiteList)