mirror of
https://gitflic.ru/project/openide/openide.git
synced 2026-09-27 10:03:11 +07:00
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:
committed by
intellij-monorepo-bot
parent
c7bba04743
commit
a2af264b63
@@ -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();
|
||||
}
|
||||
|
||||
+10
-1
@@ -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));
|
||||
}
|
||||
}
|
||||
|
||||
+42
-3
@@ -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
|
||||
Reference in New Issue
Block a user