PY-16412 Properly import named tuples for their use in type hints

For that I introduced a new method -- getElement() in PyType
that allows to find out the element that can be used to reference
this type according to PEP 484 perspective, e.g. a target assignment
for TypeVar's and NamedTuples and PSI class themselves for class types.

Also, class object types for named tuples are wrapped in typing.Type[]
as expected.
This commit is contained in:
Mikhail Golubev
2018-02-05 21:41:59 +03:00
committed by Andrey Vlasovskikh
parent d1ae008593
commit af64e90b62
23 changed files with 177 additions and 38 deletions
@@ -20,6 +20,7 @@ import com.intellij.psi.PsiElement;
import com.intellij.util.ProcessingContext;
import com.jetbrains.python.psi.AccessDirection;
import com.jetbrains.python.psi.PyExpression;
import com.jetbrains.python.psi.PyQualifiedNameOwner;
import com.jetbrains.python.psi.resolve.PyResolveContext;
import com.jetbrains.python.psi.resolve.RatedResolveResult;
import org.jetbrains.annotations.NotNull;
@@ -35,6 +36,11 @@ import java.util.Set;
*/
public interface PyType {
@Nullable
default PyQualifiedNameOwner getDeclarationElement() {
return null;
}
/**
* Resolves an attribute of type.
*
@@ -21,6 +21,7 @@ import com.intellij.util.containers.ContainerUtil;
import com.jetbrains.python.PyBundle;
import com.jetbrains.python.codeInsight.imports.AddImportHelper;
import com.jetbrains.python.codeInsight.imports.AddImportHelper.ImportPriority;
import com.jetbrains.python.codeInsight.stdlib.PyNamedTupleType;
import com.jetbrains.python.codeInsight.typing.PyTypingTypeProvider;
import com.jetbrains.python.psi.*;
import com.jetbrains.python.psi.impl.PyPsiUtils;
@@ -280,11 +281,11 @@ public class PyTypeHintGenerationUtil {
private static void addImportsForTypeAnnotations(@NotNull List<PyType> types,
@NotNull TypeEvalContext context,
@NotNull PsiFile file) {
final Set<PyClass> classes = new HashSet<>();
final Set<PsiNamedElement> symbols = new HashSet<>();
final Set<String> namesFromTyping = new HashSet<>();
for (PyType type : types) {
collectImportTargetsFromType(type, context, classes, namesFromTyping);
collectImportTargetsFromType(type, context, symbols, namesFromTyping);
}
final boolean builtinTyping = LanguageLevel.forElement(file).isAtLeast(LanguageLevel.PYTHON35);
@@ -293,24 +294,30 @@ public class PyTypeHintGenerationUtil {
AddImportHelper.addOrUpdateFromImportStatement(file, "typing", name, null, priority, null);
}
for (PyClass pyClass : classes) {
PyClassRefactoringUtil.insertImport(file, pyClass, null, true);
for (PsiNamedElement symbol : symbols) {
PyClassRefactoringUtil.insertImport(file, symbol, null, true);
}
}
private static void collectImportTargetsFromType(@Nullable PyType type,
@NotNull TypeEvalContext context,
@NotNull Set<PyClass> classes,
@NotNull Set<String> names) {
@NotNull Set<PsiNamedElement> symbols,
@NotNull Set<String> typingTypes) {
if (type == null) {
names.add("Any");
typingTypes.add("Any");
}
else if (type instanceof PyUnionType) {
final Collection<PyType> members = ((PyUnionType)type).getMembers();
final boolean isOptional = members.size() == 2 && members.contains(PyNoneType.INSTANCE);
names.add(isOptional ? "Optional" : "Union");
typingTypes.add(isOptional ? "Optional" : "Union");
for (PyType pyType : members) {
collectImportTargetsFromType(pyType, context, classes, names);
collectImportTargetsFromType(pyType, context, symbols, typingTypes);
}
}
else if (type instanceof PyNamedTupleType) {
final PyQualifiedNameOwner element = type.getDeclarationElement();
if (element instanceof PsiNamedElement) {
symbols.add((PsiNamedElement)element);
}
}
else if (type instanceof PyCollectionType) {
@@ -318,32 +325,32 @@ public class PyTypeHintGenerationUtil {
final PyClass pyClass = ((PyCollectionTypeImpl)type).getPyClass();
final String typingCollectionName = PyTypingTypeProvider.TYPING_COLLECTION_CLASSES.get(pyClass.getQualifiedName());
if (typingCollectionName != null && type.isBuiltin()) {
names.add(typingCollectionName);
typingTypes.add(typingCollectionName);
}
else {
classes.add(pyClass);
symbols.add(pyClass);
}
}
else if (type instanceof PyTupleType) {
names.add("Tuple");
typingTypes.add("Tuple");
}
for (PyType pyType : ((PyCollectionType)type).getElementTypes()) {
collectImportTargetsFromType(pyType, context, classes, names);
collectImportTargetsFromType(pyType, context, symbols, typingTypes);
}
}
else if (type instanceof PyClassType) {
classes.add(((PyClassType)type).getPyClass());
symbols.add(((PyClassType)type).getPyClass());
}
else if (type instanceof PyCallableType) {
names.add("Callable");
typingTypes.add("Callable");
final PyCallableType callableType = (PyCallableType)type;
for (PyCallableParameter parameter : ContainerUtil.notNullize(callableType.getParameters(context))) {
collectImportTargetsFromType(parameter.getType(context), context, classes, names);
collectImportTargetsFromType(parameter.getType(context), context, symbols, typingTypes);
}
collectImportTargetsFromType(callableType.getReturnType(context), context, classes, names);
collectImportTargetsFromType(callableType.getReturnType(context), context, symbols, typingTypes);
}
if (type instanceof PyInstantiableType && ((PyInstantiableType)type).isDefinition()) {
names.add("Type");
typingTypes.add("Type");
}
}
@@ -6,10 +6,7 @@ import com.intellij.psi.PsiElement;
import com.intellij.util.ArrayUtil;
import com.intellij.util.ProcessingContext;
import com.intellij.util.containers.ContainerUtil;
import java.util.HashMap;
import com.jetbrains.python.psi.PyCallSiteExpression;
import com.jetbrains.python.psi.PyClass;
import com.jetbrains.python.psi.PyExpression;
import com.jetbrains.python.psi.*;
import com.jetbrains.python.psi.types.*;
import one.util.streamex.StreamEx;
import org.jetbrains.annotations.NotNull;
@@ -32,12 +29,22 @@ public class PyNamedTupleType extends PyTupleType implements PyCallableType {
private final DefinitionLevel myDefinitionLevel;
private final boolean myTyped;
private final PyTargetExpression myTargetExpression;
public PyNamedTupleType(@NotNull PyClass tupleClass,
@NotNull String name,
@NotNull LinkedHashMap<String, FieldTypeAndDefaultValue> fields,
@NotNull DefinitionLevel definitionLevel,
boolean typed) {
this(tupleClass, name, fields, definitionLevel, typed, null);
}
public PyNamedTupleType(@NotNull PyClass tupleClass,
@NotNull String name,
@NotNull LinkedHashMap<String, FieldTypeAndDefaultValue> fields,
@NotNull DefinitionLevel definitionLevel,
boolean typed,
@Nullable PyTargetExpression target) {
super(tupleClass,
Collections.unmodifiableList(ContainerUtil.map(fields.values(), typeAndValue -> typeAndValue.getType())),
false,
@@ -47,6 +54,13 @@ public class PyNamedTupleType extends PyTupleType implements PyCallableType {
myName = name;
myDefinitionLevel = definitionLevel;
myTyped = typed;
myTargetExpression = target;
}
@NotNull
@Override
public PyQualifiedNameOwner getDeclarationElement() {
return myTargetExpression;
}
@Override
@@ -74,7 +88,7 @@ public class PyNamedTupleType extends PyTupleType implements PyCallableType {
@Override
public PyNamedTupleType getCallType(@NotNull TypeEvalContext context, @NotNull PyCallSiteExpression callSite) {
if (myDefinitionLevel == DefinitionLevel.NT_FUNCTION) {
return new PyNamedTupleType(myClass, myName, myFields, DefinitionLevel.NEW_TYPE, myTyped);
return new PyNamedTupleType(myClass, myName, myFields, DefinitionLevel.NEW_TYPE, myTyped, myTargetExpression);
}
else if (myDefinitionLevel == DefinitionLevel.NEW_TYPE) {
return getCallDefinitionType(callSite, context);
@@ -87,7 +101,7 @@ public class PyNamedTupleType extends PyTupleType implements PyCallableType {
@Override
public PyNamedTupleType toInstance() {
return myDefinitionLevel == DefinitionLevel.NEW_TYPE
? new PyNamedTupleType(myClass, myName, myFields, DefinitionLevel.INSTANCE, myTyped)
? new PyNamedTupleType(myClass, myName, myFields, DefinitionLevel.INSTANCE, myTyped, myTargetExpression)
: this;
}
@@ -96,7 +110,7 @@ public class PyNamedTupleType extends PyTupleType implements PyCallableType {
public PyNamedTupleType toClass() {
return myDefinitionLevel == DefinitionLevel.INSTANCE
? this
: new PyNamedTupleType(myClass, myName, myFields, DefinitionLevel.NEW_TYPE, myTyped);
: new PyNamedTupleType(myClass, myName, myFields, DefinitionLevel.NEW_TYPE, myTyped, myTargetExpression);
}
@Override
@@ -165,7 +179,7 @@ public class PyNamedTupleType extends PyTupleType implements PyCallableType {
}
}
return new PyNamedTupleType(myClass, myName, newFields, myDefinitionLevel, false);
return new PyNamedTupleType(myClass, myName, newFields, myDefinitionLevel, false, myTargetExpression);
}
return this;
@@ -539,7 +539,8 @@ public class PyStdlibTypeProvider extends PyTypeProviderBase {
stub.getName(),
parseNamedTupleFields(referenceTarget, fields, context),
definitionLevel,
ContainerUtil.find(fields.values(), Optional::isPresent) != null);
ContainerUtil.find(fields.values(), Optional::isPresent) != null,
as(referenceTarget, PyTargetExpression.class));
}
@Nullable
@@ -227,7 +227,17 @@ public class PyTypeModelBuilder {
myVisited.put(type, null); //mark as evaluating
TypeModel result = null;
if (type instanceof PyNamedTupleType) {
if (type instanceof PyInstantiableType && ((PyInstantiableType)type).isDefinition()) {
final PyInstantiableType instanceType = ((PyInstantiableType)type).toInstance();
// Special case: render Type[type] as just type
if (type instanceof PyClassType && instanceType.equals(PyBuiltinCache.getInstance(((PyClassType)type).getPyClass()).getTypeType())) {
result = NamedType.nameOrAny(type);
}
else {
result = new ClassObjectType(build(instanceType, allowUnions));
}
}
else if (type instanceof PyNamedTupleType) {
result = NamedType.nameOrAny(type);
}
else if (type instanceof PyTupleType) {
@@ -281,16 +291,6 @@ public class PyTypeModelBuilder {
else if (type instanceof PyCallableType && !(type instanceof PyClassLikeType)) {
result = buildCallable((PyCallableType)type);
}
else if (type instanceof PyInstantiableType && ((PyInstantiableType)type).isDefinition()) {
final PyInstantiableType instanceType = ((PyInstantiableType)type).toInstance();
// Special case: render Type[type] as just type
if (type instanceof PyClassType && instanceType.equals(PyBuiltinCache.getInstance(((PyClassType)type).getPyClass()).getTypeType())) {
result = NamedType.nameOrAny(type);
}
else {
result = new ClassObjectType(build(instanceType, allowUnions));
}
}
else if (type instanceof PyGenericType) {
result = new GenericType(type.getName());
}
@@ -85,6 +85,13 @@ public class PyClassTypeImpl extends UserDataHolderBase implements PyClassType {
return myClass;
}
@NotNull
@Override
public PyQualifiedNameOwner getDeclarationElement() {
return getPyClass();
}
/**
* @return whether this type refers to an instance or a definition of the class.
*/
@@ -20,6 +20,7 @@ import com.intellij.util.ArrayUtil;
import com.intellij.util.ProcessingContext;
import com.jetbrains.python.psi.AccessDirection;
import com.jetbrains.python.psi.PyExpression;
import com.jetbrains.python.psi.PyTargetExpression;
import com.jetbrains.python.psi.resolve.PyResolveContext;
import com.jetbrains.python.psi.resolve.RatedResolveResult;
import org.jetbrains.annotations.NotNull;
@@ -34,15 +35,27 @@ public class PyGenericType implements PyType, PyInstantiableType<PyGenericType>
@NotNull private final String myName;
@Nullable private final PyType myBound;
private boolean myIsDefinition = false;
private PyTargetExpression myTargetExpression;
public PyGenericType(@NotNull String name, @Nullable PyType bound) {
this(name, bound, false);
}
public PyGenericType(@NotNull String name, @Nullable PyType bound, boolean isDefinition) {
this(name, bound, isDefinition, null);
}
public PyGenericType(@NotNull String name, @Nullable PyType bound, boolean isDefinition, @Nullable PyTargetExpression target) {
myName = name;
myBound = bound;
myIsDefinition = isDefinition;
myTargetExpression = target;
}
@Nullable
@Override
public PyTargetExpression getDeclarationElement() {
return myTargetExpression;
}
@Nullable
@@ -0,0 +1,7 @@
from collections import namedtuple
MyTuple = namedtuple('MyTuple', ['foo'])
def func():
return MyTuple
@@ -0,0 +1,3 @@
from lib import func
va<caret>r = func()
@@ -0,0 +1,5 @@
from typing import Type
from lib import func, MyTuple
var: [Type[MyTuple]] = func()
@@ -0,0 +1,7 @@
from collections import namedtuple
MyTuple = namedtuple('MyTuple', ['foo'])
def func():
return MyTuple(foo=42)
@@ -0,0 +1,3 @@
from lib import func
va<caret>r = func()
@@ -0,0 +1,3 @@
from lib import func, MyTuple
var: [MyTuple] = func()
@@ -0,0 +1,7 @@
from typing import NamedTuple
MyTuple = NamedTuple('MyTuple', [('foo', int)])
def func():
return MyTuple
@@ -0,0 +1,3 @@
from lib import func
va<caret>r = func()
@@ -0,0 +1,5 @@
from typing import Type
from lib import func, MyTuple
var: [Type[MyTuple]] = func()
@@ -0,0 +1,9 @@
from typing import NamedTuple
class MyTuple(NamedTuple):
foo: int
def func():
return MyTuple(foo=42)
@@ -0,0 +1,3 @@
from lib import func
va<caret>r = func()
@@ -0,0 +1,3 @@
from lib import func, MyTuple
var: [MyTuple] = func()
@@ -0,0 +1,7 @@
from typing import NamedTuple
MyTuple = NamedTuple('MyTuple', [('foo', int)])
def func():
return MyTuple(foo=42)
@@ -0,0 +1,3 @@
from lib import func
va<caret>r = func()
@@ -0,0 +1,3 @@
from lib import func, MyTuple
var: [MyTuple] = func()
@@ -231,6 +231,26 @@ public class PyAnnotateVariableTypeIntentionTest extends PyIntentionTestCase {
doAnnotationTest();
}
public void testAnnotationTypingNamedTupleInOtherFile() {
doMultiFileAnnotationTest();
}
public void testAnnotationTypingNamedTupleClassInOtherFile() {
doMultiFileAnnotationTest();
}
public void testAnnotationTypingNamedTupleDirectInheritorInOtherFile() {
doMultiFileAnnotationTest();
}
public void testAnnotationCollectionsNamedTupleInOtherFile() {
doMultiFileAnnotationTest();
}
public void testAnnotationCollectionsNamedTupleClassInOtherFile() {
doMultiFileAnnotationTest();
}
public void testConflictWithAnnotationFunctionTypeIntention() {
doTest(LanguageLevel.PYTHON36);
}