From f98b4ed4da25525fc8610734ac7986869f4fe303 Mon Sep 17 00:00:00 2001 From: "Ilya.Kazakevich" Date: Wed, 29 Jan 2014 19:46:58 +0400 Subject: [PATCH] PY-11958 PullUp refactoring breaks class signature if class has more than 2 ancestors --- .../classes/PyClassRefactoringUtil.java | 59 ++++++++++++------- .../pullup/severalParents.after.py | 13 ++++ .../refactoring/pullup/severalParents.py | 13 ++++ .../classes/PyClassRefactoringTest.java | 24 +++++--- .../classes/PyExtractSuperclassTest.java | 3 - .../refactoring/classes/PyPullUpTest.java | 4 ++ .../refactoring/classes/PyPushDownTest.java | 1 - 7 files changed, 85 insertions(+), 32 deletions(-) create mode 100644 python/testData/refactoring/pullup/severalParents.after.py create mode 100644 python/testData/refactoring/pullup/severalParents.py diff --git a/python/src/com/jetbrains/python/refactoring/classes/PyClassRefactoringUtil.java b/python/src/com/jetbrains/python/refactoring/classes/PyClassRefactoringUtil.java index 9b11d159ae1f..eda33558779f 100644 --- a/python/src/com/jetbrains/python/refactoring/classes/PyClassRefactoringUtil.java +++ b/python/src/com/jetbrains/python/refactoring/classes/PyClassRefactoringUtil.java @@ -49,7 +49,8 @@ public class PyClassRefactoringUtil { private static final Key ENCODED_USE_FROM_IMPORT = Key.create("PyEncodedUseFromImport"); private static final Key ENCODED_IMPORT_AS = Key.create("PyEncodedImportAs"); - private PyClassRefactoringUtil() {} + private PyClassRefactoringUtil() { + } public static void moveSuperclasses(PyClass clazz, Set superClasses, PyClass superClass) { if (superClasses.size() == 0) return; @@ -58,7 +59,8 @@ public class PyClassRefactoringUtil { addSuperclasses(project, superClass, toAdd, superClasses); } - public static void addSuperclasses(Project project, PyClass superClass, + public static void addSuperclasses(Project project, + PyClass superClass, @Nullable Collection superClassesAsPsi, Collection superClassesAsStrings) { if (superClassesAsStrings.size() == 0) return; @@ -74,22 +76,33 @@ public class PyClassRefactoringUtil { argList.addArgument(PyElementGenerator.getInstance(project).createExpressionFromText(s)); } } - } else { + } + else { addSuperclasses(project, superClass, superClassesAsStrings); } } - public static List removeAndGetSuperClasses(PyClass clazz, Set superClasses) { - if (superClasses.size() == 0) return Collections.emptyList(); - final List toAdd = new ArrayList(); - final PyExpression[] elements = clazz.getSuperClassExpressions(); - for (PyExpression element : elements) { - if (superClasses.contains(element.getText())) { - toAdd.add(element); - PyUtil.removeListNode(element); + /** + * Removes super classes by name and returns list of removed + * + * @param clazz class to find super classes to remove + * @param superClassesToRemove list of super class names + * @return list of removed classes + */ + @NotNull + public static List removeAndGetSuperClasses(@NotNull PyClass clazz, @NotNull Set superClassesToRemove) { + if (superClassesToRemove.isEmpty()) { + return Collections.emptyList(); + } + final List result = new ArrayList(); + for (PyExpression superClassExpression : clazz.getSuperClassExpressions()) { + //TODO: We probably should use #getName() here, but #getText() was used in previous version so keeped temporary for backward comp. + if (superClassesToRemove.contains(superClassExpression.getText())) { + result.add(superClassExpression); + superClassExpression.delete(); } } - return toAdd; + return result; } public static void addSuperclasses(Project project, PyClass superClass, Collection superClasses) { @@ -106,7 +119,8 @@ public class PyClassRefactoringUtil { builder.append(")"); if (!hasChanges) return; - final PsiFile file = PsiFileFactory.getInstance(project).createFileFromText(superClass.getName() + "temp", PythonFileType.INSTANCE, builder.toString()); + final PsiFile file = + PsiFileFactory.getInstance(project).createFileFromText(superClass.getName() + "temp", PythonFileType.INSTANCE, builder.toString()); final PsiElement expression = file.getFirstChild().getFirstChild(); PsiElement colon = superClass.getFirstChild(); while (colon != null && !colon.getText().equals(":")) { @@ -148,7 +162,8 @@ public class PyClassRefactoringUtil { public static void insertPassIfNeeded(PyClass clazz) { final PyStatementList statements = clazz.getStatementList(); if (statements.getStatements().length == 0) { - statements.add(PyElementGenerator.getInstance(clazz.getProject()).createFromText(LanguageLevel.getDefault(), PyPassStatement.class, PyNames.PASS)); + statements.add( + PyElementGenerator.getInstance(clazz.getProject()).createFromText(LanguageLevel.getDefault(), PyPassStatement.class, PyNames.PASS)); } } @@ -238,7 +253,7 @@ public class PyClassRefactoringUtil { if (components.isEmpty()) { return false; } - for (String s: components) { + for (String s : components) { if (!PyNames.isIdentifier(s) || PyNames.isReserved(s)) { return false; } @@ -346,7 +361,7 @@ public class PyClassRefactoringUtil { final String name = getOriginalName(element); if (name != null) { PyImportElement importElement = null; - for (PyImportElement e: importStatement.getImportElements()) { + for (PyImportElement e : importStatement.getImportElements()) { if (name.equals(getOriginalName(e))) { importElement = e; } @@ -363,11 +378,14 @@ public class PyClassRefactoringUtil { } if (deleteImportElement) { if (importStatement.getImportElements().length == 1) { - final boolean isInjected = InjectedLanguageManager.getInstance(importElement.getProject()).isInjectedFragment(importElement.getContainingFile()); - if (!isInjected) + final boolean isInjected = + InjectedLanguageManager.getInstance(importElement.getProject()).isInjectedFragment(importElement.getContainingFile()); + if (!isInjected) { importStatement.delete(); - else + } + else { deleteImportStatementFromInjected(importStatement); + } } else { importElement.delete(); @@ -380,8 +398,7 @@ public class PyClassRefactoringUtil { private static void deleteImportStatementFromInjected(@NotNull final PyImportStatementBase importStatement) { final PsiElement sibling = importStatement.getPrevSibling(); importStatement.delete(); - if (sibling instanceof PsiWhiteSpace) - sibling.delete(); + if (sibling instanceof PsiWhiteSpace) sibling.delete(); } @Nullable diff --git a/python/testData/refactoring/pullup/severalParents.after.py b/python/testData/refactoring/pullup/severalParents.after.py new file mode 100644 index 000000000000..cd45e22d8479 --- /dev/null +++ b/python/testData/refactoring/pullup/severalParents.after.py @@ -0,0 +1,13 @@ +class Spam: + pass + +class Parent_1(object, Spam): + pass + + +class Parent_2(): + pass + + +class Child(Parent_1, Parent_2): + pass \ No newline at end of file diff --git a/python/testData/refactoring/pullup/severalParents.py b/python/testData/refactoring/pullup/severalParents.py new file mode 100644 index 000000000000..2c5f7a1b9ede --- /dev/null +++ b/python/testData/refactoring/pullup/severalParents.py @@ -0,0 +1,13 @@ +class Spam: + pass + +class Parent_1(object): + pass + + +class Parent_2(): + pass + + +class Child(Parent_1, Parent_2, Spam): + pass \ No newline at end of file diff --git a/python/testSrc/com/jetbrains/python/refactoring/classes/PyClassRefactoringTest.java b/python/testSrc/com/jetbrains/python/refactoring/classes/PyClassRefactoringTest.java index dbdca0d622a9..d6ebede4ac3d 100644 --- a/python/testSrc/com/jetbrains/python/refactoring/classes/PyClassRefactoringTest.java +++ b/python/testSrc/com/jetbrains/python/refactoring/classes/PyClassRefactoringTest.java @@ -21,6 +21,9 @@ import com.jetbrains.python.psi.PyClass; import com.jetbrains.python.psi.PyElement; import com.jetbrains.python.psi.PyFunction; import com.jetbrains.python.psi.stubs.PyClassNameIndex; +import org.hamcrest.Matchers; +import org.jetbrains.annotations.NotNull; +import org.junit.Assert; import java.util.Collection; @@ -28,22 +31,29 @@ import java.util.Collection; * @author Dennis.Ushakov */ public abstract class PyClassRefactoringTest extends PyTestCase { - protected PyElement findMember(String className, String memberName) { - if (!memberName.contains(".")) return findClass(memberName); - return findMethod(className, memberName.substring(1)); + /** + * @param className class where member should be found + * @param memberName member that starts with dot (.) is treated as method. + * It is treated parent class otherwise + * @return member or null if not found + */ + @NotNull + protected PyElement findMember(@NotNull String className, @NotNull String memberName) { + boolean findMethod = memberName.contains("."); + PyElement result = (findMethod ? findMethod(className, memberName.substring(1)) : findClass(memberName)); + Assert.assertNotNull(String.format("No member %s found in class %s", memberName, className), result); + return result; } private PyFunction findMethod(final String className, final String name) { final PyClass clazz = findClass(className); - final PyFunction method = clazz.findMethodByName(name, false); - assertNotNull(method); - return method; + return clazz.findMethodByName(name, false); } protected PyClass findClass(final String name) { final Project project = myFixture.getProject(); final Collection classes = PyClassNameIndex.find(name, project, false); - assertEquals(1, classes.size()); + Assert.assertThat(String.format("Expected one class named %s", name), classes, Matchers.hasSize(1)); return classes.iterator().next(); } } diff --git a/python/testSrc/com/jetbrains/python/refactoring/classes/PyExtractSuperclassTest.java b/python/testSrc/com/jetbrains/python/refactoring/classes/PyExtractSuperclassTest.java index cae5550dc30c..baab81cc5062 100644 --- a/python/testSrc/com/jetbrains/python/refactoring/classes/PyExtractSuperclassTest.java +++ b/python/testSrc/com/jetbrains/python/refactoring/classes/PyExtractSuperclassTest.java @@ -53,7 +53,6 @@ public class PyExtractSuperclassTest extends PyClassRefactoringTest { final List members = new ArrayList(); for (String memberName : membersName) { final PyElement member = findMember(className, memberName); - assertNotNull(member); members.add(new PyMemberInfo(member)); } @@ -82,7 +81,6 @@ public class PyExtractSuperclassTest extends PyClassRefactoringTest { final PyClass clazz = findClass(className); final List members = new ArrayList(); final PyElement member = findMember(className, ".foo"); - assertNotNull(member); members.add(new PyMemberInfo(member)); final VirtualFile base_dir = myFixture.getFile().getVirtualFile().getParent(); @@ -126,7 +124,6 @@ public class PyExtractSuperclassTest extends PyClassRefactoringTest { final PyClass clazz = findClass(className); final List members = new ArrayList(); final PyElement member = findMember(className, ".foo"); - assertNotNull(member); members.add(new PyMemberInfo(member)); final VirtualFile base_dir = myFixture.getFile().getVirtualFile().getParent(); diff --git a/python/testSrc/com/jetbrains/python/refactoring/classes/PyPullUpTest.java b/python/testSrc/com/jetbrains/python/refactoring/classes/PyPullUpTest.java index 80c6b2444095..881dac2a9766 100644 --- a/python/testSrc/com/jetbrains/python/refactoring/classes/PyPullUpTest.java +++ b/python/testSrc/com/jetbrains/python/refactoring/classes/PyPullUpTest.java @@ -45,6 +45,10 @@ public class PyPullUpTest extends PyClassRefactoringTest { doHelperTest("Boo", ".boo", "Foo"); } + public void testSeveralParents() { + doHelperTest("Child", "Spam", "Parent_1"); + } + public void testMultiFile() { // PY-2810 doMultiFileTest(); } diff --git a/python/testSrc/com/jetbrains/python/refactoring/classes/PyPushDownTest.java b/python/testSrc/com/jetbrains/python/refactoring/classes/PyPushDownTest.java index cc4ca079c785..72c3bfded44a 100644 --- a/python/testSrc/com/jetbrains/python/refactoring/classes/PyPushDownTest.java +++ b/python/testSrc/com/jetbrains/python/refactoring/classes/PyPushDownTest.java @@ -55,7 +55,6 @@ public class PyPushDownTest extends PyClassRefactoringTest { final List members = new ArrayList(); for (String memberName : membersName) { final PyElement member = findMember(className, memberName); - assertNotNull(member); members.add(new PyMemberInfo(member)); }