PY-15021: Support pytest test creation

For UnitTest you need to create class. However, pytest accepts plain function tests ("test_..").
We should not force user to use classes.

Also, many small improvement added:
* detect test folder
* create tests for everything in class
* do not create tests for tests

GitOrigin-RevId: 85dbe09e4fea47a116aaac9c00cec884f46bae9e
This commit is contained in:
Ilya.Kazakevich
2019-08-29 13:44:54 +00:00
committed by intellij-monorepo-bot
parent 2100e7c894
commit d416323604
11 changed files with 318 additions and 173 deletions
@@ -1,26 +1,17 @@
// Copyright 2000-2018 JetBrains s.r.o. Use of this source code is governed by the Apache 2.0 license that can be found in the LICENSE file.
package com.jetbrains.python.codeInsight.testIntegration;
import com.google.common.collect.Lists;
import com.intellij.codeInsight.CodeInsightBundle;
import com.intellij.codeInsight.intention.PsiElementBaseIntentionAction;
import com.intellij.openapi.command.CommandProcessor;
import com.intellij.openapi.editor.Editor;
import com.intellij.openapi.project.Project;
import com.intellij.openapi.util.text.StringUtil;
import com.intellij.psi.PsiDirectory;
import com.intellij.psi.PsiDocumentManager;
import com.intellij.psi.PsiElement;
import com.intellij.psi.PsiFile;
import com.intellij.psi.util.PsiTreeUtil;
import com.intellij.util.IncorrectOperationException;
import com.jetbrains.python.psi.PyClass;
import com.jetbrains.python.psi.PyFunction;
import com.jetbrains.python.testing.pytest.PyTestUtil;
import org.jetbrains.annotations.NotNull;
import java.util.List;
public class CreateTestAction extends PsiElementBaseIntentionAction {
@Override
@NotNull
@@ -30,67 +21,24 @@ public class CreateTestAction extends PsiElementBaseIntentionAction {
@Override
public boolean isAvailable(@NotNull Project project, Editor editor, @NotNull PsiElement element) {
PyClass psiClass = PsiTreeUtil.getParentOfType(element, PyClass.class);
if (psiClass != null && PyTestUtil.isPyTestClass(psiClass, null))
return false;
return true;
public boolean isAvailable(final @NotNull Project project, final Editor editor, final @NotNull PsiElement element) {
return PyTestCreationModel.Companion.createByElement(element) != null;
}
@Override
public void invoke(final @NotNull Project project, Editor editor, @NotNull PsiElement element) throws IncorrectOperationException {
final PyFunction srcFunction = PsiTreeUtil.getParentOfType(element, PyFunction.class);
final PyClass srcClass = PsiTreeUtil.getParentOfType(element, PyClass.class);
if (srcClass == null && srcFunction == null) return;
final PsiDirectory dir = element.getContainingFile().getContainingDirectory();
final CreateTestDialog d = new CreateTestDialog(project);
if (srcClass != null) {
d.setClassName("Test" + StringUtil.capitalize(srcClass.getName()));
d.setFileName("test_" + StringUtil.decapitalize(srcClass.getName()) + ".py");
if (dir != null)
d.setTargetDir(dir.getVirtualFile().getPath());
if (srcFunction != null) {
d.methodsSize(1);
d.addMethod("test_" + srcFunction.getName(), 0);
}
else {
final List<PyFunction> methods = Lists.newArrayList();
srcClass.visitMethods(pyFunction -> {
if (pyFunction.getName() != null && !pyFunction.getName().startsWith("__"))
methods.add(pyFunction);
return true;
}, false, null);
d.methodsSize(methods.size());
int i = 0;
for (PyFunction f : methods) {
d.addMethod("test_" + f.getName(), i);
++i;
}
}
final PyTestCreationModel model =
PyTestCreationModel.Companion.createByElement(element);
if (model == null) {
return;
}
else {
d.setClassName("Test" + StringUtil.capitalize(srcFunction.getName()));
d.setFileName("test_" + StringUtil.decapitalize(srcFunction.getName()) + ".py");
if (dir != null)
d.setTargetDir(dir.getVirtualFile().getPath());
d.methodsSize(1);
d.addMethod("test_" + srcFunction.getName(), 0);
}
if (!d.showAndGet()) {
if (!CreateTestDialog.userAcceptsTestCreation(project, model)) {
return;
}
CommandProcessor.getInstance().executeCommand(project, () -> {
PsiFile e = PyTestCreator.generateTestAndNavigate(project, d);
PsiFile e = PyTestCreator.generateTestAndNavigate(project, model);
final PsiDocumentManager documentManager = PsiDocumentManager.getInstance(project);
documentManager.commitAllDocuments();
}, CodeInsightBundle.message("intention.create.test"), this);
}
}
}
@@ -8,6 +8,8 @@ import com.intellij.openapi.ui.TextFieldWithBrowseButton;
import com.intellij.openapi.util.text.StringUtil;
import com.intellij.ui.BooleanTableCellRenderer;
import com.intellij.ui.TableUtil;
import one.util.streamex.StreamEx;
import org.jetbrains.annotations.NotNull;
import javax.swing.*;
import javax.swing.event.DocumentEvent;
@@ -18,18 +20,25 @@ import java.awt.event.ActionEvent;
import java.awt.event.ActionListener;
import java.util.ArrayList;
import java.util.List;
import java.util.Vector;
public class CreateTestDialog extends DialogWrapper {
public final class CreateTestDialog extends DialogWrapper {
@NotNull
private final PyTestCreationModel myModel;
private final boolean myClassRequired;
private TextFieldWithBrowseButton myTargetDir;
private JTextField myClassName;
private JPanel myMainPanel;
private JTextField myFileName;
private JTable myMethodsTable;
private DefaultTableModel myTableModel;
@NotNull
private final DefaultTableModel myTableModel;
protected CreateTestDialog(Project project) {
private CreateTestDialog(@NotNull final Project project, @NotNull final PyTestCreationModel model) {
super(project);
init();
myClassRequired = StringUtil.isNotEmpty(model.getClassName());
myModel = model;
myTargetDir.addBrowseFolderListener("Select target directory", null, project,
FileChooserDescriptorFactory.createSingleFolderDescriptor());
myTargetDir.setEditable(false);
@@ -47,10 +56,21 @@ public class CreateTestDialog extends DialogWrapper {
addUpdater(myFileName);
addUpdater(myClassName);
}
public void methodsSize(int methods) {
myTableModel = new DefaultTableModel(methods, 2);
//Fill UI with model
myTargetDir.setText(model.getTargetDir());
myFileName.setText(model.getFileName());
final String clazz = model.getClassName();
myClassName.setText(clazz);
final List<String> methods = model.getMethods();
final String[] columnNames = new String[]{"", "Test function"};
myTableModel = new DefaultTableModel(
methods.stream().map(name -> new Object[]{Boolean.FALSE, name}).toArray(size -> new Object[size][columnNames.length]),
columnNames
);
// If only one method, then select it by default
if (methods.size() == 1) {
myTableModel.setValueAt(Boolean.TRUE, myTableModel.getRowCount() - 1, 0);
}
myMethodsTable.setModel(myTableModel);
TableColumn checkColumn = myMethodsTable.getColumnModel().getColumn(0);
@@ -58,14 +78,31 @@ public class CreateTestDialog extends DialogWrapper {
checkColumn.setCellRenderer(new BooleanTableCellRenderer());
checkColumn.setCellEditor(new DefaultCellEditor(new JCheckBox()));
myMethodsTable.getColumnModel().getColumn(1).setHeaderValue("Test method");
checkColumn.setHeaderValue("");
getOKAction().setEnabled(true);
getOKAction().setEnabled(isValid());
}
protected void addUpdater(JTextField field) {
field.getDocument().addDocumentListener(new MyDocumentListener());
static boolean userAcceptsTestCreation(@NotNull final Project project, @NotNull final PyTestCreationModel model) {
final CreateTestDialog dialog = new CreateTestDialog(project, model);
if (!dialog.showAndGet()) {
return false;
}
dialog.copyToModel();
return true;
}
private void copyToModel() {
myModel.setClassName(myClassName.getText());
myModel.setFileName(myFileName.getText());
myModel.setTargetDir(myTargetDir.getText());
@SuppressWarnings("unchecked")
StreamEx<Vector<Object>> methods = StreamEx.of(myTableModel.getDataVector().stream());
myModel.setMethods(new ArrayList<>(methods.map(v -> (v.get(0) == Boolean.TRUE) ? v.get(1).toString() : null).nonNull().toList()));
}
private void addUpdater(JTextField field) {
field.getDocument().addDocumentListener(new MyDocumentListener());
}
private class MyDocumentListener implements DocumentListener {
@Override
public void insertUpdate(DocumentEvent documentEvent) {
@@ -84,8 +121,9 @@ public class CreateTestDialog extends DialogWrapper {
}
private boolean isValid() {
return !StringUtil.isEmptyOrSpaces(getTargetDir()) && !StringUtil.isEmptyOrSpaces(getClassName())
&& !StringUtil.isEmptyOrSpaces(getFileName());
return !StringUtil.isEmptyOrSpaces(getTargetDir())
&& (!myClassRequired || !StringUtil.isEmptyOrSpaces(getClassName()))
&& !StringUtil.isEmptyOrSpaces(getFileName());
}
@Override
@@ -93,48 +131,20 @@ public class CreateTestDialog extends DialogWrapper {
return myMainPanel;
}
public String getTargetDir() {
private String getTargetDir() {
return myTargetDir.getText().trim();
}
public void setTargetDir(String text) {
myTargetDir.setText(text);
}
public void setClassName(String text) {
myClassName.setText(text);
}
public void setFileName(String text) {
myFileName.setText(text);
}
public String getClassName() {
private String getClassName() {
return myClassName.getText().trim();
}
public String getFileName() {
private String getFileName() {
return myFileName.getText().trim();
}
public void addMethod(String name, int row) {
myTableModel.setValueAt(name, row, 1);
myTableModel.setValueAt(Boolean.FALSE, row, 0);
}
public List<String> getMethods() {
List<String> res = new ArrayList<>();
for (int i = 0; i != myTableModel.getRowCount(); ++i) {
Object val = myTableModel.getValueAt(i, 0);
if (val != null && (Boolean)val == true)
res.add((String)myTableModel.getValueAt(i, 1));
}
return res;
}
@Override
protected String getHelpId() {
return "reference.dialogs.createTestsFromGoTo";
}
}
}
@@ -0,0 +1,66 @@
// Copyright 2000-2019 JetBrains s.r.o. Use of this source code is governed by the Apache 2.0 license that can be found in the LICENSE file.
package com.jetbrains.python.codeInsight.testIntegration
import com.intellij.openapi.module.ModuleUtil
import com.intellij.openapi.vfs.VirtualFile
import com.intellij.psi.PsiElement
import com.intellij.psi.PsiNamedElement
import com.intellij.psi.search.FilenameIndex
import com.intellij.psi.util.PsiTreeUtil
import com.jetbrains.python.codeInsight.testIntegration.PyTestCreationModel.Companion.createByElement
import com.jetbrains.python.psi.PyClass
import com.jetbrains.python.psi.PyFile
import com.jetbrains.python.psi.PyFunction
import com.jetbrains.python.testing.PythonUnitTestUtil
/**
* Created with [createByElement], then modified my user and provided to [PyTestCreator.createTest] to create test
*/
class PyTestCreationModel(var fileName: String,
var targetDir: String,
var className: String,
var methods: List<String>) {
init {
assert(methods.isNotEmpty()) { "Provide at least one method" }
}
companion object {
/**
* @return model of null if no test could be created for this element
*/
fun createByElement(element: PsiElement): PyTestCreationModel? {
if (PythonUnitTestUtil.isTestElement(element, null)) return null //Can't create tests for tests
val file = element.containingFile as? PyFile ?: return null
val pyClass = PsiTreeUtil.getParentOfType(element, PyClass::class.java, false)
val function = PsiTreeUtil.getParentOfType(element, PyFunction::class.java, false)
val elementsToTest: Sequence<PsiNamedElement> = when {
function != null -> listOf(function)
pyClass != null -> pyClass.methods.asList()
else -> (file.topLevelFunctions + file.topLevelClasses) as List<PsiNamedElement>
}.asSequence().filterNot { PythonUnitTestUtil.isTestElement(it, null) }
val functionNames = elementsToTest
.filterNot { it.name?.startsWith("__") == true }
.mapNotNull { it.name }
.map { "test_${it.toLowerCase()}" }.toList()
return if (functionNames.isEmpty()) null
else {
val className = if (PythonUnitTestUtil.isTestCaseClassRequired(file)) "Test${pyClass?.name ?: ""}" else ""
PyTestCreationModel(fileName = "test_${file.name}",
targetDir = getTestFolder(element).path,
className = className,
methods = functionNames)
}
}
private fun getTestFolder(element: PsiElement): VirtualFile =
ModuleUtil.findModuleForPsiElement(element)?.let { module ->
FilenameIndex.getVirtualFilesByName(element.project, "tests", module.moduleContentScope).firstOrNull()
} ?: element.containingFile.containingDirectory.virtualFile
}
}
@@ -7,6 +7,7 @@ import com.intellij.openapi.editor.Editor;
import com.intellij.openapi.fileEditor.ex.IdeDocumentHistory;
import com.intellij.openapi.project.Project;
import com.intellij.openapi.util.Computable;
import com.intellij.openapi.util.text.StringUtil;
import com.intellij.psi.PsiElement;
import com.intellij.psi.PsiFile;
import com.intellij.psi.codeStyle.CodeStyleManager;
@@ -38,7 +39,7 @@ public class PyTestCreator implements TestCreator {
try {
CreateTestAction action = new CreateTestAction();
PsiElement element = file.findElementAt(editor.getCaretModel().getOffset());
if (action.isAvailable(project, editor, element)) {
if (element != null && action.isAvailable(project, editor, element)) {
action.invoke(project, editor, file.getContainingFile());
}
}
@@ -52,11 +53,11 @@ public class PyTestCreator implements TestCreator {
*
* @return file with test
*/
static PsiFile generateTestAndNavigate(@NotNull final Project project, @NotNull final CreateTestDialog dialog) {
static PsiFile generateTestAndNavigate(@NotNull final Project project, @NotNull final PyTestCreationModel creationModel) {
return PostprocessReformattingAspect.getInstance(project).postponeFormattingInside(
() -> ApplicationManager.getApplication().runWriteAction((Computable<PsiFile>)() -> {
try {
final PyElement testClass = generateTest(project, dialog);
final PyElement testClass = generateTest(project, creationModel);
testClass.navigate(false);
return testClass.getContainingFile();
}
@@ -68,41 +69,49 @@ public class PyTestCreator implements TestCreator {
}
/**
* Generates test, puts it into file and returns class element for test
* Generates test, puts it into file and returns element to navigate to
*
* @return newly created test class
*/
@NotNull
static PyElement generateTest(@NotNull final Project project, @NotNull final CreateTestDialog dialog) {
static PyElement generateTest(@NotNull final Project project, @NotNull final PyTestCreationModel model) {
IdeDocumentHistory.getInstance(project).includeCurrentPlaceAsChangePlace();
String fileName = dialog.getFileName();
String fileName = model.getFileName();
if (!fileName.endsWith(".py")) {
fileName = fileName + "." + PythonFileType.INSTANCE.getDefaultExtension();
}
final PyFile psiFile = PyUtil.getOrCreateFile(model.getTargetDir() + "/" + fileName, project);
StringBuilder fileText = new StringBuilder();
fileText.append("class ").append(dialog.getClassName()).append("(TestCase):\n\t");
List<String> methods = dialog.getMethods();
if (methods.size() == 0) {
fileText.append("pass\n");
final String className = model.getClassName();
final List<String> 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 StringBuilder fileText = new StringBuilder();
fileText.append("class ").append(className).append("(TestCase):\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");
}
AddImportHelper.addOrUpdateFromImportStatement(psiFile, "unittest", "TestCase", null, AddImportHelper.ImportPriority.BUILTIN,
null);
result = (PyElement)psiFile.addAfter(generator.createFromText(level, PyClass.class, fileText.toString()), psiFile.getLastChild());
}
for (String method : methods) {
fileText.append("def ").append(method).append("(self):\n\tself.fail()\n\n\t");
else {
for (final String method : methods) {
final PyFunction fun = generator.createFromText(level, PyFunction.class, String.format("def %s():\n\tassert False\n\n", method));
psiFile.addAfter(fun, psiFile.getLastChild());
}
}
PsiFile psiFile = PyUtil.getOrCreateFile(
dialog.getTargetDir() + "/" + fileName, project);
AddImportHelper.addOrUpdateFromImportStatement(psiFile, "unittest", "TestCase", null, AddImportHelper.ImportPriority.BUILTIN,
null);
PyElement createdClass = PyElementGenerator.getInstance(project).createFromText(
LanguageLevel.forElement(psiFile), PyClass.class, fileText.toString());
createdClass = (PyElement)psiFile.addAfter(createdClass, psiFile.getLastChild());
PostprocessReformattingAspect.getInstance(project).doPostponedFormatting(psiFile.getViewProvider());
CodeStyleManager.getInstance(project).reformat(psiFile);
return createdClass;
return result;
}
}
@@ -16,7 +16,6 @@ import com.jetbrains.python.psi.stubs.PyClassNameIndex;
import com.jetbrains.python.psi.stubs.PyFunctionNameIndex;
import com.jetbrains.python.testing.PythonUnitTestUtil;
import com.jetbrains.python.testing.doctest.PythonDocTestUtil;
import com.jetbrains.python.testing.pytest.PyTestUtil;
import org.jetbrains.annotations.NotNull;
import java.util.ArrayList;
@@ -47,10 +46,11 @@ public class PyTestFinder implements TestFinder {
Collection<String> names = PyClassNameIndex.allKeys(element.getProject());
for (String eachName : names) {
if (eachName.contains(sourceName)) {
for (PyClass eachClass : PyClassNameIndex.find(eachName, element.getProject(), GlobalSearchScope.projectScope(element.getProject()))) {
for (PyClass eachClass : PyClassNameIndex
.find(eachName, element.getProject(), GlobalSearchScope.projectScope(element.getProject()))) {
if (PythonUnitTestUtil.isTestClass(eachClass, ThreeState.UNSURE, null) || PythonDocTestUtil.isDocTestClass(eachClass)) {
classesWithProximities.add(
new Pair<PsiNamedElement, Integer>(eachClass, TestFinderHelper.calcTestNameProximity(sourceName, eachName)));
new Pair<PsiNamedElement, Integer>(eachClass, TestFinderHelper.calcTestNameProximity(sourceName, eachName)));
}
}
}
@@ -60,7 +60,8 @@ public class PyTestFinder implements TestFinder {
Collection<String> names = PyFunctionNameIndex.allKeys(element.getProject());
for (String eachName : names) {
if (eachName.contains(sourceName)) {
for (PyFunction eachFunction : PyFunctionNameIndex.find(eachName, element.getProject(), GlobalSearchScope.projectScope(element.getProject()))) {
for (PyFunction eachFunction : PyFunctionNameIndex
.find(eachName, element.getProject(), GlobalSearchScope.projectScope(element.getProject()))) {
if (PythonUnitTestUtil.isTestFunction(
eachFunction, ThreeState.UNSURE, null) || PythonDocTestUtil.isDocTestFunction(eachFunction)) {
classesWithProximities.add(
@@ -80,34 +81,34 @@ public class PyTestFinder implements TestFinder {
final PyClass source = PsiTreeUtil.getParentOfType(element, PyClass.class);
if (sourceFunction == null && source == null) return Collections.emptySet();
List<Pair<? extends PsiNamedElement, Integer>> classesWithWeights = new ArrayList<>();
List<Pair<? extends PsiNamedElement, Integer>> testsWithWeights = new ArrayList<>();
final List<Pair<String, Integer>> possibleNames = new ArrayList<>();
if (source != null)
if (source != null) {
possibleNames.addAll(TestFinderHelper.collectPossibleClassNamesWithWeights(source.getName()));
if (sourceFunction != null)
}
if (sourceFunction != null) {
possibleNames.addAll(TestFinderHelper.collectPossibleClassNamesWithWeights(sourceFunction.getName()));
}
for (Pair<String, Integer> eachNameWithWeight : possibleNames) {
for (final Pair<String, Integer> eachNameWithWeight : possibleNames) {
for (PyClass eachClass : PyClassNameIndex.find(eachNameWithWeight.first, element.getProject(),
GlobalSearchScope.projectScope(element.getProject()))) {
if (!PyTestUtil.isPyTestClass(eachClass, null))
classesWithWeights.add(new Pair<PsiNamedElement, Integer>(eachClass, eachNameWithWeight.second));
if (!PythonUnitTestUtil.isTestClass(eachClass, ThreeState.NO, null)) {
testsWithWeights.add(new Pair<PsiNamedElement, Integer>(eachClass, eachNameWithWeight.second));
}
}
for (PyFunction function : PyFunctionNameIndex.find(eachNameWithWeight.first, element.getProject(),
GlobalSearchScope.projectScope(element.getProject()))) {
if (!PyTestUtil.isPyTestFunction(function))
classesWithWeights.add(new Pair<PsiNamedElement, Integer>(function, eachNameWithWeight.second));
GlobalSearchScope.projectScope(element.getProject()))) {
if (!PythonUnitTestUtil.isTestFunction(function, ThreeState.UNSURE, null)) {
testsWithWeights.add(new Pair<PsiNamedElement, Integer>(function, eachNameWithWeight.second));
}
}
}
return TestFinderHelper.getSortedElements(classesWithWeights, false);
return TestFinderHelper.getSortedElements(testsWithWeights, false);
}
@Override
public boolean isTest(@NotNull PsiElement element) {
PyClass cl = PsiTreeUtil.getParentOfType(element, PyClass.class, false);
if (cl != null)
return PyTestUtil.isPyTestClass(cl, null);
return false;
return PythonUnitTestUtil.isTestElement(element, null);
}
}
@@ -15,6 +15,7 @@ import com.intellij.openapi.vfs.VirtualFile;
import com.intellij.psi.PsiElement;
import com.intellij.psi.PsiFile;
import com.intellij.psi.PsiManager;
import com.intellij.psi.util.PsiTreeUtil;
import com.intellij.util.ThreeState;
import com.jetbrains.extensions.python.PyClassExtKt;
import com.jetbrains.python.psi.PyClass;
@@ -50,7 +51,6 @@ public final class PythonUnitTestUtil {
private PythonUnitTestUtil() {
}
public static boolean isTestFile(@NotNull final PyFile file,
@NotNull final ThreeState testCaseClassRequired,
@Nullable final TypeEvalContext context) {
@@ -74,6 +74,25 @@ public final class PythonUnitTestUtil {
return isTestClass(cls, ThreeState.YES, context);
}
/**
* If element itself is test or situated inside of test
*/
public static boolean isTestElement(@NotNull final PsiElement element, @Nullable final TypeEvalContext context) {
final PyFunction fun = PsiTreeUtil.getParentOfType(element, PyFunction.class, false);
if (fun != null) {
if (isTestFunction(fun, ThreeState.UNSURE, context)) {
return true;
}
}
final PyClass clazz = PsiTreeUtil.getParentOfType(element, PyClass.class, false);
if (clazz != null && isTestClass(clazz, ThreeState.UNSURE, context)) {
return true;
}
return element instanceof PyFile && isTestFile((PyFile)element, ThreeState.UNSURE, context);
}
public static boolean isTestClass(@NotNull final PyClass cls,
@NotNull final ThreeState testCaseClassRequired,
@Nullable TypeEvalContext context) {
@@ -145,6 +164,10 @@ public final class PythonUnitTestUtil {
if (userProvidedValue != ThreeState.UNSURE) {
return userProvidedValue.toBoolean();
}
return isTestCaseClassRequired(anchor);
}
public static boolean isTestCaseClassRequired(@NotNull final PsiElement anchor) {
final Module module = ModuleUtilCore.findModuleForPsiElement(anchor);
if (module == null) {
return true;
@@ -1,9 +1,6 @@
from unittest import TestCase
def eggs():
assert False
class Spam(TestCase):
def eggs(self):
self.fail()
def eggs_and_ham(self):
self.fail()
def eggs_and_ham():
assert False
+6 -1
View File
@@ -3,4 +3,9 @@ class Spam:
pass
def eggs_and_ham(self):
pass
pass
def test_foo():
pass
@@ -0,0 +1,9 @@
from unittest import TestCase
class Spam(TestCase):
def eggs(self):
self.fail()
def eggs_and_ham(self):
self.fail()
@@ -0,0 +1,77 @@
// Copyright 2000-2019 JetBrains s.r.o. Use of this source code is governed by the Apache 2.0 license that can be found in the LICENSE file.
package com.jetbrains.python.codeInsight.testIntegration
import com.intellij.openapi.application.ApplicationManager
import com.intellij.openapi.application.WriteAction
import com.intellij.openapi.vfs.VfsUtil
import com.intellij.psi.PsiElement
import com.jetbrains.python.PyNames
import com.jetbrains.python.fixtures.PyTestCase
import com.jetbrains.python.psi.PyFile
import com.jetbrains.python.testing.PyTestFrameworkService
import com.jetbrains.python.testing.PythonTestConfigurationsModel
import com.jetbrains.python.testing.TestRunnerService
import org.junit.Assert
class PyTestCreationModelTest : PyTestCase() {
private val dir get() = myFixture.file.containingDirectory.virtualFile
private val dirPath get() = dir.path
private val service: TestRunnerService get() = TestRunnerService.getInstance(myFixture.module)
private val testsFolderName = "tests"
fun testWithUnitTest() {
service.projectConfiguration = PythonTestConfigurationsModel.PYTHONS_UNITTEST_NAME
val modelToTestClass = getModel(true)!!
Assert.assertEquals("test_create_tst.py", modelToTestClass.fileName)
Assert.assertEquals("TestSpam", modelToTestClass.className)
Assert.assertEquals(dirPath, modelToTestClass.targetDir)
Assert.assertEquals(modelToTestClass.methods, listOf("test_eggs", "test_eggs_and_ham"))
val modelToTestFunction = getModel(false)!!
Assert.assertEquals("test_create_tst.py", modelToTestFunction.fileName)
Assert.assertEquals("Test", modelToTestFunction.className)
Assert.assertEquals(dirPath, modelToTestClass.targetDir)
Assert.assertEquals(modelToTestFunction.methods, listOf("test_test_foo"))
}
fun testWithPyTest() {
service.projectConfiguration = PyTestFrameworkService.getSdkReadableNameByFramework(PyNames.PY_TEST)
val modelToTestClass = getModel(true)!!
Assert.assertEquals("test_create_tst.py", modelToTestClass.fileName)
Assert.assertEquals("", modelToTestClass.className)
Assert.assertEquals(dirPath, modelToTestClass.targetDir)
Assert.assertEquals(modelToTestClass.methods, listOf("test_eggs", "test_eggs_and_ham"))
Assert.assertNull("test_foo is test from pytest point of view, can't test it", getModel(false))
}
fun testTestFolderDetected() {
ApplicationManager.getApplication().invokeAndWait {
WriteAction.runAndWait<Throwable> {
VfsUtil.createDirectoryIfMissing(dir, testsFolderName)
}
}
val modelToTestClass = getModel(true)!!
Assert.assertEquals(dir.findChild(testsFolderName)!!.path, modelToTestClass.targetDir)
}
override fun setUp() {
super.setUp()
myFixture.configureByFile("/create_tests/create_tst.py")
}
override fun tearDown() {
ApplicationManager.getApplication().invokeAndWait {
WriteAction.runAndWait<Throwable> {
dir.findChild(testsFolderName)?.delete(this)
}
}
super.tearDown()
}
private fun getModel(forClass: Boolean): PyTestCreationModel? {
val pyFile = myFixture.file as PyFile
val element: PsiElement = if (forClass) pyFile.topLevelClasses[0] else pyFile.findTopLevelFunction("test_foo")!!
return PyTestCreationModel.createByElement(element)
}
}
@@ -20,13 +20,10 @@ import com.intellij.openapi.roots.ModuleRootManager;
import com.intellij.openapi.vfs.VirtualFile;
import com.intellij.psi.PsiFile;
import com.jetbrains.python.fixtures.PyTestCase;
import org.easymock.IMocksControl;
import org.jetbrains.annotations.NotNull;
import java.util.Arrays;
import static org.easymock.EasyMock.createNiceControl;
import static org.easymock.EasyMock.expect;
/**
* Checks how test classes are created
*
@@ -41,20 +38,23 @@ public final class PyTestCreatorTest extends PyTestCase {
assert roots.length > 0 : "Empty roots for module " + myFixture.getModule();
final VirtualFile root = roots[0];
final IMocksControl mockControl = createNiceControl();
final CreateTestDialog dialog = mockControl.createMock(CreateTestDialog.class);
expect(dialog.getFileName()).andReturn("tests.py").anyTimes();
expect(dialog.getClassName()).andReturn("Spam").anyTimes();
// Target dir is first module source
expect(dialog.getTargetDir()).andReturn(root.getCanonicalPath()).anyTimes();
expect(dialog.getMethods()).andReturn(Arrays.asList("eggs", "eggs_and_ham")).anyTimes();
mockControl.replay();
final PyTestCreationModel model =
new PyTestCreationModel("tests.py", root.getCanonicalPath(), "Spam", Arrays.asList("eggs", "eggs_and_ham"));
checkResult(model, "create_tst_class.expected.py");
model.setClassName("");
model.setFileName("tests_no_class.py");
checkResult(model, "create_tst.expected.py");
}
private void checkResult(@NotNull final PyTestCreationModel model, @NotNull final String fileName) {
WriteCommandAction.runWriteCommandAction(myFixture.getProject(), () -> {
final PsiFile file = PyTestCreator.generateTest(myFixture.getProject(), dialog).getContainingFile();
final PsiFile file = PyTestCreator.generateTest(myFixture.getProject(), model).getContainingFile();
myFixture.configureByText(file.getFileType(), file.getText());
myFixture.checkResultByFile("/create_tests/create_tst.expected.py");
myFixture.checkResultByFile("/create_tests/" + fileName);
});
}
}