PY-18816 Use actual AST instead for annotations if context allows it

This way we resolve type hints more precisely inside e.g. the file
actually opened in the editor, in particular, those referencing local
classes and type aliases.
This commit is contained in:
Mikhail Golubev
2017-07-19 19:28:31 +03:00
parent d2011f7147
commit 09030f96ad
4 changed files with 58 additions and 23 deletions
@@ -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();
@@ -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
@@ -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<PyType> getParameterTypeFromAnnotation(@NotNull PyNamedParameter parameter, @NotNull TypeEvalContext context) {
final Ref<PyType> 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<PyType> 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<PyType> 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<PyType> 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<PyType> 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<PyType> 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<PsiElement> tryResolving(@NotNull PyExpression expression, @NotNull TypeEvalContext context) {
final List<PsiElement> elements = Lists.newArrayList();
if (expression instanceof PyReferenceExpression) {
final List<PsiElement> results = tryResolvingOnStubs((PyReferenceExpression)expression, context);
final List<PsiElement> 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);
@@ -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());