PY-11958 PullUp refactoring breaks class signature if class has more than 2 ancestors

This commit is contained in:
Ilya.Kazakevich
2014-01-29 19:46:58 +04:00
parent 6c989e646d
commit f98b4ed4da
7 changed files with 85 additions and 32 deletions
@@ -49,7 +49,8 @@ public class PyClassRefactoringUtil {
private static final Key<Boolean> ENCODED_USE_FROM_IMPORT = Key.create("PyEncodedUseFromImport");
private static final Key<String> ENCODED_IMPORT_AS = Key.create("PyEncodedImportAs");
private PyClassRefactoringUtil() {}
private PyClassRefactoringUtil() {
}
public static void moveSuperclasses(PyClass clazz, Set<String> 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<PyExpression> superClassesAsPsi,
Collection<String> 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<PyExpression> removeAndGetSuperClasses(PyClass clazz, Set<String> superClasses) {
if (superClasses.size() == 0) return Collections.emptyList();
final List<PyExpression> toAdd = new ArrayList<PyExpression>();
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<PyExpression> removeAndGetSuperClasses(@NotNull PyClass clazz, @NotNull Set<String> superClassesToRemove) {
if (superClassesToRemove.isEmpty()) {
return Collections.emptyList();
}
final List<PyExpression> result = new ArrayList<PyExpression>();
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<String> 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
@@ -0,0 +1,13 @@
class Spam:
pass
class Parent_1(object, Spam):
pass
class Parent_2():
pass
class Child(Parent_1, Parent_2):
pass
@@ -0,0 +1,13 @@
class Spam:
pass
class Parent_1(object):
pass
class Parent_2():
pass
class Child(Parent_1, Parent_2, Spam):
pass
@@ -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 (<code>.</code>) 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<PyClass> 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();
}
}
@@ -53,7 +53,6 @@ public class PyExtractSuperclassTest extends PyClassRefactoringTest {
final List<PyMemberInfo> members = new ArrayList<PyMemberInfo>();
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<PyMemberInfo> members = new ArrayList<PyMemberInfo>();
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<PyMemberInfo> members = new ArrayList<PyMemberInfo>();
final PyElement member = findMember(className, ".foo");
assertNotNull(member);
members.add(new PyMemberInfo(member));
final VirtualFile base_dir = myFixture.getFile().getVirtualFile().getParent();
@@ -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();
}
@@ -55,7 +55,6 @@ public class PyPushDownTest extends PyClassRefactoringTest {
final List<PyMemberInfo> members = new ArrayList<PyMemberInfo>();
for (String memberName : membersName) {
final PyElement member = findMember(className, memberName);
assertNotNull(member);
members.add(new PyMemberInfo(member));
}