From 8894ce938aca59afd61ba5d49897e00fb400bf3b Mon Sep 17 00:00:00 2001 From: Andrey Vlasovskikh Date: Thu, 2 Apr 2015 18:19:31 +0300 Subject: [PATCH] Infer types from '# type: ...' comments as in PEP 484 (PY-15206) --- .../codeInsight/PyTypingTypeProvider.java | 54 ++++++++++++++++--- .../com/jetbrains/python/PyTypingTest.java | 6 +++ 2 files changed, 52 insertions(+), 8 deletions(-) diff --git a/python/src/com/jetbrains/python/codeInsight/PyTypingTypeProvider.java b/python/src/com/jetbrains/python/codeInsight/PyTypingTypeProvider.java index 87947bc01cb9..648dd4d426b3 100644 --- a/python/src/com/jetbrains/python/codeInsight/PyTypingTypeProvider.java +++ b/python/src/com/jetbrains/python/codeInsight/PyTypingTypeProvider.java @@ -19,8 +19,10 @@ import com.google.common.collect.ImmutableMap; import com.google.common.collect.ImmutableSet; import com.intellij.openapi.project.Project; import com.intellij.openapi.util.Ref; +import com.intellij.psi.PsiComment; import com.intellij.psi.PsiElement; import com.intellij.psi.PsiPolyVariantReference; +import com.intellij.psi.util.PsiTreeUtil; import com.intellij.psi.util.QualifiedName; import com.jetbrains.python.PyNames; import com.jetbrains.python.psi.*; @@ -33,11 +35,14 @@ import org.jetbrains.annotations.Nullable; import java.util.ArrayList; import java.util.Collections; import java.util.List; +import java.util.regex.Matcher; +import java.util.regex.Pattern; /** * @author vlan */ public class PyTypingTypeProvider extends PyTypeProviderBase { + public static final Pattern TYPE_COMMENT_PATTERN = Pattern.compile("# *type: *(.*)"); private static ImmutableMap BUILTIN_COLLECTIONS = ImmutableMap.builder() .put("typing.List", "list") .put("typing.Dict", "dict") @@ -104,6 +109,33 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { return null; } + @Override + public PyType getReferenceType(@NotNull PsiElement referenceTarget, TypeEvalContext context, @Nullable PsiElement anchor) { + if (referenceTarget instanceof PyTargetExpression && context.maySwitchToAST(referenceTarget)) { + final String comment = getTypeComment((PyTargetExpression)referenceTarget); + if (comment != null) { + return getStringBasedType(comment, referenceTarget, context); + } + } + return null; + } + + @Nullable + private static String getTypeComment(@NotNull PyTargetExpression target) { + final PyAssignmentStatement assignment = PsiTreeUtil.getParentOfType(target, PyAssignmentStatement.class); + if (assignment != null) { + final PsiElement lastChild = assignment.getLastChild(); + if (lastChild instanceof PsiComment) { + final String text = lastChild.getText(); + final Matcher m = TYPE_COMMENT_PATTERN.matcher(text); + if (m.matches()) { + return m.group(1); + } + } + } + return null; + } + private static boolean isAny(@NotNull PyType type) { return type instanceof PyClassType && "typing.Any".equals(((PyClassType)type).getPyClass().getQualifiedName()); } @@ -257,14 +289,20 @@ public class PyTypingTypeProvider extends PyTypeProviderBase { if (expression instanceof PyStringLiteralExpression) { // XXX: Requires switching from stub to AST final String contents = ((PyStringLiteralExpression)expression).getStringValue(); - final Project project = expression.getProject(); - final PyExpressionCodeFragmentImpl codeFragment = new PyExpressionCodeFragmentImpl(project, "dummy.py", contents, false); - codeFragment.setContext(expression.getContainingFile()); - final PsiElement element = codeFragment.getFirstChild(); - if (element instanceof PyExpressionStatement) { - final PyExpression dummyExpr = ((PyExpressionStatement)element).getExpression(); - return getType(dummyExpr, context); - } + return getStringBasedType(contents, expression, context); + } + return null; + } + + @Nullable + private static PyType getStringBasedType(@NotNull String contents, @NotNull PsiElement anchor, @NotNull TypeEvalContext context) { + final Project project = anchor.getProject(); + final PyExpressionCodeFragmentImpl codeFragment = new PyExpressionCodeFragmentImpl(project, "dummy.py", contents, false); + codeFragment.setContext(anchor.getContainingFile()); + final PsiElement element = codeFragment.getFirstChild(); + if (element instanceof PyExpressionStatement) { + final PyExpression dummyExpr = ((PyExpressionStatement)element).getExpression(); + return getType(dummyExpr, context); } return null; } diff --git a/python/testSrc/com/jetbrains/python/PyTypingTest.java b/python/testSrc/com/jetbrains/python/PyTypingTest.java index 48cee02c68d9..a941e1dc68db 100644 --- a/python/testSrc/com/jetbrains/python/PyTypingTest.java +++ b/python/testSrc/com/jetbrains/python/PyTypingTest.java @@ -279,6 +279,12 @@ public class PyTypingTest extends PyTestCase { " expr = cast(str, x)\n"); } + public void testComment() { + doTest("int", + "def foo(x):\n" + + " expr = x # type: int\n"); + } + private void doTest(@NotNull String expectedType, @NotNull String text) { myFixture.copyDirectoryToProject("typing", ""); myFixture.configureByText(PythonFileType.INSTANCE, text);