From d41632360481555dfec415d934e3121d3d2b836b Mon Sep 17 00:00:00 2001 From: "Ilya.Kazakevich" Date: Fri, 9 Aug 2019 04:35:22 +0300 Subject: [PATCH] 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 --- .../testIntegration/CreateTestAction.java | 70 ++---------- .../testIntegration/CreateTestDialog.java | 102 ++++++++++-------- .../testIntegration/PyTestCreationModel.kt | 66 ++++++++++++ .../testIntegration/PyTestCreator.java | 57 +++++----- .../testIntegration/PyTestFinder.java | 39 +++---- .../python/testing/PythonUnitTestUtil.java | 25 ++++- .../create_tests/create_tst.expected.py | 11 +- python/testData/create_tests/create_tst.py | 7 +- .../create_tests/create_tst_class.expected.py | 9 ++ .../PyTestCreationModelTest.kt | 77 +++++++++++++ .../testIntegration/PyTestCreatorTest.java | 28 ++--- 11 files changed, 318 insertions(+), 173 deletions(-) create mode 100644 python/src/com/jetbrains/python/codeInsight/testIntegration/PyTestCreationModel.kt create mode 100644 python/testData/create_tests/create_tst_class.expected.py create mode 100644 python/testSrc/com/jetbrains/python/codeInsight/testIntegration/PyTestCreationModelTest.kt diff --git a/python/src/com/jetbrains/python/codeInsight/testIntegration/CreateTestAction.java b/python/src/com/jetbrains/python/codeInsight/testIntegration/CreateTestAction.java index a001daf88e25..cfa443ada339 100644 --- a/python/src/com/jetbrains/python/codeInsight/testIntegration/CreateTestAction.java +++ b/python/src/com/jetbrains/python/codeInsight/testIntegration/CreateTestAction.java @@ -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 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); } -} \ No newline at end of file +} diff --git a/python/src/com/jetbrains/python/codeInsight/testIntegration/CreateTestDialog.java b/python/src/com/jetbrains/python/codeInsight/testIntegration/CreateTestDialog.java index 8812b8161122..d163d1064a83 100644 --- a/python/src/com/jetbrains/python/codeInsight/testIntegration/CreateTestDialog.java +++ b/python/src/com/jetbrains/python/codeInsight/testIntegration/CreateTestDialog.java @@ -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 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> 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 getMethods() { - List 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"; } -} \ No newline at end of file +} diff --git a/python/src/com/jetbrains/python/codeInsight/testIntegration/PyTestCreationModel.kt b/python/src/com/jetbrains/python/codeInsight/testIntegration/PyTestCreationModel.kt new file mode 100644 index 000000000000..7ba559fde3a1 --- /dev/null +++ b/python/src/com/jetbrains/python/codeInsight/testIntegration/PyTestCreationModel.kt @@ -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) { + + 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 = when { + function != null -> listOf(function) + pyClass != null -> pyClass.methods.asList() + else -> (file.topLevelFunctions + file.topLevelClasses) as List + }.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 + + } + +} diff --git a/python/src/com/jetbrains/python/codeInsight/testIntegration/PyTestCreator.java b/python/src/com/jetbrains/python/codeInsight/testIntegration/PyTestCreator.java index e7b8bfa88cff..62f2b9d4e460 100644 --- a/python/src/com/jetbrains/python/codeInsight/testIntegration/PyTestCreator.java +++ b/python/src/com/jetbrains/python/codeInsight/testIntegration/PyTestCreator.java @@ -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)() -> { 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 methods = dialog.getMethods(); - if (methods.size() == 0) { - fileText.append("pass\n"); + 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 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; } } diff --git a/python/src/com/jetbrains/python/codeInsight/testIntegration/PyTestFinder.java b/python/src/com/jetbrains/python/codeInsight/testIntegration/PyTestFinder.java index f476c890165a..9d2143afa96a 100644 --- a/python/src/com/jetbrains/python/codeInsight/testIntegration/PyTestFinder.java +++ b/python/src/com/jetbrains/python/codeInsight/testIntegration/PyTestFinder.java @@ -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 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(eachClass, TestFinderHelper.calcTestNameProximity(sourceName, eachName))); + new Pair(eachClass, TestFinderHelper.calcTestNameProximity(sourceName, eachName))); } } } @@ -60,7 +60,8 @@ public class PyTestFinder implements TestFinder { Collection 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> classesWithWeights = new ArrayList<>(); + List> testsWithWeights = new ArrayList<>(); final List> 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 eachNameWithWeight : possibleNames) { + for (final Pair eachNameWithWeight : possibleNames) { for (PyClass eachClass : PyClassNameIndex.find(eachNameWithWeight.first, element.getProject(), GlobalSearchScope.projectScope(element.getProject()))) { - if (!PyTestUtil.isPyTestClass(eachClass, null)) - classesWithWeights.add(new Pair(eachClass, eachNameWithWeight.second)); + if (!PythonUnitTestUtil.isTestClass(eachClass, ThreeState.NO, null)) { + testsWithWeights.add(new Pair(eachClass, eachNameWithWeight.second)); + } } for (PyFunction function : PyFunctionNameIndex.find(eachNameWithWeight.first, element.getProject(), - GlobalSearchScope.projectScope(element.getProject()))) { - if (!PyTestUtil.isPyTestFunction(function)) - classesWithWeights.add(new Pair(function, eachNameWithWeight.second)); + GlobalSearchScope.projectScope(element.getProject()))) { + if (!PythonUnitTestUtil.isTestFunction(function, ThreeState.UNSURE, null)) { + testsWithWeights.add(new Pair(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); } } diff --git a/python/src/com/jetbrains/python/testing/PythonUnitTestUtil.java b/python/src/com/jetbrains/python/testing/PythonUnitTestUtil.java index d5316521cf0a..8219e9bd5fa4 100644 --- a/python/src/com/jetbrains/python/testing/PythonUnitTestUtil.java +++ b/python/src/com/jetbrains/python/testing/PythonUnitTestUtil.java @@ -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; diff --git a/python/testData/create_tests/create_tst.expected.py b/python/testData/create_tests/create_tst.expected.py index b09ba847eef6..fedc099427f5 100644 --- a/python/testData/create_tests/create_tst.expected.py +++ b/python/testData/create_tests/create_tst.expected.py @@ -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 diff --git a/python/testData/create_tests/create_tst.py b/python/testData/create_tests/create_tst.py index 97ec6bd2e5fa..09b05f02b12d 100644 --- a/python/testData/create_tests/create_tst.py +++ b/python/testData/create_tests/create_tst.py @@ -3,4 +3,9 @@ class Spam: pass def eggs_and_ham(self): - pass \ No newline at end of file + pass + + + +def test_foo(): + pass diff --git a/python/testData/create_tests/create_tst_class.expected.py b/python/testData/create_tests/create_tst_class.expected.py new file mode 100644 index 000000000000..b09ba847eef6 --- /dev/null +++ b/python/testData/create_tests/create_tst_class.expected.py @@ -0,0 +1,9 @@ +from unittest import TestCase + + +class Spam(TestCase): + def eggs(self): + self.fail() + + def eggs_and_ham(self): + self.fail() diff --git a/python/testSrc/com/jetbrains/python/codeInsight/testIntegration/PyTestCreationModelTest.kt b/python/testSrc/com/jetbrains/python/codeInsight/testIntegration/PyTestCreationModelTest.kt new file mode 100644 index 000000000000..08942ca35763 --- /dev/null +++ b/python/testSrc/com/jetbrains/python/codeInsight/testIntegration/PyTestCreationModelTest.kt @@ -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 { + 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 { + 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) + } +} diff --git a/python/testSrc/com/jetbrains/python/codeInsight/testIntegration/PyTestCreatorTest.java b/python/testSrc/com/jetbrains/python/codeInsight/testIntegration/PyTestCreatorTest.java index 01fce57e0cd9..dd86d3503dd9 100644 --- a/python/testSrc/com/jetbrains/python/codeInsight/testIntegration/PyTestCreatorTest.java +++ b/python/testSrc/com/jetbrains/python/codeInsight/testIntegration/PyTestCreatorTest.java @@ -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); }); } }