diff --git a/python/src/com/jetbrains/python/codeInsight/intentions/PyTypeHintGenerationUtil.java b/python/src/com/jetbrains/python/codeInsight/intentions/PyTypeHintGenerationUtil.java index 15592d3af782..1327854dcdce 100644 --- a/python/src/com/jetbrains/python/codeInsight/intentions/PyTypeHintGenerationUtil.java +++ b/python/src/com/jetbrains/python/codeInsight/intentions/PyTypeHintGenerationUtil.java @@ -349,6 +349,12 @@ public class PyTypeHintGenerationUtil { } collectImportTargetsFromType(callableType.getReturnType(context), context, symbols, typingTypes); } + else if (type instanceof PyGenericType) { + final PyTargetExpression target = as(type.getDeclarationElement(), PyTargetExpression.class); + if (target != null) { + symbols.add(target); + } + } if (type instanceof PyInstantiableType && ((PyInstantiableType)type).isDefinition()) { typingTypes.add("Type"); } diff --git a/python/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.java b/python/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.java index f2df1e78668b..949797b52bd7 100644 --- a/python/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.java +++ b/python/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.java @@ -5,6 +5,7 @@ import com.google.common.collect.ImmutableMap; import com.google.common.collect.ImmutableSet; import com.google.common.collect.Lists; import com.google.common.collect.Sets; +import com.intellij.openapi.util.Pair; import com.intellij.openapi.util.Ref; import com.intellij.psi.PsiElement; import com.intellij.psi.PsiFile; @@ -721,8 +722,8 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { private static Ref getType(@NotNull PyExpression expression, @NotNull Context context) { final List members = Lists.newArrayList(); boolean foundAny = false; - for (PsiElement resolved : tryResolving(expression, context.getTypeContext())) { - final Ref typeRef = getTypeForResolvedElement(resolved, context); + for (Pair pair : tryResolvingWithAliases(expression, context.getTypeContext())) { + final Ref typeRef = getTypeForResolvedElement(pair.getFirst(), pair.getSecond(), context); if (typeRef != null) { final PyType type = typeRef.get(); if (type == null) { @@ -736,7 +737,9 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { } @Nullable - private static Ref getTypeForResolvedElement(@NotNull PsiElement resolved, @NotNull Context context) { + private static Ref getTypeForResolvedElement(@Nullable PyTargetExpression alias, + @NotNull PsiElement resolved, + @NotNull Context context) { if (context.getExpressionCache().contains(resolved)) { // Recursive types are not yet supported return null; @@ -758,7 +761,7 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { } final Ref classObjType = getClassObjectType(resolved, context); if (classObjType != null) { - return classObjType; + return Ref.create(addTypeVarAlias(classObjType.get(), alias)); } final PyType parameterizedType = getParameterizedType(resolved, context); if (parameterizedType != null) { @@ -770,7 +773,7 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { } final PyType genericType = getGenericTypeFromTypeVar(resolved, context); if (genericType != null) { - return Ref.create(genericType); + return Ref.create(addTypeVarAlias(genericType, alias)); } final PyType stringBasedType = getStringLiteralType(resolved, context); if (stringBasedType != null) { @@ -791,6 +794,15 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { } } + @Nullable + private static PyType addTypeVarAlias(@Nullable PyType type, @Nullable PyTargetExpression alias) { + final PyGenericType typeVar = as(type, PyGenericType.class); + if (typeVar != null) { + return new PyGenericType(typeVar.getName(), typeVar.getBound(), typeVar.isDefinition(), alias); + } + return type; + } + @Nullable private static Ref getClassObjectType(@NotNull PsiElement resolved, @NotNull Context context) { if (resolved instanceof PySubscriptionExpression) { @@ -1074,7 +1086,13 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { @NotNull private static List tryResolving(@NotNull PyExpression expression, @NotNull TypeEvalContext context) { - final List elements = Lists.newArrayList(); + return ContainerUtil.map(tryResolvingWithAliases(expression, context), x -> x.getSecond()); + } + + @NotNull + private static List> tryResolvingWithAliases(@NotNull PyExpression expression, + @NotNull TypeEvalContext context) { + final List> elements = Lists.newArrayList(); if (expression instanceof PyReferenceExpression) { final List results; if (context.maySwitchToAST(expression)) { @@ -1090,14 +1108,14 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { if (PyUtil.isInit(function)) { final PyClass cls = function.getContainingClass(); if (cls != null) { - elements.add(cls); + elements.add(Pair.create(null, cls)); continue; } } } final String name = element != null ? getQualifiedName(element) : null; if (name != null && OPAQUE_NAMES.contains(name)) { - elements.add(element); + elements.add(Pair.create(null, element)); continue; } // Presumably, a TypeVar definition or a type alias @@ -1111,7 +1129,7 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { assignedValue = PyTypingAliasStubType.getAssignedValueStubLike(targetExpr); } if (assignedValue != null) { - elements.add(assignedValue); + elements.add(Pair.create(targetExpr, assignedValue)); continue; } } @@ -1121,16 +1139,16 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { final QualifiedName osPathLikeQName = QualifiedName.fromComponents("os", PyNames.PATH_LIKE); final PsiElement osPathLike = PyResolveImportUtil.resolveTopLevelMember(osPathLikeQName, PyResolveImportUtil.fromFoothold(element)); if (osPathLike != null) { - elements.add(osPathLike); + elements.add(Pair.create(null, osPathLike)); continue; } } if (element != null) { - elements.add(element); + elements.add(Pair.create(null, element)); } } } - return !elements.isEmpty() ? elements : Collections.singletonList(expression); + return !elements.isEmpty() ? elements : Collections.singletonList(Pair.create(null, expression)); } @NotNull diff --git a/python/src/com/jetbrains/python/psi/types/PyGenericType.java b/python/src/com/jetbrains/python/psi/types/PyGenericType.java index 4db2801039c5..a577531ea0e7 100644 --- a/python/src/com/jetbrains/python/psi/types/PyGenericType.java +++ b/python/src/com/jetbrains/python/psi/types/PyGenericType.java @@ -123,13 +123,13 @@ public class PyGenericType implements PyType, PyInstantiableType @NotNull @Override public PyGenericType toInstance() { - return myIsDefinition ? new PyGenericType(myName, myBound, false) : this; + return myIsDefinition ? new PyGenericType(myName, myBound, false, myTargetExpression) : this; } @NotNull @Override public PyGenericType toClass() { - return myIsDefinition ? this : new PyGenericType(myName, myBound, true); + return myIsDefinition ? this : new PyGenericType(myName, myBound, true, myTargetExpression); } @Override diff --git a/python/testData/intentions/PyAnnotateVariableTypeIntentionTest/AnnotationTypeVarInOtherFile/lib.py b/python/testData/intentions/PyAnnotateVariableTypeIntentionTest/AnnotationTypeVarInOtherFile/lib.py new file mode 100644 index 000000000000..f58d06f49978 --- /dev/null +++ b/python/testData/intentions/PyAnnotateVariableTypeIntentionTest/AnnotationTypeVarInOtherFile/lib.py @@ -0,0 +1,5 @@ +from typing import TypeVar + +T = TypeVar('T') + +target: T = 42 \ No newline at end of file diff --git a/python/testData/intentions/PyAnnotateVariableTypeIntentionTest/AnnotationTypeVarInOtherFile/main.py b/python/testData/intentions/PyAnnotateVariableTypeIntentionTest/AnnotationTypeVarInOtherFile/main.py new file mode 100644 index 000000000000..3bab0a10b432 --- /dev/null +++ b/python/testData/intentions/PyAnnotateVariableTypeIntentionTest/AnnotationTypeVarInOtherFile/main.py @@ -0,0 +1,3 @@ +from lib import target + +var = target \ No newline at end of file diff --git a/python/testData/intentions/PyAnnotateVariableTypeIntentionTest/AnnotationTypeVarInOtherFile/main_after.py b/python/testData/intentions/PyAnnotateVariableTypeIntentionTest/AnnotationTypeVarInOtherFile/main_after.py new file mode 100644 index 000000000000..45ab51682ea1 --- /dev/null +++ b/python/testData/intentions/PyAnnotateVariableTypeIntentionTest/AnnotationTypeVarInOtherFile/main_after.py @@ -0,0 +1,3 @@ +from lib import target, T + +var: [T] = target \ No newline at end of file diff --git a/python/testSrc/com/jetbrains/python/intentions/PyAnnotateVariableTypeIntentionTest.java b/python/testSrc/com/jetbrains/python/intentions/PyAnnotateVariableTypeIntentionTest.java index 9a79ec6e5e27..619467220034 100644 --- a/python/testSrc/com/jetbrains/python/intentions/PyAnnotateVariableTypeIntentionTest.java +++ b/python/testSrc/com/jetbrains/python/intentions/PyAnnotateVariableTypeIntentionTest.java @@ -243,6 +243,10 @@ public class PyAnnotateVariableTypeIntentionTest extends PyIntentionTestCase { doMultiFileAnnotationTest(); } + public void testAnnotationTypeVarInOtherFile() { + doMultiFileAnnotationTest(); + } + public void testAnnotationCollectionsNamedTupleInOtherFile() { doMultiFileAnnotationTest(); }