PY-45588 Add necessary imports on completing methods of superclasses to override

I had to move addImports to PyClassRefactoringUtil because PySuperMethodCompletionContributor
resides in python-psi-impl and, thus, has no access to PyOverrideImplementUtil.

GitOrigin-RevId: cf2ac19da779977649144b2477bac3f8ae78bbcd
This commit is contained in:
Mikhail Golubev
2023-08-09 20:53:35 +00:00
committed by intellij-monorepo-bot
parent c7bba04743
commit a2af264b63
7 changed files with 80 additions and 93 deletions
@@ -331,6 +331,11 @@ public abstract class PythonCommonCompletionTest extends PythonCommonTestCase {
runWithLanguageLevel(LanguageLevel.getLatest(), this::doTest);
}
// PY-45588
public void testSuperMethodWithAnnotationInsertingImports() {
runWithLanguageLevel(LanguageLevel.getLatest(), this::doMultiFileTest);
}
public void testSuperMethodWithCommentAnnotation() {
doTest();
}
@@ -19,6 +19,7 @@ import com.intellij.codeInsight.TailType;
import com.intellij.codeInsight.completion.*;
import com.intellij.codeInsight.lookup.LookupElementBuilder;
import com.intellij.codeInsight.lookup.TailTypeDecorator;
import com.intellij.openapi.command.WriteCommandAction;
import com.intellij.openapi.project.DumbAware;
import com.intellij.psi.PsiElement;
import com.intellij.psi.PsiWhiteSpace;
@@ -29,6 +30,7 @@ import com.jetbrains.python.PyTokenTypes;
import com.jetbrains.python.psi.*;
import com.jetbrains.python.psi.types.TypeEvalContext;
import com.jetbrains.python.refactoring.PyPsiRefactoringUtil;
import com.jetbrains.python.refactoring.classes.PyClassRefactoringUtil;
import org.jetbrains.annotations.NotNull;
import java.util.HashSet;
@@ -93,7 +95,14 @@ public class PySuperMethodCompletionContributor extends CompletionContributor im
builder.append(":");
}
}
LookupElementBuilder element = LookupElementBuilder.create(builder.toString());
LookupElementBuilder element = LookupElementBuilder.create(builder.toString())
.withInsertHandler((insertionContext, item) -> {
PsiElement methodName = insertionContext.getFile().findElementAt(insertionContext.getStartOffset());
if (methodName == null || !(methodName.getParent() instanceof PyFunction insertedMethod)) return;
WriteCommandAction.writeCommandAction(insertionContext.getFile()).run(() -> {
PyClassRefactoringUtil.transplantImportsFromSignature(superMethod, insertedMethod);
});
});
result.addElement(TailTypeDecorator.withTail(element, TailType.NONE));
}
}
@@ -236,9 +236,9 @@ public final class PyClassRefactoringUtil {
}
public static void restoreReference(@NotNull PsiElement sourceNode,
@NotNull PsiElement targetNode,
PsiElement @NotNull [] otherMovedElements) {
private static void restoreReference(@NotNull PsiElement sourceNode,
@NotNull PsiElement targetNode,
PsiElement @NotNull [] otherMovedElements) {
try {
if (sourceNode instanceof PyReferenceExpression) {
doRestoreReference((PyReferenceExpression)sourceNode, targetNode, otherMovedElements);
@@ -555,6 +555,45 @@ public final class PyClassRefactoringUtil {
return (PyFile)psi;
}
/**
* Transfer necessary imports for references in the signature of one function to another.
*/
public static void transplantImportsFromSignature(@NotNull PyFunction sourceFunction, @NotNull PyFunction targetFunction) {
Map<String, PyReferenceExpression> unresolvedNames = new HashMap<>();
targetFunction.accept(new PyRecursiveElementVisitor() {
@Override
public void visitPyReferenceExpression(@NotNull PyReferenceExpression referenceExpression) {
super.visitPyReferenceExpression(referenceExpression);
if (!referenceExpression.isQualified() && referenceExpression.getReference().multiResolve(false).length == 0) {
unresolvedNames.put(referenceExpression.getName(), referenceExpression);
}
}
@Override
public void visitPyStatementList(@NotNull PyStatementList pruned) {
}
});
sourceFunction.accept(new PyRecursiveElementVisitor() {
@Override
public void visitPyReferenceExpression(@NotNull PyReferenceExpression referenceExpression) {
super.visitPyReferenceExpression(referenceExpression);
if (!referenceExpression.isQualified()) {
PyReferenceExpression unresolvedReference = unresolvedNames.get(referenceExpression.getName());
if (unresolvedReference != null) {
rememberNamedReferences(referenceExpression);
restoreReference(referenceExpression, unresolvedReference, PsiElement.EMPTY_ARRAY);
}
}
}
@Override
public void visitPyStatementList(@NotNull PyStatementList pruned) {
}
});
}
private static final class DynamicNamedElement extends LightElement implements PsiNamedElement {
private final PsiFile myFile;
private final String myName;
@@ -13,11 +13,9 @@ import com.intellij.openapi.editor.ScrollType;
import com.intellij.openapi.project.Project;
import com.intellij.openapi.ui.DialogWrapper;
import com.intellij.openapi.util.text.StringUtil;
import com.intellij.openapi.vfs.VirtualFile;
import com.intellij.psi.PsiElement;
import com.intellij.psi.PsiFile;
import com.intellij.psi.PsiWhiteSpace;
import com.intellij.psi.search.ProjectScope;
import com.intellij.psi.util.PsiTreeUtil;
import com.jetbrains.python.PyBundle;
import com.jetbrains.python.PyNames;
@@ -25,19 +23,18 @@ import com.jetbrains.python.psi.*;
import com.jetbrains.python.psi.impl.ParamHelper;
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.PyCallableParameter;
import com.jetbrains.python.psi.types.PyNoneType;
import com.jetbrains.python.psi.types.TypeEvalContext;
import com.jetbrains.python.pyi.PyiUtil;
import com.jetbrains.python.refactoring.PyPsiRefactoringUtil;
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.*;
import java.util.function.Function;
import java.util.ArrayList;
import java.util.Collection;
import java.util.List;
public final class PyOverrideImplementUtil {
@@ -169,14 +166,14 @@ public final class PyOverrideImplementUtil {
final TypeEvalContext context = TypeEvalContext.userInitiated(cls.getProject(), cls.getContainingFile());
final PyFunction function = buildOverriddenFunction(cls, baseFunction, implement).addFunctionAfter(statementList, anchor);
addImports(baseFunction, function);
PyClassRefactoringUtil.transplantImportsFromSignature(baseFunction, function);
PyiUtil
.getOverloads(baseFunction, context)
.forEach(
baseOverload -> {
final PyFunction overload = (PyFunction)statementList.addBefore(baseOverload, function);
addImports(baseOverload, overload);
PyClassRefactoringUtil.transplantImportsFromSignature(baseOverload, overload);
}
);
@@ -329,85 +326,4 @@ public final class PyOverrideImplementUtil {
}
return toClass.getName();
}
/**
* Adds imports for type hints and decorators in overridden function.
*
* @param baseFunction base function used to resolve types
* @param function overridden function
*/
private static void addImports(@NotNull PyFunction baseFunction, @NotNull PyFunction function) {
UnresolvedExpressionVisitor unresolvedExpressionVisitor = new UnresolvedExpressionVisitor() {
@Override
public void visitPyStatementList(@NotNull PyStatementList node) {
}
};
function.accept(unresolvedExpressionVisitor);
ResolveExpressionVisitor resolveExpressionVisitor = new ResolveExpressionVisitor(unresolvedExpressionVisitor.getUnresolved()) {
@Override
public void visitPyStatementList(@NotNull PyStatementList node) {
}
};
baseFunction.accept(resolveExpressionVisitor);
}
/**
* Collects unresolved {@link PyReferenceExpression} objects.
*/
private static class UnresolvedExpressionVisitor extends PyRecursiveElementVisitor {
private final List<PyReferenceExpression> myUnresolved = new ArrayList<>();
@Override
public void visitPyReferenceExpression(final @NotNull PyReferenceExpression referenceExpression) {
super.visitPyReferenceExpression(referenceExpression);
final var context = TypeEvalContext.codeInsightFallback(referenceExpression.getProject());
final PyResolveContext resolveContext = PyResolveContext.defaultContext(context);
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<PyReferenceExpression> 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<String, PyReferenceExpression> myExpressionsToResolve;
/**
* {@link PyReferenceExpression} objects to resolve.
*
* @param toResolve collection of references to resolve.
*/
ResolveExpressionVisitor(@NotNull Collection<PyReferenceExpression> toResolve) {
myExpressionsToResolve = StreamEx.of(toResolve)
.toMap(PyReferenceExpression::getName, Function.identity(),
(expression1, expression2) -> expression2);
}
@Override
public void visitPyReferenceExpression(final @NotNull PyReferenceExpression referenceExpression) {
super.visitPyReferenceExpression(referenceExpression);
if (myExpressionsToResolve.containsKey(referenceExpression.getName())) {
PyClassRefactoringUtil.rememberNamedReferences(referenceExpression);
PyClassRefactoringUtil.restoreReference(referenceExpression,
myExpressionsToResolve.get(referenceExpression.getName()),
PsiElement.EMPTY_ARRAY);
}
}
}
}
@@ -0,0 +1,7 @@
from typing import Optional
from mod import Super
class Sub(Super):
def method(self) -> Optional[int]:<caret>
@@ -0,0 +1,5 @@
from mod import Super
class Sub(Super):
def metho<caret>
@@ -0,0 +1,6 @@
from typing import Optional
class Super:
def method(self) -> Optional[int]:
pass