PY-34493 Don't copy annotations from .pyi stubs and libraries on super method completion

GitOrigin-RevId: 65787827f5df5aa80986107dda0ba555b0942d40
This commit is contained in:
Mikhail Golubev
2023-08-09 20:53:35 +00:00
committed by intellij-monorepo-bot
parent dac1ef32bb
commit b50c617b7c
14 changed files with 122 additions and 19 deletions
@@ -328,7 +328,7 @@ public abstract class PythonCommonCompletionTest extends PythonCommonTestCase {
}
public void testSuperMethodWithAnnotation() {
doTest();
runWithLanguageLevel(LanguageLevel.getLatest(), this::doTest);
}
public void testSuperMethodWithCommentAnnotation() {
@@ -340,6 +340,33 @@ public abstract class PythonCommonCompletionTest extends PythonCommonTestCase {
doTest();
}
// PY-34493
public void testSuperMethodAnnotationsNotCopiedFromPyiStub() {
doMultiFileTest();
}
// PY-34493
public void testSuperMethodAnnotationsNotCopiedFromThirdPartyLibrary() {
runWithLanguageLevel(LanguageLevel.getLatest(), () -> {
runWithAdditionalClassEntryInSdkRoots(getTestName(true) + "/lib", () -> {
myFixture.copyDirectoryToProject(getTestName(true) + "/src", "");
myFixture.configureByFile("a.py");
myFixture.completeBasic();
myFixture.checkResultByFile(getTestName(true) + "/src/a.after.py");
});
});
}
// PY-34493
public void testSuperMethodAnnotationsCopiedFromPyiStubToPyiStub() {
runWithLanguageLevel(LanguageLevel.getLatest(), () -> {
myFixture.copyDirectoryToProject(getTestName(true), "");
myFixture.configureByFile("a.pyi");
myFixture.complete(CompletionType.BASIC, 1);
myFixture.checkResultByFile(getTestName(true) + "/a.after.pyi");
});
}
public void testLocalVarInDictKey() { // PY-2558
doTest();
}
@@ -2076,7 +2103,8 @@ public abstract class PythonCommonCompletionTest extends PythonCommonTestCase {
});
}
protected String getTestDataPath() {
return PythonTestUtil.getTestDataPath() + "/completion";
@Override
protected @NotNull String getTestDataPath() {
return super.getTestDataPath() + "/completion";
}
}
@@ -6,12 +6,14 @@ import com.intellij.openapi.projectRoots.Sdk;
import com.intellij.openapi.projectRoots.SdkModificator;
import com.intellij.openapi.roots.OrderRootType;
import com.intellij.openapi.util.text.StringUtil;
import com.intellij.openapi.vfs.StandardFileSystems;
import com.intellij.openapi.vfs.VfsUtil;
import com.intellij.openapi.vfs.VirtualFile;
import com.intellij.psi.PsiFile;
import com.intellij.psi.codeStyle.CodeStyleSettings;
import com.intellij.psi.codeStyle.CommonCodeStyleSettings;
import com.jetbrains.python.PythonLanguage;
import com.jetbrains.python.PythonTestUtil;
import com.jetbrains.python.documentation.PyDocumentationSettings;
import com.jetbrains.python.documentation.docstrings.DocStringFormat;
import com.jetbrains.python.psi.LanguageLevel;
@@ -365,6 +367,13 @@ public abstract class PythonCommonTestCase extends TestCase {
runWithAdditionalRoot(sdk, directory, OrderRootType.CLASSES, (__) -> runnable.run());
}
protected void runWithAdditionalClassEntryInSdkRoots(@NotNull String relativeTestDataPath, @NotNull Runnable runnable) {
final String absPath = getTestDataPath() + "/" + relativeTestDataPath;
final VirtualFile testDataDir = StandardFileSystems.local().findFileByPath(absPath);
assertNotNull("Additional class entry directory '" + absPath + "' not found", testDataDir);
runWithAdditionalClassEntryInSdkRoots(testDataDir, runnable);
}
private static void runWithAdditionalRoot(@NotNull Sdk sdk,
@NotNull VirtualFile root,
@NotNull OrderRootType rootType,
@@ -387,4 +396,8 @@ public abstract class PythonCommonTestCase extends TestCase {
});
}
}
protected @NotNull String getTestDataPath() {
return PythonTestUtil.getTestDataPath();
}
}
@@ -26,10 +26,7 @@ import com.intellij.psi.util.PsiTreeUtil;
import com.intellij.util.ProcessingContext;
import com.jetbrains.python.PyNames;
import com.jetbrains.python.PyTokenTypes;
import com.jetbrains.python.psi.LanguageLevel;
import com.jetbrains.python.psi.PyClass;
import com.jetbrains.python.psi.PyFunction;
import com.jetbrains.python.psi.PyParameterList;
import com.jetbrains.python.psi.*;
import com.jetbrains.python.psi.types.TypeEvalContext;
import com.jetbrains.python.refactoring.PyPsiRefactoringUtil;
import org.jetbrains.annotations.NotNull;
@@ -74,8 +71,16 @@ public class PySuperMethodCompletionContributor extends CompletionContributor im
StringBuilder builder = new StringBuilder();
builder.append(superMethod.getName());
if (!(nextElement instanceof PyParameterList)) {
builder.append(superMethod.getParameterList().getText());
if (superMethod.getAnnotation() != null) {
PyParameterList parameterList;
boolean copyAnnotations = PyPsiRefactoringUtil.shouldCopyAnnotations(superMethod, parameters.getOriginalFile());
if (copyAnnotations) {
parameterList = superMethod.getParameterList();
}
else {
parameterList = stripAnnotations(superMethod.getParameterList());
}
builder.append(parameterList.getText());
if (superMethod.getAnnotation() != null && copyAnnotations) {
builder.append(" ")
.append(superMethod.getAnnotation().getText())
.append(":");
@@ -94,4 +99,15 @@ public class PySuperMethodCompletionContributor extends CompletionContributor im
}
});
}
private static <T extends PsiElement> @NotNull T stripAnnotations(@NotNull T element) {
@SuppressWarnings("unchecked") T result = (T)element.copy();
result.accept(new PyRecursiveElementVisitor() {
@Override
public void visitPyAnnotation(@NotNull PyAnnotation node) {
node.delete();
}
});
return result;
}
}
@@ -10,6 +10,7 @@ import com.intellij.openapi.util.Pair;
import com.intellij.openapi.vfs.VirtualFile;
import com.intellij.psi.*;
import com.intellij.psi.impl.light.LightElement;
import com.intellij.psi.search.ProjectScope;
import com.intellij.psi.util.PsiTreeUtil;
import com.intellij.psi.util.PsiUtilBase;
import com.intellij.psi.util.QualifiedName;
@@ -386,4 +387,13 @@ public final class PyPsiRefactoringUtil {
addSuperClassExpressions(project, clazz, superClassNames, null);
}
public static boolean shouldCopyAnnotations(@NotNull PsiElement copiedElement, @NotNull PsiFile destFile) {
if (LanguageLevel.forElement(copiedElement).isPython2() ||
(PyiUtil.isInsideStub(copiedElement) && !PyiUtil.isPyiFileOfPackage(destFile))) {
return false;
}
VirtualFile virtualFile = copiedElement.getContainingFile().getVirtualFile();
return virtualFile != null && ProjectScope.getProjectScope(copiedElement.getProject()).contains(virtualFile);
}
}
@@ -208,7 +208,8 @@ public final class PyOverrideImplementUtil {
}
}
PyAnnotation anno = baseFunction.getAnnotation();
if (anno != null && shouldCopyAnnotations(baseFunction, pyClass)) {
boolean copyAnnotations = PyPsiRefactoringUtil.shouldCopyAnnotations(baseFunction, pyClass.getContainingFile());
if (anno != null && copyAnnotations) {
pyFunctionBuilder.annotation(anno.getText());
}
if (baseFunction.isAsync()) {
@@ -224,7 +225,7 @@ public final class PyOverrideImplementUtil {
final StringBuilder parameterBuilder = new StringBuilder();
parameterBuilder.append(ParamHelper.getNameInSignature(namedParameter));
final PyAnnotation annotation = namedParameter.getAnnotation();
if (annotation != null && shouldCopyAnnotations(baseFunction, pyClass)) {
if (annotation != null && copyAnnotations) {
parameterBuilder.append(annotation.getText());
}
final PyExpression defaultValue = namedParameter.getDefaultValue();
@@ -315,14 +316,6 @@ public final class PyOverrideImplementUtil {
return pyFunctionBuilder;
}
private static boolean shouldCopyAnnotations(@NotNull PyFunction baseFunction, @NotNull PyClass subClass) {
if (LanguageLevel.forElement(baseFunction).isPython2() || (PyiUtil.isInsideStub(baseFunction) && !PyiUtil.isInsideStub(subClass))) {
return false;
}
VirtualFile virtualFile = baseFunction.getContainingFile().getVirtualFile();
return virtualFile != null && ProjectScope.getProjectScope(baseFunction.getProject()).contains(virtualFile);
}
// TODO find a better place for this logic
private static String getReferenceText(PyClass fromClass, PyClass toClass) {
final PyExpression[] superClassExpressions = fromClass.getSuperClassExpressions();
@@ -0,0 +1,5 @@
from mod import Super
class Sub(Super):
def method(self, p: int) -> str:
@@ -0,0 +1,5 @@
from mod import Super
class Sub(Super):
def meth<caret>
@@ -0,0 +1,3 @@
class Super:
def method(self, p: int) -> str:
...
@@ -0,0 +1,7 @@
class Super:
def method(self, x):
pass
class Sub(Super):
def method(self, x):<caret>
@@ -0,0 +1,7 @@
class Super:
def method(self, x):
pass
class Sub(Super):
def meth<caret>
@@ -0,0 +1,3 @@
class Super:
def method(self, x: int) -> str:
...
@@ -0,0 +1,3 @@
class Super:
def method(self, x: int) -> str:
pass
@@ -0,0 +1,5 @@
from mod import Super
class Sub(Super):
def method(self, x):<caret>
@@ -0,0 +1,5 @@
from mod import Super
class Sub(Super):
def meth<caret>