Fix type inference for types with Final qualifier (PEP 591) (PY-34945)

GitOrigin-RevId: 1b2b580273df4d6edc0411849183b142392745b8
This commit is contained in:
Semyon Proshev
2019-07-02 06:52:16 +03:00
committed by intellij-monorepo-bot
parent c6e4181228
commit 0c63a3d7ca
7 changed files with 104 additions and 27 deletions
@@ -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
@@ -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<String> GENERIC_CLASSES = ImmutableSet.<String>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<PyType> 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<PyType> getFinalType(@NotNull PsiElement resolved, @NotNull Context context) {
if (resolved instanceof PySubscriptionExpression) {
final PySubscriptionExpression subscriptionExpr = (PySubscriptionExpression)resolved;
final Collection<String> 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)) {
@@ -1,2 +0,0 @@
class Protocol:
pass
@@ -1,2 +0,0 @@
class Protocol:
pass
@@ -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<TypeEvalContext> getTypeEvalContexts(@NotNull PyExpression element) {
return ImmutableList.of(TypeEvalContext.codeAnalysis(element.getProject(), element.getContainingFile()).withTracing(),
TypeEvalContext.userInitiated(element.getProject(), element.getContainingFile()).withTracing());
@@ -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<? extends PyInspection> getInspectionClass() {
@@ -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)