mirror of
https://gitflic.ru/project/openide/openide.git
synced 2026-09-27 10:03:11 +07:00
more tests and improvements for python requirements generation (PY-18847)
GitOrigin-RevId: 2176648fd9021b3734e29c72b9daaa98e81a108a
This commit is contained in:
committed by
intellij-monorepo-bot
parent
539dab2069
commit
6cb5115bed
@@ -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
|
||||
+1
@@ -0,0 +1 @@
|
||||
import docker
|
||||
+2
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user