one-line PyStatementList handles insertion correctly (PY-149)

This commit is contained in:
Dmitry Jemerov
2012-09-05 23:36:16 +02:00
parent 805753403c
commit 61dd175ed2
5 changed files with 66 additions and 3 deletions
@@ -63,6 +63,9 @@ public abstract class PyElementGenerator {
@NotNull
public abstract <T> T createFromText(LanguageLevel langLevel, Class<T> aClass, final String text);
@NotNull
public abstract <T> T createPhysicalFromText(LanguageLevel langLevel, Class<T> aClass, final String text);
/**
* Creates an arbitrary PSI element from text, by creating a bigger construction and then cutting the proper subelement.
* Will produce all kinds of exceptions if the path or class would not match the PSI tree.
@@ -42,11 +42,15 @@ public class PyElementGeneratorImpl extends PyElementGenerator {
@Override
public PsiFile createDummyFile(LanguageLevel langLevel, String contents) {
return createDummyFile(langLevel, contents, false);
}
public PsiFile createDummyFile(LanguageLevel langLevel, String contents, boolean physical) {
final PsiFileFactory factory = PsiFileFactory.getInstance(myProject);
final String name = "dummy." + PythonFileType.INSTANCE.getDefaultExtension();
final LightVirtualFile virtualFile = new LightVirtualFile(name, PythonFileType.INSTANCE, contents);
virtualFile.putUserData(LanguageLevel.KEY, langLevel);
final PsiFile psiFile = ((PsiFileFactoryImpl)factory).trySetupPsiForFile(virtualFile, PythonLanguage.getInstance(), false, true);
final PsiFile psiFile = ((PsiFileFactoryImpl)factory).trySetupPsiForFile(virtualFile, PythonLanguage.getInstance(), physical, true);
assert psiFile != null;
return psiFile;
}
@@ -233,6 +237,12 @@ public class PyElementGeneratorImpl extends PyElementGenerator {
return createFromText(langLevel, aClass, text, FROM_ROOT);
}
@NotNull
@Override
public <T> T createPhysicalFromText(LanguageLevel langLevel, Class<T> aClass, String text) {
return createFromText(langLevel, aClass, text, FROM_ROOT, true);
}
static int[] PATH_PARAMETER = {0, 3, 1};
public PyNamedParameter createParameter(@NotNull String name) {
@@ -247,7 +257,12 @@ public class PyElementGeneratorImpl extends PyElementGenerator {
@NotNull
public <T> T createFromText(LanguageLevel langLevel, Class<T> aClass, final String text, final int[] path) {
PsiElement ret = createDummyFile(langLevel, text);
return createFromText(langLevel, aClass, text, path, false);
}
@NotNull
public <T> T createFromText(LanguageLevel langLevel, Class<T> aClass, final String text, final int[] path, boolean physical) {
PsiElement ret = createDummyFile(langLevel, text, physical);
for (int skip : path) {
if (ret != null) {
ret = ret.getFirstChild();
@@ -1,6 +1,8 @@
package com.jetbrains.python.psi.impl;
import com.intellij.lang.ASTFactory;
import com.intellij.lang.ASTNode;
import com.intellij.psi.TokenType;
import com.jetbrains.python.PythonDialectsTokenSetProvider;
import com.jetbrains.python.psi.PyElementVisitor;
import com.jetbrains.python.psi.PyStatement;
@@ -22,4 +24,16 @@ public class PyStatementListImpl extends PyElementImpl implements PyStatementLis
public PyStatement[] getStatements() {
return childrenToPsi(PythonDialectsTokenSetProvider.INSTANCE.getStatementTokens(), PyStatement.EMPTY_ARRAY);
}
@Override
public ASTNode addInternal(ASTNode first, ASTNode last, ASTNode anchor, Boolean before) {
if (first.getPsi() instanceof PyStatement && getStatements().length == 1) {
ASTNode treePrev = getNode().getTreePrev();
if (treePrev != null && treePrev.getElementType() == TokenType.WHITE_SPACE && !treePrev.textContains('\n')) {
ASTNode lineBreak = ASTFactory.whitespace("\n");
treePrev.getTreeParent().replaceChild(treePrev, lineBreak);
}
}
return super.addInternal(first, last, anchor, before);
}
}
@@ -0,0 +1,30 @@
package com.jetbrains.python;
import com.intellij.openapi.command.WriteCommandAction;
import com.jetbrains.python.fixtures.PyTestCase;
import com.jetbrains.python.psi.LanguageLevel;
import com.jetbrains.python.psi.PyElementGenerator;
import com.jetbrains.python.psi.PyFunction;
import com.jetbrains.python.psi.PyStatementList;
/**
* @author yole
*/
public class PyStatementListTest extends PyTestCase {
public void testOneLineList() {
PyElementGenerator generator = PyElementGenerator.getInstance(myFixture.getProject());
PyFunction function = generator.createPhysicalFromText(LanguageLevel.PYTHON27, PyFunction.class, "def foo(): print 1");
PyFunction function2 = generator.createPhysicalFromText(LanguageLevel.PYTHON27, PyFunction.class, "def foo(): print 2");
final PyStatementList list1 = function.getStatementList();
final PyStatementList list2 = function2.getStatementList();
new WriteCommandAction.Simple(myFixture.getProject()) {
@Override
protected void run() throws Throwable {
list1.add(list2.getStatements()[0]);
}
}.execute();
assertEquals("def foo():\n print 1\n print 2", function.getText());
}
}
@@ -99,7 +99,8 @@ public class PythonAllTestsSuite {
PyPropertyAccessInspectionTest.class,
Jinja2ParserTest.class,
DjangoTemplateParserTest.class,
PyJoinLinesTest.class
PyJoinLinesTest.class,
PyStatementListTest.class
};
public static TestSuite suite() {