diff --git a/python/src/com/jetbrains/python/codeInsight/testIntegration/CreateTestAction.java b/python/src/com/jetbrains/python/codeInsight/testIntegration/CreateTestAction.java index cfa443ada339..1b21b9916726 100644 --- a/python/src/com/jetbrains/python/codeInsight/testIntegration/CreateTestAction.java +++ b/python/src/com/jetbrains/python/codeInsight/testIntegration/CreateTestAction.java @@ -36,7 +36,7 @@ public class CreateTestAction extends PsiElementBaseIntentionAction { return; } CommandProcessor.getInstance().executeCommand(project, () -> { - PsiFile e = PyTestCreator.generateTestAndNavigate(project, model); + PsiFile e = PyTestCreator.generateTestAndNavigate(element, model); final PsiDocumentManager documentManager = PsiDocumentManager.getInstance(project); documentManager.commitAllDocuments(); }, CodeInsightBundle.message("intention.create.test"), this); diff --git a/python/src/com/jetbrains/python/codeInsight/testIntegration/PyTestCreator.java b/python/src/com/jetbrains/python/codeInsight/testIntegration/PyTestCreator.java index 62f2b9d4e460..fd279f0e2641 100644 --- a/python/src/com/jetbrains/python/codeInsight/testIntegration/PyTestCreator.java +++ b/python/src/com/jetbrains/python/codeInsight/testIntegration/PyTestCreator.java @@ -17,6 +17,7 @@ import com.intellij.util.IncorrectOperationException; import com.jetbrains.python.PythonFileType; import com.jetbrains.python.codeInsight.imports.AddImportHelper; import com.jetbrains.python.psi.*; +import com.jetbrains.python.testing.PythonUnitTestUtil; import org.jetbrains.annotations.NotNull; import java.util.List; @@ -53,11 +54,12 @@ public class PyTestCreator implements TestCreator { * * @return file with test */ - static PsiFile generateTestAndNavigate(@NotNull final Project project, @NotNull final PyTestCreationModel creationModel) { + static PsiFile generateTestAndNavigate(@NotNull final PsiElement anchor, @NotNull final PyTestCreationModel creationModel) { + final Project project = anchor.getProject(); return PostprocessReformattingAspect.getInstance(project).postponeFormattingInside( () -> ApplicationManager.getApplication().runWriteAction((Computable)() -> { try { - final PyElement testClass = generateTest(project, creationModel); + final PyElement testClass = generateTest(anchor, creationModel); testClass.navigate(false); return testClass.getContainingFile(); } @@ -74,7 +76,8 @@ public class PyTestCreator implements TestCreator { * @return newly created test class */ @NotNull - static PyElement generateTest(@NotNull final Project project, @NotNull final PyTestCreationModel model) { + static PyElement generateTest(@NotNull final PsiElement anchor, @NotNull final PyTestCreationModel model) { + final Project project = anchor.getProject(); IdeDocumentHistory.getInstance(project).includeCurrentPlaceAsChangePlace(); String fileName = model.getFileName(); @@ -85,22 +88,42 @@ public class PyTestCreator implements TestCreator { final String className = model.getClassName(); final List methods = model.getMethods(); + final boolean classBased = !StringUtil.isEmptyOrSpaces(className); final LanguageLevel level = LanguageLevel.forElement(psiFile); final PyElementGenerator generator = PyElementGenerator.getInstance(project); PyElement result = psiFile; if (classBased) { + final boolean unitTestClassRequired = PythonUnitTestUtil.isTestCaseClassRequired(anchor); final StringBuilder fileText = new StringBuilder(); - fileText.append("class ").append(className).append("(TestCase):\n\t"); + fileText.append("class ").append(className); + if (unitTestClassRequired) { + fileText.append("(TestCase)"); + } + else if (LanguageLevel.forElement(anchor).isPython2()) { + fileText.append("(object)"); + } + fileText.append(":\n\t"); + + if (methods.isEmpty()) { fileText.append("pass\n"); } for (final String method : methods) { - fileText.append("def ").append(method).append("(self):\n\tself.fail()\n\n\t"); + fileText.append("def ").append(method).append("(self):\n\t"); + if (unitTestClassRequired) { + fileText.append("self.fail()"); + } + else { + fileText.append("assert False"); + } + fileText.append("\n\n\t"); + } + if (unitTestClassRequired) { + AddImportHelper.addOrUpdateFromImportStatement(psiFile, "unittest", "TestCase", null, AddImportHelper.ImportPriority.BUILTIN, + null); } - AddImportHelper.addOrUpdateFromImportStatement(psiFile, "unittest", "TestCase", null, AddImportHelper.ImportPriority.BUILTIN, - null); result = (PyElement)psiFile.addAfter(generator.createFromText(level, PyClass.class, fileText.toString()), psiFile.getLastChild()); } else { diff --git a/python/testData/create_tests/create_tst_class.expected_pytest_2k.py b/python/testData/create_tests/create_tst_class.expected_pytest_2k.py new file mode 100644 index 000000000000..327ec499b4f4 --- /dev/null +++ b/python/testData/create_tests/create_tst_class.expected_pytest_2k.py @@ -0,0 +1,6 @@ +class Spam(object): + def eggs(self): + assert False + + def eggs_and_ham(self): + assert False diff --git a/python/testData/create_tests/create_tst_class.expected_pytest_3k.py b/python/testData/create_tests/create_tst_class.expected_pytest_3k.py new file mode 100644 index 000000000000..791b7536b6b0 --- /dev/null +++ b/python/testData/create_tests/create_tst_class.expected_pytest_3k.py @@ -0,0 +1,6 @@ +class Spam: + def eggs(self): + assert False + + def eggs_and_ham(self): + assert False diff --git a/python/testData/create_tests/create_tst_class.expected.py b/python/testData/create_tests/create_tst_class.expected_unittest.py similarity index 100% rename from python/testData/create_tests/create_tst_class.expected.py rename to python/testData/create_tests/create_tst_class.expected_unittest.py diff --git a/python/testSrc/com/jetbrains/python/codeInsight/testIntegration/PyTestCreatorTest.java b/python/testSrc/com/jetbrains/python/codeInsight/testIntegration/PyTestCreatorTest.java index dd86d3503dd9..29d3c76ddbfd 100644 --- a/python/testSrc/com/jetbrains/python/codeInsight/testIntegration/PyTestCreatorTest.java +++ b/python/testSrc/com/jetbrains/python/codeInsight/testIntegration/PyTestCreatorTest.java @@ -19,7 +19,11 @@ import com.intellij.openapi.command.WriteCommandAction; import com.intellij.openapi.roots.ModuleRootManager; import com.intellij.openapi.vfs.VirtualFile; import com.intellij.psi.PsiFile; +import com.jetbrains.python.PyNames; import com.jetbrains.python.fixtures.PyTestCase; +import com.jetbrains.python.psi.LanguageLevel; +import com.jetbrains.python.testing.PythonTestConfigurationsModel; +import com.jetbrains.python.testing.TestRunnerService; import org.jetbrains.annotations.NotNull; import java.util.Arrays; @@ -30,19 +34,21 @@ import java.util.Arrays; * @author Ilya.Kazakevich */ public final class PyTestCreatorTest extends PyTestCase { - public void testCreateTest() { - myFixture.configureByFile("/create_tests/create_tst.py"); + public void testCreateUnitTest() { + final PyTestCreationModel model = prepareAndCreateModel(); + TestRunnerService testRunnerService = TestRunnerService.getInstance(myFixture.getModule()); + testRunnerService.setProjectConfiguration(PythonTestConfigurationsModel.PYTHONS_UNITTEST_NAME); + checkResult(model, "create_tst_class.expected_unittest.py"); + } - final VirtualFile[] roots = ModuleRootManager.getInstance(myFixture.getModule()).getSourceRoots(); - assert roots.length > 0 : "Empty roots for module " + myFixture.getModule(); - final VirtualFile root = roots[0]; + public void testCreatePyTest() { + final PyTestCreationModel model = prepareAndCreateModel(); + boolean p2k = LanguageLevel.forElement(myFixture.getFile()).isPython2(); + TestRunnerService testRunnerService = TestRunnerService.getInstance(myFixture.getModule()); + testRunnerService.setProjectConfiguration(PyNames.PY_TEST); - final PyTestCreationModel model = - new PyTestCreationModel("tests.py", root.getCanonicalPath(), "Spam", Arrays.asList("eggs", "eggs_and_ham")); - - - checkResult(model, "create_tst_class.expected.py"); + checkResult(model, (p2k ? "create_tst_class.expected_pytest_2k.py" : "create_tst_class.expected_pytest_3k.py")); model.setClassName(""); model.setFileName("tests_no_class.py"); @@ -50,9 +56,20 @@ public final class PyTestCreatorTest extends PyTestCase { checkResult(model, "create_tst.expected.py"); } + @NotNull + private PyTestCreationModel prepareAndCreateModel() { + myFixture.configureByFile("/create_tests/create_tst.py"); + + final VirtualFile[] roots = ModuleRootManager.getInstance(myFixture.getModule()).getSourceRoots(); + assert roots.length > 0 : "Empty roots for module " + myFixture.getModule(); + final VirtualFile root = roots[0]; + + return new PyTestCreationModel("tests.py", root.getCanonicalPath(), "Spam", Arrays.asList("eggs", "eggs_and_ham")); + } + private void checkResult(@NotNull final PyTestCreationModel model, @NotNull final String fileName) { WriteCommandAction.runWriteCommandAction(myFixture.getProject(), () -> { - final PsiFile file = PyTestCreator.generateTest(myFixture.getProject(), model).getContainingFile(); + final PsiFile file = PyTestCreator.generateTest(myFixture.getFile(), model).getContainingFile(); myFixture.configureByText(file.getFileType(), file.getText()); myFixture.checkResultByFile("/create_tests/" + fileName); });