more tests and improvements for python requirements generation (PY-18847)

GitOrigin-RevId: 2176648fd9021b3734e29c72b9daaa98e81a108a
This commit is contained in:
Aleksei Kniazev
2020-03-12 09:39:39 +00:00
committed by intellij-monorepo-bot
parent 539dab2069
commit 6cb5115bed
27 changed files with 190 additions and 108 deletions
@@ -63,7 +63,9 @@ public interface PyRequirement {
default boolean isEditable() {
return getInstallOptions().size() > 0 && "-e".equals(getInstallOptions().get(0));
if (getInstallOptions().isEmpty()) return false;
String firstOption = getInstallOptions().get(0);
return "-e".equals(firstOption) || "--editable".equals(firstOption);
}
/**
@@ -100,7 +100,9 @@ public class PyPackageRequirementsInspection extends PyInspection {
registerProblem(file, PyBundle.message("python.requirements.file.empty"),
ProblemHighlightType.GENERIC_ERROR_OR_WARNING, null, new PyGenerateRequirementsFileQuickFix(module));
}
else checkPackagesHaveBeenInstalled(file, module);
else {
checkPackagesHaveBeenInstalled(file, module);
}
}
}
@@ -2,6 +2,7 @@
package com.jetbrains.python.packaging
import com.intellij.openapi.vfs.VirtualFile
import com.intellij.psi.PsiDirectory
import com.intellij.psi.PsiFile
import com.jetbrains.python.PyBundle
@@ -16,6 +17,11 @@ class PyRequirementsFileVisitor(private val importedPackages: MutableMap<String,
fun visitRequirementsFile(requirementsFile: PsiFile): PyRequirementsAnalysisResult {
doVisitFile(requirementsFile, mutableSetOf(requirementsFile.virtualFile))
val currentFileOutput = collectedOutput.remove(requirementsFile.virtualFile)!!
importedPackages.values.asSequence()
.map { if (settings.specifyVersion) "${it.name}${settings.versionSpecifier.separator}${it.version}" else it.name }
.forEach { currentFileOutput.add(it) }
return PyRequirementsAnalysisResult(currentFileOutput, collectedOutput, unmatchedLines, unchangedInBaseFiles)
}
@@ -43,15 +49,9 @@ class PyRequirementsFileVisitor(private val importedPackages: MutableMap<String,
}
if (settings.modifyBaseFiles && isFileReference(line)) {
// collecting requirements from base files for modification
val filename = line.split(" ")[1]
val dir = requirementsFile.containingDirectory.virtualFile
val virtualFile = dir.findFileByRelativePath(filename)
if (virtualFile != null && virtualFile !in visitedFiles) {
visitedFiles.add(virtualFile)
val baseRequirementsFile = requirementsFile.manager.findFile(virtualFile)!!
doVisitFile(baseRequirementsFile, visitedFiles)
visitedFiles.remove(virtualFile)
}
visitBaseFile(filename, requirementsFile.containingDirectory, visitedFiles)
outputLines.addAll(lines) // always keeping base file reference
continue
}
@@ -63,7 +63,7 @@ class PyRequirementsFileVisitor(private val importedPackages: MutableMap<String,
continue
}
// report those requirements that were not changed in base files
// base files cannot be modified, so we report requirements with different version
parsed.asSequence()
.filter { it.name.toLowerCase() in importedPackages }
.map { it to importedPackages.remove(it.name.toLowerCase()) }
@@ -79,20 +79,10 @@ class PyRequirementsFileVisitor(private val importedPackages: MutableMap<String,
else {
val requirement = parsed.first()
val name = requirement.name.toLowerCase()
// processing only requirements that are used in the project, discarding others
if (name in importedPackages) {
val pkg = importedPackages.remove(name)!!
if (requirement.isEditable || vcsPrefixes.any { line.startsWith(it) }) {
// keeping editable and vcs requirements
outputLines.addAll(lines)
}
else if (settings.keepMatchingSpecifier && compatibleVersion(requirement, pkg.version, settings.specifyVersion)) {
// existing version separators match the current package version, keeping them
outputLines.addAll(lines)
}
else {
outputLines.add(convertToRequirementsEntry(requirement, settings, pkg.version))
}
val formatted = formatRequirement(requirement, pkg, lines)
outputLines.addAll(formatted)
}
else if (!settings.removeUnused) {
outputLines.addAll(lines)
@@ -102,6 +92,24 @@ class PyRequirementsFileVisitor(private val importedPackages: MutableMap<String,
collectedOutput[requirementsFile.virtualFile] = outputLines
}
private fun formatRequirement(requirement: PyRequirement, pkg: PyPackage, lines: List<String>): List<String> = when {
// keeping editable and vcs requirements
requirement.isEditable || vcsPrefixes.any { lines.first().startsWith(it) } -> lines
// existing version separators match the current package version
settings.keepMatchingSpecifier && compatibleVersion(requirement, pkg.version, settings.specifyVersion) -> lines
// requirement does not match package version and settings
else -> listOf(convertToRequirementsEntry(requirement, settings, pkg.version))
}
private fun visitBaseFile(filename: String, directory: PsiDirectory, visitedFiles: MutableSet<VirtualFile>) {
val referencedFile = directory.virtualFile.findFileByRelativePath(filename)
if (referencedFile != null && visitedFiles.add(referencedFile)) {
val baseRequirementsFile = directory.manager.findFile(referencedFile)!!
doVisitFile(baseRequirementsFile, visitedFiles)
visitedFiles.remove(referencedFile)
}
}
private fun splitByRequirementsEntries(requirementsText: String): MutableList<List<String>> {
var splitList = mutableListOf<String>()
val resultList = mutableListOf<List<String>>()
@@ -1,6 +1,7 @@
// Copyright 2000-2020 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.packaging
import com.intellij.notification.NotificationAction
import com.intellij.notification.NotificationGroup
import com.intellij.notification.NotificationType
import com.intellij.openapi.application.ApplicationManager
@@ -8,11 +9,13 @@ import com.intellij.openapi.application.ReadAction
import com.intellij.openapi.command.WriteCommandAction
import com.intellij.openapi.fileChooser.FileChooserDescriptor
import com.intellij.openapi.fileEditor.FileDocumentManager
import com.intellij.openapi.fileTypes.FileTypeRegistry
import com.intellij.openapi.module.Module
import com.intellij.openapi.progress.ProgressIndicator
import com.intellij.openapi.progress.Task
import com.intellij.openapi.project.Project
import com.intellij.openapi.project.rootManager
import com.intellij.openapi.projectRoots.Sdk
import com.intellij.openapi.util.Ref
import com.intellij.openapi.util.text.StringUtil
import com.intellij.openapi.vfs.LocalFileSystem
@@ -22,36 +25,72 @@ import com.intellij.ui.components.dialog
import com.intellij.ui.layout.*
import com.jetbrains.python.PyBundle
import com.jetbrains.python.PyPsiPackageUtil
import com.jetbrains.python.PythonFileType
import com.jetbrains.python.psi.PyFile
import com.jetbrains.python.sdk.PySdkPopupFactory
import com.jetbrains.python.sdk.PythonSdkUtil
import java.nio.file.Paths
import javax.swing.DefaultComboBoxModel
fun syncWithImports(module: Module) {
/**
* Holder class for generated requirements.
* @param currentFileOutput content of existing requirements file, if any, with added missing entries
* @param baseFilesOutput content of base files, referenced from the original file
* @param unhandledLines lines that we failed to analyze
* @param unchangedInBaseFiles packages with different versions to notify the user about if modification of base files is not allowed
*/
data class PyRequirementsAnalysisResult(val currentFileOutput: List<String>,
val baseFilesOutput: Map<VirtualFile, List<String>>,
val unhandledLines: List<String>,
val unchangedInBaseFiles: List<String>) {
companion object {
fun empty() = PyRequirementsAnalysisResult(mutableListOf(), mutableMapOf(), mutableListOf(), mutableListOf())
}
}
private class PyCollectImportsTask(private val module: Module,
private val psiManager: PsiManager,
title: String) : Task.WithResult<Set<String>, Exception>(module.project, title, true) {
override fun compute(indicator: ProgressIndicator): Set<String> {
val imported = mutableSetOf<String>()
ReadAction.run<Throwable> {
module.rootManager.fileIndex.iterateContent {
indicator.checkCanceled()
if (PythonFileType.INSTANCE == FileTypeRegistry.getInstance().getFileTypeByFileName(it.name)) {
addImports(psiManager.findFile(it) as PyFile, imported)
}
return@iterateContent true
}
}
return imported
}
}
internal fun syncWithImports(module: Module) {
val notificationGroup = NotificationGroup.balloonGroup(PyBundle.message("python.requirements.balloon"))
val sdk = PythonSdkUtil.findPythonSdk(module)
if (sdk == null) {
showNotification(notificationGroup, NotificationType.ERROR, PyBundle.message("python.requirements.error.no.interpreter"), module.project)
val configureSdkAction = NotificationAction.createSimpleExpiring(PyBundle.message("configure.python.interpreter")) {
PySdkPopupFactory.createAndShow(module.project, module)
}
showNotification(notificationGroup,
NotificationType.ERROR,
PyBundle.message("python.requirements.error.no.interpreter"),
module.project,
configureSdkAction)
return
}
val settings = PyPackageRequirementsSettings.getInstance(module)
if (!ApplicationManager.getApplication().isUnitTestMode) {
val proceed = showSpecifyRequirementsFileDialog(module.project, settings)
val proceed = showSyncSettingsDialog(module.project, settings)
if (!proceed) return
}
var requirementsFile = PyPackageUtil.findRequirementsTxt(module)
val matchResult = try {
prepareRequirementsText(module, settings)
} catch (e: IllegalStateException) {
notificationGroup
.createNotification(PyBundle.message("python.requirements.balloon"), e.message!!, NotificationType.ERROR)
.notify(module.project)
return
}
if (matchResult == null) return
val matchResult = prepareRequirementsText(module, sdk, settings)
val psiManager = PsiManager.getInstance(module.project)
WriteCommandAction.runWriteCommandAction(module.project, PyBundle.message("python.requirements.action.name"), null, {
@@ -85,39 +124,26 @@ fun syncWithImports(module: Module) {
}
}
private fun showNotification(notificationGroup: NotificationGroup, type: NotificationType, text: String, project: Project) {
notificationGroup
.createNotification(PyBundle.message("python.requirements.balloon"), text, type)
.notify(project)
private fun showNotification(notificationGroup: NotificationGroup,
type: NotificationType,
text: String,
project: Project,
action: NotificationAction? = null) {
val notification = notificationGroup.createNotification(PyBundle.message("python.requirements.balloon"), text, type)
if (action != null) notification.addAction(action)
notification.notify(project)
}
fun prepareRequirementsText(module: Module, settings: PyPackageRequirementsSettings): PyRequirementsAnalysisResult? {
val sdk = PythonSdkUtil.findPythonSdk(module) ?: return PyRequirementsAnalysisResult.empty()
private fun prepareRequirementsText(module: Module, sdk: Sdk, settings: PyPackageRequirementsSettings): PyRequirementsAnalysisResult {
val psiManager = PsiManager.getInstance(module.project)
val imported = mutableSetOf<String>()
var canceled = false
val dialogTitle = PyBundle.message("python.requirements.analyzing.imports.title")
val task = PyCollectImportsTask(module, psiManager, dialogTitle)
task.queue()
object : Task.Modal(module.project, PyBundle.message("python.requirements.analyzing.imports.title"), true) {
override fun run(indicator: ProgressIndicator) {
ReadAction.run<Throwable> {
module.rootManager.fileIndex.iterateContent {
indicator.checkCanceled()
if (!it.isDirectory && it.extension == "py") {
addImports(psiManager.findFile(it) as PyFile, imported)
}
return@iterateContent true
}
}
}
override fun onCancel() {
canceled = true
}
}.queue()
if (canceled) return null
val installedPackages = PyPackageManager.getInstance(sdk).packages ?: error(PyBundle.message("python.requirements.error.no.packages"))
val importedPackages = imported.asSequence()
val installedPackages = PyPackageManager.getInstance(sdk).refreshAndGetPackages(false)
val importedPackages = task.result.asSequence()
.flatMap { topLevelPackage ->
val aliases = PyPsiPackageUtil.PACKAGES_TOPLEVEL[topLevelPackage]?.toTypedArray() ?: emptyArray()
sequenceOf(topLevelPackage, *aliases)
@@ -126,21 +152,12 @@ fun prepareRequirementsText(module: Module, settings: PyPackageRequirementsSetti
.map { it.name.toLowerCase() to it }
.toMap(mutableMapOf())
val requirementsFile = PyPackageUtil.findRequirementsTxt(module)
val analysisResult = when {
requirementsFile != null -> PyRequirementsFileVisitor(importedPackages, settings).visitRequirementsFile(psiManager.findFile(requirementsFile)!!)
else -> PyRequirementsAnalysisResult.empty()
}
importedPackages.values.asSequence()
.map { if (settings.specifyVersion) "${it.name}${settings.versionSpecifier.separator}${it.version}" else it.name }
.forEach { analysisResult.currentFileOutput.add(it) }
return analysisResult
val requirementsFile = PyPackageUtil.findRequirementsTxt(module) ?: return PyRequirementsAnalysisResult.empty()
val visitor = PyRequirementsFileVisitor(importedPackages, settings)
return visitor.visitRequirementsFile(psiManager.findFile(requirementsFile)!!)
}
private fun showSpecifyRequirementsFileDialog(project: Project, settings: PyPackageRequirementsSettings): Boolean {
private fun showSyncSettingsDialog(project: Project, settings: PyPackageRequirementsSettings): Boolean {
val ref = Ref.create(false)
val descriptor = FileChooserDescriptor(true, false, false, false, false, false)
val panel = panel {
@@ -153,8 +170,8 @@ private fun showSpecifyRequirementsFileDialog(project: Project, settings: PyPack
row {
label(PyBundle.message("python.requirements.version.label"))
comboBox(DefaultComboBoxModel(PyRequirementsVersionSpecifierType.values()),
{ settings.versionSpecifier },
{ settings.versionSpecifier = it }).constraints(growX)
{ settings.versionSpecifier },
{ settings.versionSpecifier = it }).constraints(growX)
}
row {
checkBox(PyBundle.message("python.requirements.remove.unused"),
@@ -183,18 +200,8 @@ private fun showSpecifyRequirementsFileDialog(project: Project, settings: PyPack
return ref.get()
}
private fun addImports(file: PyFile, imported: MutableSet<String>) {
(file.importTargets.asSequence().mapNotNull { it.importedQName?.firstComponent } +
file.fromImports.asSequence().mapNotNull { it.importSourceQName?.firstComponent })
.forEach { imported.add(it) }
}
data class PyRequirementsAnalysisResult(val currentFileOutput: MutableList<String>,
val baseFilesOutput: MutableMap<VirtualFile, MutableList<String>>,
val unhandledLines: MutableList<String>,
val unchangedInBaseFiles: MutableList<String>) {
companion object {
fun empty() = PyRequirementsAnalysisResult(mutableListOf(), mutableMapOf(), mutableListOf(), mutableListOf())
}
}
@@ -8,7 +8,7 @@ import com.intellij.openapi.actionSystem.LangDataKeys
class PySyncPythonRequirementsAction : AnAction() {
override fun actionPerformed(e: AnActionEvent) {
val module = LangDataKeys.MODULE.getData(e.dataContext) ?: return
val module = e.getData(LangDataKeys.MODULE) ?: return
syncWithImports(module)
}
}
@@ -0,0 +1,4 @@
Django==1.1.1
numpy>=1.18.1
requests==3.3.3
tox
@@ -0,0 +1,3 @@
import Django
import pandas
import cookiecutter
@@ -0,0 +1 @@
Django==3.0.0
@@ -0,0 +1,3 @@
-r base_requirements.txt
pandas==1.0.1
cookiecutter==1.7.0
@@ -0,0 +1,4 @@
-r base_requirements.txt
pandas
cookiecutter
Jinja2
@@ -0,0 +1,3 @@
# versions do not match, but will not be changed
Django==1.1.1
requests==3.3.3
@@ -0,0 +1,4 @@
import requests
import Django
import pandas
import cookiecutter
@@ -0,0 +1,3 @@
# versions do not match, but will not be changed
Django==1.1.1
requests==3.3.3
@@ -0,0 +1,3 @@
-r base_requirements.txt
pandas==1.0.1
cookiecutter==1.7.0
@@ -0,0 +1,3 @@
-r base_requirements.txt
pandas
cookiecutter
@@ -0,0 +1,2 @@
Django==1.1.1
requests==3.3.3
@@ -0,0 +1,4 @@
import requests
import Django
import pandas
import cookiecutter
@@ -0,0 +1,2 @@
Django==3.0.0
requests==2.22.0
@@ -0,0 +1,3 @@
-r base_requirements.txt
pandas==1.0.1
cookiecutter==1.7.0
@@ -0,0 +1,3 @@
-r base_requirements.txt
pandas
cookiecutter
@@ -0,0 +1 @@
import docker
@@ -0,0 +1 @@
docker-py==1.10.6
@@ -0,0 +1,2 @@
docker==3.7.0
docker-py==1.10.6
@@ -1,6 +1,11 @@
// Copyright 2000-2020 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.packaging
import com.intellij.ide.DataManager
import com.intellij.openapi.actionSystem.ActionManager
import com.intellij.openapi.actionSystem.AnActionEvent
import com.intellij.openapi.actionSystem.LangDataKeys
import com.intellij.openapi.actionSystem.impl.SimpleDataContext
import com.intellij.openapi.application.ApplicationManager
import com.intellij.openapi.project.Project
import com.intellij.openapi.projectRoots.Sdk
@@ -11,6 +16,14 @@ import org.easymock.EasyMock
class PyRequirementsGenerationTest : PyTestCase() {
private var oldPackageManagers: PyPackageManagers? = null
private val installedPackages = mapOf("Django" to "3.0.0",
"requests" to "2.22.0",
"Jinja2" to "2.11.1",
"pandas" to "1.0.1" ,
"cookiecutter" to "1.7.0",
"numpy" to "1.18.1",
"tox" to "3.14.4",
"docker-py" to "1.10.6")
fun testNewFileGeneration() = doTest()
fun testNewFileWithoutVersion() = doTest(PyRequirementsVersionSpecifierType.NO_VERSION)
@@ -27,11 +40,16 @@ class PyRequirementsGenerationTest : PyTestCase() {
fun testRemoveUnused() = doTest(removeUnused = true)
fun testUpdateVersionKeepInstallOptions() = doTest()
fun testCompatibleFileReference() = doTest()
fun testDifferentTopLevelImport() = doTest()
fun testDifferentTopLevelImportWithOriginalPackage() = doTest(packages = installedPackages + mapOf("docker" to "3.7.0"))
fun testBaseFileUnchanged() = doTest()
fun testBaseFileUpdate() = doTest(modifyBaseFiles = true)
fun testBaseFileCleanup() = doTest(modifyBaseFiles = true, removeUnused = true)
private fun doTest(versionSpecifier: PyRequirementsVersionSpecifierType = PyRequirementsVersionSpecifierType.STRONG_EQ,
removeUnused: Boolean = false,
modifyBaseFiles: Boolean = false,
block: ((PyRequirementsAnalysisResult) -> Unit)? = null) {
packages: Map<String, String> = installedPackages) {
val settings = PyPackageRequirementsSettings.getInstance(myFixture.module)
val oldRequirementsPath = settings.requirementsPath
val oldVersionSpecifier = settings.versionSpecifier
@@ -39,15 +57,7 @@ class PyRequirementsGenerationTest : PyTestCase() {
val oldModifyBaseFiles = settings.modifyBaseFiles
try {
overrideInstalledPackages(
"Django" to "3.0.0",
"requests" to "2.22.0",
"Jinja2" to "2.11.1",
"pandas" to "1.0.1" ,
"cookiecutter" to "1.7.0",
"numpy" to "1.18.1",
"tox" to "3.14.4"
)
overrideInstalledPackages(packages)
settings.requirementsPath = "requirements.txt"
settings.versionSpecifier = versionSpecifier
@@ -57,12 +67,15 @@ class PyRequirementsGenerationTest : PyTestCase() {
val testName = getTestName(true)
myFixture.copyDirectoryToProject(testName, "")
myFixture.configureFromTempProjectFile(settings.requirementsPath)
if (block != null) {
block(prepareRequirementsText(myFixture.module, settings)!!)
}
else {
syncWithImports(myFixture.module)
myFixture.checkResultByFile("$testName/new_${settings.requirementsPath}")
val action = ActionManager.getInstance().getAction("PySyncPythonRequirements")
val parentContext = DataManager.getInstance().getDataContext(myFixture.editor.component)
val context = SimpleDataContext.getSimpleContext(LangDataKeys.MODULE.name, myFixture.module, parentContext)
val event = AnActionEvent.createFromAnAction(action, null, "", context)
action.actionPerformed(event)
myFixture.checkResultByFile("$testName/new_${settings.requirementsPath}", true)
if (modifyBaseFiles) {
myFixture.checkResultByFile("base_${settings.requirementsPath}", "$testName/new_base_${settings.requirementsPath}", true)
}
assertProjectFilesNotParsed(myFixture.file)
}
@@ -93,15 +106,15 @@ class PyRequirementsGenerationTest : PyTestCase() {
override fun getTestDataPath(): String = super.getTestDataPath() + "/requirement/generation"
private fun overrideInstalledPackages(vararg packages: Pair<String, String>) {
ApplicationManager.getApplication().registerServiceInstance(PyPackageManagers::class.java, MockPyPackageManagers(packages.toMap()))
private fun overrideInstalledPackages(packages: Map<String, String>) {
ApplicationManager.getApplication().registerServiceInstance(PyPackageManagers::class.java, MockPyPackageManagers(packages))
}
private class MockPyPackageManagers(val packages: Map<String, String>) : PyPackageManagers() {
override fun forSdk(sdk: Sdk): PyPackageManager {
val packageManager = EasyMock.createMock<PyPackageManager>(PyPackageManager::class.java)
EasyMock
.expect(packageManager.packages)
.expect(packageManager.refreshAndGetPackages(false))
.andReturn(packages.map { PyPackage(it.key, it.value, null, emptyList()) })
EasyMock.replay(packageManager)