diff --git a/python/psi-api/src/com/jetbrains/python/psi/PyAnnotationOwner.java b/python/psi-api/src/com/jetbrains/python/psi/PyAnnotationOwner.java index f8d6d5d9f75b..f114be63a266 100644 --- a/python/psi-api/src/com/jetbrains/python/psi/PyAnnotationOwner.java +++ b/python/psi-api/src/com/jetbrains/python/psi/PyAnnotationOwner.java @@ -15,12 +15,13 @@ */ package com.jetbrains.python.psi; +import com.intellij.psi.PsiElement; import org.jetbrains.annotations.Nullable; /** * @author Mikhail Golubev */ -public interface PyAnnotationOwner { +public interface PyAnnotationOwner extends PsiElement { @Nullable PyAnnotation getAnnotation(); diff --git a/python/psi-api/src/com/jetbrains/python/psi/PyTypeCommentOwner.java b/python/psi-api/src/com/jetbrains/python/psi/PyTypeCommentOwner.java index 55ee2d802050..0c9134850331 100644 --- a/python/psi-api/src/com/jetbrains/python/psi/PyTypeCommentOwner.java +++ b/python/psi-api/src/com/jetbrains/python/psi/PyTypeCommentOwner.java @@ -16,12 +16,13 @@ package com.jetbrains.python.psi; import com.intellij.psi.PsiComment; +import com.intellij.psi.PsiElement; import org.jetbrains.annotations.Nullable; /** * @author Mikhail Golubev */ -public interface PyTypeCommentOwner { +public interface PyTypeCommentOwner extends PsiElement { /** * Returns a special comment that follows element definition and starts with conventional "type:" prefix. * It is supposed to contain type annotation in PEP 484 compatible format. For further details see sections diff --git a/python/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.java b/python/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.java index c7b0bd06bf65..2673c12db103 100644 --- a/python/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.java +++ b/python/src/com/jetbrains/python/codeInsight/typing/PyTypingTypeProvider.java @@ -40,6 +40,7 @@ import com.jetbrains.python.psi.impl.PyBuiltinCache; import com.jetbrains.python.psi.impl.PyPsiUtils; import com.jetbrains.python.psi.impl.stubs.PyClassElementType; import com.jetbrains.python.psi.impl.stubs.PyTypingAliasStubType; +import com.jetbrains.python.psi.resolve.PyResolveContext; import com.jetbrains.python.psi.resolve.PyResolveImportUtil; import com.jetbrains.python.psi.resolve.PyResolveUtil; import com.jetbrains.python.psi.stubs.PyClassStub; @@ -203,8 +204,8 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { @Nullable private static Ref getParameterTypeFromAnnotation(@NotNull PyNamedParameter parameter, @NotNull TypeEvalContext context) { final Ref annotationValueTypeRef = Optional - .ofNullable(parameter.getAnnotationValue()) - .map(text -> getStringBasedType(text, parameter, context)) + .ofNullable(getAnnotationValue(parameter, context)) + .map(text -> getType(text, new Context(context))) .orElse(null); if (annotationValueTypeRef != null) { @@ -244,7 +245,7 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { public Ref getReturnType(@NotNull PyCallable callable, @NotNull TypeEvalContext context) { if (callable instanceof PyFunction) { final PyFunction function = (PyFunction)callable; - final PyExpression value = getReturnTypeAnnotation(function); + final PyExpression value = getReturnTypeAnnotation(function, context); if (value != null) { final Ref typeRef = getType(value, new Context(context)); if (typeRef != null) { @@ -259,10 +260,10 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { } @Nullable - private static PyExpression getReturnTypeAnnotation(@NotNull PyFunction function) { - final String annotation = function.getAnnotationValue(); - if (annotation != null) { - return PyUtil.createExpressionFromFragment(annotation, function); + private static PyExpression getReturnTypeAnnotation(@NotNull PyFunction function, TypeEvalContext context) { + final PyExpression returnAnnotation = getAnnotationValue(function, context); + if (returnAnnotation != null) { + return returnAnnotation; } final PyFunctionTypeAnnotation functionAnnotation = getFunctionTypeAnnotation(function); if (functionAnnotation != null) { @@ -305,9 +306,9 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { if (GENERIC.equals(target.getQualifiedName())) { return createTypingGenericType(); } - final String annotation = target.getAnnotationValue(); + final PyExpression annotation = getAnnotationValue(target, context); if (annotation != null) { - return Ref.deref(getStringBasedType(annotation, target, context)); + return Ref.deref(getType(annotation, new Context(context))); } final String comment = target.getTypeCommentAnnotation(); if (comment != null) { @@ -519,7 +520,7 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { if (genericType != null) { return Ref.create(genericType); } - final PyType stringBasedType = getStringBasedType(resolved, context); + final PyType stringBasedType = getStringLiteralType(resolved, context); if (stringBasedType != null) { return Ref.create(stringBasedType); } @@ -621,18 +622,25 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { } @Nullable - public static Ref getStringBasedType(@NotNull String contents, @NotNull PsiElement anchor, @NotNull TypeEvalContext context) { - return getStringBasedType(contents, anchor, new Context(context)); + private static PyExpression getAnnotationValue(@NotNull PyAnnotationOwner owner, @NotNull TypeEvalContext context) { + if (context.maySwitchToAST(owner)) { + final PyAnnotation annotation = owner.getAnnotation(); + if (annotation != null) { + return annotation.getValue(); + } + } + else { + final String annotationText = owner.getAnnotationValue(); + if (annotationText != null) { + return PyUtil.createExpressionFromFragment(annotationText, owner.getContainingFile()); + } + } + return null; } @Nullable - private static PyType getStringBasedType(@NotNull PsiElement element, @NotNull Context context) { - if (element instanceof PyStringLiteralExpression) { - // XXX: Requires switching from stub to AST - final String contents = ((PyStringLiteralExpression)element).getStringValue(); - return Ref.deref(getStringBasedType(contents, element, context)); - } - return null; + public static Ref getStringBasedType(@NotNull String contents, @NotNull PsiElement anchor, @NotNull TypeEvalContext context) { + return getStringBasedType(contents, anchor, new Context(context)); } @Nullable @@ -645,6 +653,16 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { return expr != null ? getType(expr, context) : null; } + @Nullable + private static PyType getStringLiteralType(@NotNull PsiElement element, @NotNull Context context) { + if (element instanceof PyStringLiteralExpression) { + // XXX: Requires switching from stub to AST + final String contents = ((PyStringLiteralExpression)element).getStringValue(); + return Ref.deref(getStringBasedType(contents, element, context)); + } + return null; + } + @Nullable private static Ref getVariableTypeCommentType(@NotNull String contents, @NotNull PsiElement anchor, @NotNull Context context) { final PyExpression expr = PyUtil.createExpressionFromFragment(contents, anchor); @@ -802,7 +820,14 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { private static List tryResolving(@NotNull PyExpression expression, @NotNull TypeEvalContext context) { final List elements = Lists.newArrayList(); if (expression instanceof PyReferenceExpression) { - final List results = tryResolvingOnStubs((PyReferenceExpression)expression, context); + final List results; + if (context.maySwitchToAST(expression)) { + final PyResolveContext resolveContext = PyResolveContext.noImplicits().withTypeEvalContext(context); + results = PyUtil.multiResolveTopPriority(expression, resolveContext); + } + else { + results = tryResolvingOnStubs((PyReferenceExpression)expression, context); + } for (PsiElement element : results) { if (element instanceof PyFunction) { final PyFunction function = (PyFunction)element; @@ -816,7 +841,7 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { } final String name = element != null ? getQualifiedName(element) : null; // For the following names we shouldn't go to the RHS of assignments, - // since in typing.py there are not type aliases already and assigned to + // since in typing.py they are not type aliases already and assigned to // something not so useful. if (name != null && OPAQUE_NAMES.contains(name)) { elements.add(element); diff --git a/python/testSrc/com/jetbrains/python/PyTypingTest.java b/python/testSrc/com/jetbrains/python/PyTypingTest.java index 5901737cc18e..8dfe11baee07 100644 --- a/python/testSrc/com/jetbrains/python/PyTypingTest.java +++ b/python/testSrc/com/jetbrains/python/PyTypingTest.java @@ -966,6 +966,14 @@ public class PyTypingTest extends PyTestCase { " expr = MyClass(x)\n"); } + // PY-18816 + public void testLocalTypeAlias() { + doTest("int", + "def func(g):\n" + + " Alias = int\n" + + " expr: Alias = g()"); + } + private void doTestNoInjectedText(@NotNull String text) { myFixture.configureByText(PythonFileType.INSTANCE, text); final InjectedLanguageManager languageManager = InjectedLanguageManager.getInstance(myFixture.getProject());