fixed PY-8565 Unwrap/Remove is not available in for loop

This commit is contained in:
Ekaterina Tuzova
2013-01-24 11:13:29 +04:00
parent 6f3d0a2dd9
commit d2c5e0106a
7 changed files with 46 additions and 1 deletions
@@ -221,6 +221,7 @@ surround.with.try.except.template=try / except
##########################################################################################################################
unwrap.if=Unwrap if...
unwrap.while=Unwrap while...
unwrap.for=Unwrap for...
unwrap.try=Unwrap try...
unwrap.else=Unwrap else...
unwrap.elif=Unwrap elif...
@@ -0,0 +1,34 @@
package com.jetbrains.python.refactoring.unwrap;
import com.intellij.psi.PsiElement;
import com.intellij.util.IncorrectOperationException;
import com.jetbrains.python.PyBundle;
import com.jetbrains.python.psi.*;
/**
* User : ktisha
*/
public class PyForUnwrapper extends PyUnwrapper {
public PyForUnwrapper() {
super(PyBundle.message("unwrap.for"));
}
public boolean isApplicableTo(PsiElement e) {
if (e instanceof PyForStatement) {
final PyStatementList statementList = ((PyForStatement)e).getForPart().getStatementList();
if (statementList != null) {
final PyStatement[] statements = statementList.getStatements();
return statements.length == 1 && !(statements[0] instanceof PyPassStatement) || statements.length > 1;
}
}
return false;
}
@Override
protected void doUnwrap(final PsiElement element, final Context context) throws IncorrectOperationException {
final PyForStatement forStatement = (PyForStatement)element;
context.extractPart(forStatement);
context.delete(forStatement);
}
}
@@ -16,7 +16,8 @@ public class PyUnwrapDescriptor extends UnwrapDescriptorBase{
new PyElseUnwrapper(),
new PyElIfUnwrapper(),
new PyElIfRemover(),
new PyTryUnwrapper()
new PyTryUnwrapper(),
new PyForUnwrapper()
};
}
}
@@ -65,6 +65,10 @@ public abstract class PyUnwrapper extends AbstractUnwrapper<PyUnwrapper.Context>
final PyTryPart part = ((PyTryExceptStatement)from).getTryPart();
statementList = part.getStatementList();
}
else if (from instanceof PyForStatement) {
final PyForPart part = ((PyForStatement)from).getForPart();
statementList = part.getStatementList();
}
if (statementList != null)
extract(statementList.getFirstChild(), statementList.getLastChild(), from);
}
@@ -0,0 +1 @@
print 1
@@ -0,0 +1,2 @@
for item in range(1):
prin<caret>t 1
@@ -35,6 +35,8 @@ public class PyUnwrapperTest extends PyTestCase {
public void testTryUnwrap() throws Throwable {doTest();}
public void testForUnwrap() throws Throwable {doTest();}
private void doTest() {
doTest(0);
}