From 724f0a385a045c61fc839d1f8a301f3b256aa157 Mon Sep 17 00:00:00 2001 From: abragin Date: Tue, 12 Dec 2017 12:07:34 +0300 Subject: [PATCH] PY-18553 Hinted types import on override added When base function is located in different module its type hints may refer types that are inaccessible in the module where overridden function is located. To prevent unresolved reference error import statements are added to the module. --- .../override/PyOverrideImplementUtil.java | 112 +++++++++++++++++- .../classes/PyClassRefactoringUtil.java | 28 +++-- .../override/importsForTypeAnnotations1.py | 5 + .../importsForTypeAnnotations1_after.py | 6 + .../importsForTypeAnnotations1_import.py | 3 + .../override/importsForTypeAnnotations2.py | 5 + .../importsForTypeAnnotations2_after.py | 8 ++ .../importsForTypeAnnotations2_import.py | 6 + .../override/importsForTypeAnnotations3.py | 5 + .../importsForTypeAnnotations3_after.py | 6 + .../importsForTypeAnnotations3_import.py | 11 ++ .../com/jetbrains/python/PyOverrideTest.java | 35 ++++++ 12 files changed, 213 insertions(+), 17 deletions(-) create mode 100644 python/testData/override/importsForTypeAnnotations1.py create mode 100644 python/testData/override/importsForTypeAnnotations1_after.py create mode 100644 python/testData/override/importsForTypeAnnotations1_import.py create mode 100644 python/testData/override/importsForTypeAnnotations2.py create mode 100644 python/testData/override/importsForTypeAnnotations2_after.py create mode 100644 python/testData/override/importsForTypeAnnotations2_import.py create mode 100644 python/testData/override/importsForTypeAnnotations3.py create mode 100644 python/testData/override/importsForTypeAnnotations3_after.py create mode 100644 python/testData/override/importsForTypeAnnotations3_import.py diff --git a/python/src/com/jetbrains/python/codeInsight/override/PyOverrideImplementUtil.java b/python/src/com/jetbrains/python/codeInsight/override/PyOverrideImplementUtil.java index 713b5826f7b6..0fac595c0334 100644 --- a/python/src/com/jetbrains/python/codeInsight/override/PyOverrideImplementUtil.java +++ b/python/src/com/jetbrains/python/codeInsight/override/PyOverrideImplementUtil.java @@ -3,6 +3,7 @@ package com.jetbrains.python.codeInsight.override; import com.google.common.collect.Lists; import com.google.common.collect.Maps; +import com.google.common.collect.Streams; import com.intellij.codeInsight.CodeInsightUtilCore; import com.intellij.featureStatistics.FeatureUsageTracker; import com.intellij.featureStatistics.ProductivityFeatureNames; @@ -26,14 +27,17 @@ import com.jetbrains.python.PyNames; import com.jetbrains.python.psi.*; import com.jetbrains.python.psi.impl.PyFunctionBuilder; import com.jetbrains.python.psi.impl.PyPsiUtils; +import com.jetbrains.python.psi.resolve.PyResolveContext; import com.jetbrains.python.psi.types.*; +import com.jetbrains.python.refactoring.classes.PyClassRefactoringUtil; +import one.util.streamex.StreamEx; import org.jetbrains.annotations.NotNull; import org.jetbrains.annotations.Nullable; -import java.util.ArrayList; -import java.util.Collection; -import java.util.List; -import java.util.Map; +import java.util.*; +import java.util.function.Function; +import java.util.stream.Collectors; +import java.util.stream.Stream; /** * @author Alexey.Ivanov @@ -164,6 +168,8 @@ public class PyOverrideImplementUtil { final PyFunction baseFunction = (PyFunction)newMember.getPsiElement(); final PyFunctionBuilder builder = buildOverriddenFunction(pyClass, baseFunction, implement); final PyFunction function = builder.addFunctionAfter(statementList, anchor, languageLevel); + + addImports(baseFunction, function); element = CodeInsightUtilCore.forcePsiPostprocessAndRestoreElement(function); } @@ -358,4 +364,102 @@ public class PyOverrideImplementUtil { } return Lists.newArrayList(functions.values()); } + + /** + * Adds imports for type hints in overridden function (PY-18553). + * + * @param baseFunction base function used to resolve types + * @param function overridden function + */ + private static void addImports(@NotNull PyFunction baseFunction, @NotNull PyFunction function) { + final TypeEvalContext typeEvalContext = TypeEvalContext.userInitiated(baseFunction.getProject(), baseFunction.getContainingFile()); + + final UnresolvedExpressionVisitor unresolvedExpressionVisitor = new UnresolvedExpressionVisitor(); + final List annotations = getAnnotations(function, typeEvalContext); + annotations.forEach(annotation -> unresolvedExpressionVisitor.visitPyElement(annotation)); + final List unresolved = unresolvedExpressionVisitor.getUnresolved(); + + final ResolveExpressionVisitor resolveExpressionVisitor = new ResolveExpressionVisitor(unresolved); + final List baseAnnotations = getAnnotations(baseFunction, typeEvalContext); + baseAnnotations.forEach(annotation -> resolveExpressionVisitor.visitPyElement(annotation)); + } + + /** + * Collect annotations from function parameters and return. + * + * @param function + * @param typeEvalContext + * @return + */ + private static List getAnnotations(@NotNull PyFunction function, @NotNull TypeEvalContext typeEvalContext) { + return Streams.concat( + function.getParameters(typeEvalContext).stream() + .map(PyCallableParameter::getParameter) + .filter(PyNamedParameter.class::isInstance) + .map(PyNamedParameter.class::cast) + .filter(parameter -> !parameter.isSelf()) + .map(pyNamedParameter -> pyNamedParameter.getAnnotation()), + Stream.of(function.getAnnotation()) + ) + .filter(Objects::nonNull) + .collect(Collectors.toList()); + } + + /** + * Collects unresolved {@link PyReferenceExpression} objects. + */ + private static class UnresolvedExpressionVisitor extends PyRecursiveElementVisitor { + + private final List myUnresolved = new ArrayList<>(); + + @Override + public void visitPyReferenceExpression(final PyReferenceExpression referenceExpression) { + super.visitPyReferenceExpression(referenceExpression); + final PyResolveContext resolveContext = PyResolveContext.noImplicits(); + if (referenceExpression.getReference(resolveContext).multiResolve(false).length == 0) { + myUnresolved.add(referenceExpression); + } + } + + /** + * Get list of {@link PyReferenceExpression} that left myUnresolved after function override. + * + * @return list of {@link PyReferenceExpression} elements. + */ + @NotNull + List getUnresolved() { + return Collections.unmodifiableList(myUnresolved); + } + } + + /** + * Resolves reference expressions by name and adds imports for them using references being visited. + */ + private static class ResolveExpressionVisitor extends PyRecursiveElementVisitor { + + private final Map myExpressionsToResolve; + + /** + * {@link PyReferenceExpression} objects to resolve. + * + * @param toResolve collection of references to resolve. + */ + ResolveExpressionVisitor(@NotNull Collection toResolve) { + myExpressionsToResolve = StreamEx.of(toResolve) + .toMap(PyReferenceExpression::getName, Function.identity(), + (expression1, expression2) -> expression2); + } + + @Override + public void visitPyReferenceExpression(final PyReferenceExpression referenceExpression) { + super.visitPyReferenceExpression(referenceExpression); + + if (myExpressionsToResolve.containsKey(referenceExpression.getName())) { + PyClassRefactoringUtil.rememberNamedReferences(referenceExpression); + PyClassRefactoringUtil.restoreReference(referenceExpression, + myExpressionsToResolve.get(referenceExpression.getName()), + PsiElement.EMPTY_ARRAY); + } + } + } } diff --git a/python/src/com/jetbrains/python/refactoring/classes/PyClassRefactoringUtil.java b/python/src/com/jetbrains/python/refactoring/classes/PyClassRefactoringUtil.java index c838da234fba..5c31e0595b0a 100644 --- a/python/src/com/jetbrains/python/refactoring/classes/PyClassRefactoringUtil.java +++ b/python/src/com/jetbrains/python/refactoring/classes/PyClassRefactoringUtil.java @@ -224,7 +224,7 @@ public final class PyClassRefactoringUtil { @Override public void visitPyReferenceExpression(PyReferenceExpression node) { super.visitPyReferenceExpression(node); - restoreReference(node, otherMovedElements); + restoreReference(node, node, otherMovedElements); } @Override @@ -240,13 +240,15 @@ public final class PyClassRefactoringUtil { } - private static void restoreReference(@NotNull PyReferenceExpression node, @NotNull PsiElement[] otherMovedElements) { + public static void restoreReference(@NotNull PyReferenceExpression sourceNode, + @NotNull PyReferenceExpression targetNode, + @NotNull PsiElement[] otherMovedElements) { try { - PsiNamedElement target = node.getCopyableUserData(ENCODED_IMPORT); - final String asName = node.getCopyableUserData(ENCODED_IMPORT_AS); - final Boolean useFromImport = node.getCopyableUserData(ENCODED_USE_FROM_IMPORT); + PsiNamedElement target = sourceNode.getCopyableUserData(ENCODED_IMPORT); + final String asName = sourceNode.getCopyableUserData(ENCODED_IMPORT_AS); + final Boolean useFromImport = sourceNode.getCopyableUserData(ENCODED_USE_FROM_IMPORT); if (target instanceof PsiDirectory) { - target = (PsiNamedElement)PyUtil.getPackageElement((PsiDirectory)target, node); + target = (PsiNamedElement)PyUtil.getPackageElement((PsiDirectory)target, sourceNode); } if (target instanceof PyFunction) { final PyFunction f = (PyFunction)target; @@ -256,19 +258,19 @@ public final class PyClassRefactoringUtil { } } if (target == null) return; - if (PsiTreeUtil.isAncestor(node.getContainingFile(), target, false)) return; + if (PsiTreeUtil.isAncestor(targetNode.getContainingFile(), target, false)) return; if (ArrayUtil.contains(target, otherMovedElements)) return; if (target instanceof PyFile || target instanceof PsiDirectory) { - insertImport(node, target, asName, useFromImport != null ? useFromImport : true); + insertImport(targetNode, target, asName, useFromImport != null ? useFromImport : true); } else { - insertImport(node, target, asName, true); + insertImport(targetNode, target, asName, true); } } finally { - node.putCopyableUserData(ENCODED_IMPORT, null); - node.putCopyableUserData(ENCODED_IMPORT_AS, null); - node.putCopyableUserData(ENCODED_USE_FROM_IMPORT, null); + sourceNode.putCopyableUserData(ENCODED_IMPORT, null); + sourceNode.putCopyableUserData(ENCODED_IMPORT_AS, null); + sourceNode.putCopyableUserData(ENCODED_USE_FROM_IMPORT, null); } } @@ -344,7 +346,7 @@ public final class PyClassRefactoringUtil { * @param namesToSkip if reference inside of element has one of this names, it will not be saved. */ public static void rememberNamedReferences(@NotNull final PsiElement element, @NotNull final String... namesToSkip) { - element.acceptChildren(new PyRecursiveElementVisitor() { + element.accept(new PyRecursiveElementVisitor() { @Override public void visitPyReferenceExpression(PyReferenceExpression node) { super.visitPyReferenceExpression(node); diff --git a/python/testData/override/importsForTypeAnnotations1.py b/python/testData/override/importsForTypeAnnotations1.py new file mode 100644 index 000000000000..5221f01f2d57 --- /dev/null +++ b/python/testData/override/importsForTypeAnnotations1.py @@ -0,0 +1,5 @@ +from .importsForTypeAnnotations1_import import Foo + + +class Bar(Foo): + pass diff --git a/python/testData/override/importsForTypeAnnotations1_after.py b/python/testData/override/importsForTypeAnnotations1_after.py new file mode 100644 index 000000000000..65ff1b958d5e --- /dev/null +++ b/python/testData/override/importsForTypeAnnotations1_after.py @@ -0,0 +1,6 @@ +from .importsForTypeAnnotations1_import import Foo + + +class Bar(Foo): + def func(self, arg: int) -> int: + return super().func(arg) diff --git a/python/testData/override/importsForTypeAnnotations1_import.py b/python/testData/override/importsForTypeAnnotations1_import.py new file mode 100644 index 000000000000..406e2f9929c0 --- /dev/null +++ b/python/testData/override/importsForTypeAnnotations1_import.py @@ -0,0 +1,3 @@ +class Foo: + def func(self, arg: int) -> int: + pass diff --git a/python/testData/override/importsForTypeAnnotations2.py b/python/testData/override/importsForTypeAnnotations2.py new file mode 100644 index 000000000000..7f2c03e41b1d --- /dev/null +++ b/python/testData/override/importsForTypeAnnotations2.py @@ -0,0 +1,5 @@ +from .importsForTypeAnnotations2_import import Foo + + +class Bar(Foo): + pass diff --git a/python/testData/override/importsForTypeAnnotations2_after.py b/python/testData/override/importsForTypeAnnotations2_after.py new file mode 100644 index 000000000000..c8dfc86ad1bf --- /dev/null +++ b/python/testData/override/importsForTypeAnnotations2_after.py @@ -0,0 +1,8 @@ +from typing import Union + +from .importsForTypeAnnotations2_import import Foo + + +class Bar(Foo): + def something(self, arg: Union[dict, int]) -> Union[None, int]: + return super().something(arg) diff --git a/python/testData/override/importsForTypeAnnotations2_import.py b/python/testData/override/importsForTypeAnnotations2_import.py new file mode 100644 index 000000000000..f62e636785c7 --- /dev/null +++ b/python/testData/override/importsForTypeAnnotations2_import.py @@ -0,0 +1,6 @@ +from typing import Union + + +class Foo: + def something(self, arg: Union[dict, int]) -> Union[None, int]: + pass diff --git a/python/testData/override/importsForTypeAnnotations3.py b/python/testData/override/importsForTypeAnnotations3.py new file mode 100644 index 000000000000..96608a91b1fb --- /dev/null +++ b/python/testData/override/importsForTypeAnnotations3.py @@ -0,0 +1,5 @@ +from override.importsForTypeAnnotations3_import import Foo + + +class Bar(Foo): + pass diff --git a/python/testData/override/importsForTypeAnnotations3_after.py b/python/testData/override/importsForTypeAnnotations3_after.py new file mode 100644 index 000000000000..39b9cd95a357 --- /dev/null +++ b/python/testData/override/importsForTypeAnnotations3_after.py @@ -0,0 +1,6 @@ +from override.importsForTypeAnnotations3_import import Foo, Param, Return + + +class Bar(Foo): + def func(self, arg: Param) -> Return: + return super().func(arg) diff --git a/python/testData/override/importsForTypeAnnotations3_import.py b/python/testData/override/importsForTypeAnnotations3_import.py new file mode 100644 index 000000000000..e42650751847 --- /dev/null +++ b/python/testData/override/importsForTypeAnnotations3_import.py @@ -0,0 +1,11 @@ +class Param: + pass + + +class Return: + pass + + +class Foo: + def func(self, arg: Param) -> Return: + pass diff --git a/python/testSrc/com/jetbrains/python/PyOverrideTest.java b/python/testSrc/com/jetbrains/python/PyOverrideTest.java index 5f8ab391d60d..ba0c4b3d93a2 100644 --- a/python/testSrc/com/jetbrains/python/PyOverrideTest.java +++ b/python/testSrc/com/jetbrains/python/PyOverrideTest.java @@ -17,6 +17,8 @@ import com.jetbrains.python.psi.stubs.PyClassNameIndex; import java.util.Arrays; import java.util.Collections; +import java.util.List; +import java.util.stream.Collectors; /** * @author yole @@ -163,6 +165,39 @@ public class PyOverrideTest extends PyTestCase { doTest3k(); } + // PY-18553 + public void testImportsForTypeAnnotations1() { + testImportsForTypeAnnotations(getTestName(true), 0); + } + + public void testImportsForTypeAnnotations2() { + testImportsForTypeAnnotations(getTestName(true), 0); + } + + public void testImportsForTypeAnnotations3() { + testImportsForTypeAnnotations(getTestName(true), 2); + } + + private void testImportsForTypeAnnotations(String testName, int orderOfClassToOverride) { + + runWithLanguageLevel(LanguageLevel.PYTHON35, () -> { + final String initialFilePath = String.format("override/%s.py", testName); + final String importFilePath = String.format("override/%s_import.py", testName); + final String resultFilePath = String.format("override/%s_after.py", testName); + + List pyFiles = Arrays.stream( + myFixture.configureByFiles(initialFilePath, importFilePath)) + .map(PyFile.class::cast) + .collect(Collectors.toList()); + + PyFunction toOverride = pyFiles.get(1).getTopLevelClasses().get(orderOfClassToOverride).getMethods()[0]; + PyOverrideImplementUtil.overrideMethods(myFixture.getEditor(), getTopLevelClass(0), + Collections.singletonList(new PyMethodMember(toOverride)), false); + myFixture.checkResultByFile(resultFilePath, true); + }); + + } + public void testSingleStar() { // PY-6455 doTest3k(); }