diff --git a/python/src/com/jetbrains/python/codeInsight/imports/PyImportOptimizer.java b/python/src/com/jetbrains/python/codeInsight/imports/PyImportOptimizer.java index 314a1dd23567..4daf1172959d 100644 --- a/python/src/com/jetbrains/python/codeInsight/imports/PyImportOptimizer.java +++ b/python/src/com/jetbrains/python/codeInsight/imports/PyImportOptimizer.java @@ -18,10 +18,14 @@ package com.jetbrains.python.codeInsight.imports; import com.google.common.collect.Ordering; import com.intellij.codeInspection.LocalInspectionToolSession; import com.intellij.lang.ImportOptimizer; +import com.intellij.openapi.util.Comparing; +import com.intellij.openapi.util.text.StringUtil; import com.intellij.psi.PsiElement; import com.intellij.psi.PsiFile; import com.intellij.psi.codeStyle.CodeStyleSettingsManager; +import com.intellij.psi.util.QualifiedName; import com.intellij.util.containers.ContainerUtil; +import com.intellij.util.containers.MultiMap; import com.jetbrains.python.codeInsight.imports.AddImportHelper.ImportPriority; import com.jetbrains.python.formatter.PyBlock; import com.jetbrains.python.formatter.PyCodeStyleSettings; @@ -31,6 +35,8 @@ import org.jetbrains.annotations.NotNull; import java.util.*; +import static com.jetbrains.python.psi.PyUtil.as; + /** * @author yole */ @@ -47,7 +53,7 @@ public class PyImportOptimizer implements ImportOptimizer { final LocalInspectionToolSession session = new LocalInspectionToolSession(file, 0, file.getTextLength()); final PyUnresolvedReferencesInspection.Visitor visitor = new PyUnresolvedReferencesInspection.Visitor(null, session, - Collections.emptyList()); + Collections.emptyList()); file.accept(new PyRecursiveElementVisitor() { @Override public void visitElement(PsiElement node) { @@ -65,6 +71,9 @@ public class PyImportOptimizer implements ImportOptimizer { private static class ImportSorter { + private static final Comparator IMPORT_ELEMENT_COMPARATOR = (o1, o2) -> Comparing.compare(o1.getImportedQName(), + o2.getImportedQName()); + private final PyFile myFile; private final List myImportBlock; private final Map> myGroups; @@ -84,29 +93,97 @@ public class PyImportOptimizer implements ImportOptimizer { if (myImportBlock.isEmpty()) { return; } - boolean hasSplittedImports = false; - final LanguageLevel langLevel = LanguageLevel.forElement(myFile); - final PyElementGenerator generator = PyElementGenerator.getInstance(myFile.getProject()); + for (PyImportStatementBase importStatement : myImportBlock) { final ImportPriority priority = AddImportHelper.getImportPriority(importStatement); - if (importStatement instanceof PyImportStatement && importStatement.getImportElements().length > 1) { - for (PyImportElement importElement : importStatement.getImportElements()) { - hasSplittedImports = true; - // getText() for ImportElement includes alias - final PyImportStatement splitImport = generator.createImportStatement(langLevel, importElement.getText(), null); - myGroups.get(priority).add(splitImport); - } - } - else { - myGroups.get(priority).add(importStatement); - } + myGroups.get(priority).add(importStatement); } - if (hasSplittedImports || needBlankLinesBetweenGroups() || groupsNotSorted()) { + + boolean hasTransformedImports = false; + for (ImportPriority priority : ImportPriority.values()) { + final List original = myGroups.get(priority); + final List transformed = transformImportStatements(original); + hasTransformedImports |= !original.equals(transformed); + myGroups.put(priority, transformed); + } + + if (hasTransformedImports || needBlankLinesBetweenGroups() || groupsNotSorted()) { applyResults(); } } + + @NotNull + private List transformImportStatements(@NotNull List imports) { + final List result = new ArrayList<>(); + + final PyElementGenerator generator = PyElementGenerator.getInstance(myFile.getProject()); + final LanguageLevel langLevel = LanguageLevel.forElement(myFile); + + final MultiMap fromImportSources = MultiMap.create(); + for (PyImportStatementBase statement : imports) { + final PyFromImportStatement fromImport = as(statement, PyFromImportStatement.class); + if (fromImport != null) { + fromImportSources.putValue(fromImport.getImportSourceQName(), fromImport); + } + } + + for (PyImportStatementBase statement : imports) { + if (statement instanceof PyImportStatement) { + final PyImportStatement importStatement = (PyImportStatement)statement; + final PyImportElement[] importElements = importStatement.getImportElements(); + // Split combined imports like "import foo, bar as b" + if (importElements.length > 1) { + for (PyImportElement importElement : importElements) { + // getText() for ImportElement includes alias + final PyImportStatement splitted = generator.createImportStatement(langLevel, importElement.getText(), null); + result.add(splitted); + } + } + else { + result.add(importStatement); + } + } + else if (statement instanceof PyFromImportStatement) { + final PyFromImportStatement fromImportStatement = (PyFromImportStatement)statement; + final QualifiedName source = fromImportStatement.getImportSourceQName(); + final String sourceText = Objects.toString(source, ""); + if (myPySettings.OPTIMIZE_IMPORTS_JOIN_FROM_IMPORTS_WITH_SAME_SOURCE) { + final Collection sameSourceImports = fromImportSources.get(source); + if (!sameSourceImports.isEmpty()) { + final List allImportElements = new ArrayList<>(); + for (PyFromImportStatement sameSourceImport : sameSourceImports) { + ContainerUtil.addAll(allImportElements, sameSourceImport.getImportElements()); + } + if (myPySettings.OPTIMIZE_IMPORTS_SORT_NAMES_IN_FROM_IMPORTS) { + Collections.sort(allImportElements, IMPORT_ELEMENT_COMPARATOR); + } + final String importedNames = StringUtil.join(allImportElements, PsiElement::getText, ", "); + result.add(generator.createFromImportStatement(langLevel, sourceText, importedNames, null)); + + // remember that we have checked imports from this source already + fromImportSources.remove(source); + } + } + else if (myPySettings.OPTIMIZE_IMPORTS_SORT_NAMES_IN_FROM_IMPORTS) { + final PyImportElement[] importElements = fromImportStatement.getImportElements(); + Arrays.sort(importElements, IMPORT_ELEMENT_COMPARATOR); + final String importedNames = StringUtil.join(importElements, PsiElement::getText, ", "); + result.add(generator.createFromImportStatement(langLevel, sourceText, importedNames, null)); + } + else { + result.add(fromImportStatement); + } + } + } + + + return result; + } private boolean groupsNotSorted() { + if (!myPySettings.OPTIMIZE_IMPORTS_SORT_ALPHABETICALLY) { + return false; + } final Ordering importOrdering = Ordering.from(AddImportHelper.IMPORT_TYPE_THEN_NAME_COMPARATOR); return ContainerUtil.exists(myGroups.values(), imports -> !importOrdering.isOrdered(imports)); } diff --git a/python/testData/optimizeImports/disableAlphabeticalOrder.after.py b/python/testData/optimizeImports/disableAlphabeticalOrder.after.py new file mode 100644 index 000000000000..46e2c5d80bbe --- /dev/null +++ b/python/testData/optimizeImports/disableAlphabeticalOrder.after.py @@ -0,0 +1,24 @@ +from __future__ import unicode_literals +from __future__ import absolute_import + +import sys +from datetime import timedelta + +import z +import b +import a +from a import C1 +from alphabet import D +from b import func +from +import foo # broken +from . import m1 +import # broken +from alphabet import * +from .. import m2 +from alphabet import C +from alphabet import B, A +from .pkg import m3 +from . import m4, m5 + +print(z, b, a, C1, func, sys, abc, foo, timedelta, A, B, C, D, m1, m2, m3, m4, m5) \ No newline at end of file diff --git a/python/testData/optimizeImports/disableAlphabeticalOrder.py b/python/testData/optimizeImports/disableAlphabeticalOrder.py new file mode 100644 index 000000000000..cb560cc4f45b --- /dev/null +++ b/python/testData/optimizeImports/disableAlphabeticalOrder.py @@ -0,0 +1,23 @@ +from __future__ import unicode_literals +from __future__ import absolute_import + +import z +import b +import a +from a import C1 +from alphabet import D +from alphabet import A +from b import func +from import foo # broken +import sys +from . import m1 +from datetime import timedelta +import # broken +from alphabet import * +from .. import m2 +from alphabet import C +from alphabet import B, A +from .pkg import m3 +from . import m4, m5 + +print(z, b, a, C1, func, sys, abc, foo, timedelta, A, B, C, D, m1, m2, m3, m4, m5) \ No newline at end of file diff --git a/python/testData/optimizeImports/joinFromImportsForSameSource.after.py b/python/testData/optimizeImports/joinFromImportsForSameSource.after.py new file mode 100644 index 000000000000..633b8b9eab4e --- /dev/null +++ b/python/testData/optimizeImports/joinFromImportsForSameSource.after.py @@ -0,0 +1,3 @@ +from module import B as Z, A, C + +print(A, C, Z) diff --git a/python/testData/optimizeImports/joinFromImportsForSameSource.py b/python/testData/optimizeImports/joinFromImportsForSameSource.py new file mode 100644 index 000000000000..872b22b7ce71 --- /dev/null +++ b/python/testData/optimizeImports/joinFromImportsForSameSource.py @@ -0,0 +1,5 @@ +from module import B as Z +from module import A +from module import C + +print(A, C, Z) diff --git a/python/testData/optimizeImports/joinFromImportsForSameSourceAndSortNames.after.py b/python/testData/optimizeImports/joinFromImportsForSameSourceAndSortNames.after.py new file mode 100644 index 000000000000..10bc3c51d4ca --- /dev/null +++ b/python/testData/optimizeImports/joinFromImportsForSameSourceAndSortNames.after.py @@ -0,0 +1,3 @@ +from module import A, B as Z, C + +print(A, C, Z) diff --git a/python/testData/optimizeImports/joinFromImportsForSameSourceAndSortNames.py b/python/testData/optimizeImports/joinFromImportsForSameSourceAndSortNames.py new file mode 100644 index 000000000000..872b22b7ce71 --- /dev/null +++ b/python/testData/optimizeImports/joinFromImportsForSameSourceAndSortNames.py @@ -0,0 +1,5 @@ +from module import B as Z +from module import A +from module import C + +print(A, C, Z) diff --git a/python/testData/optimizeImports/orderNamesInsideFromImport.after.py b/python/testData/optimizeImports/orderNamesInsideFromImport.after.py new file mode 100644 index 000000000000..48dd9c280e37 --- /dev/null +++ b/python/testData/optimizeImports/orderNamesInsideFromImport.after.py @@ -0,0 +1,3 @@ +from module import A as Z, C, a, b, c + +print(C, Z, a, b, c) diff --git a/python/testData/optimizeImports/orderNamesInsideFromImport.py b/python/testData/optimizeImports/orderNamesInsideFromImport.py new file mode 100644 index 000000000000..bb963b45d06c --- /dev/null +++ b/python/testData/optimizeImports/orderNamesInsideFromImport.py @@ -0,0 +1,3 @@ +from module import C, A as Z, a, c, b + +print(C, Z, a, b, c) diff --git a/python/testSrc/com/jetbrains/python/PyOptimizeImportsTest.java b/python/testSrc/com/jetbrains/python/PyOptimizeImportsTest.java index c812c354c016..a97a33e08da9 100644 --- a/python/testSrc/com/jetbrains/python/PyOptimizeImportsTest.java +++ b/python/testSrc/com/jetbrains/python/PyOptimizeImportsTest.java @@ -132,6 +132,31 @@ public class PyOptimizeImportsTest extends PyTestCase { }); } } + + // PY-18792 + public void testDisableAlphabeticalOrder() { + getPythonCodeStyleSettings().OPTIMIZE_IMPORTS_SORT_ALPHABETICALLY = false; + doTest(); + } + + // PY-18792, PY-19292 + public void testOrderNamesInsideFromImport() { + getPythonCodeStyleSettings().OPTIMIZE_IMPORTS_SORT_NAMES_IN_FROM_IMPORTS = true; + doTest(); + } + + // PY-18792, PY-12926 + public void testJoinFromImportsForSameSource() { + getPythonCodeStyleSettings().OPTIMIZE_IMPORTS_JOIN_FROM_IMPORTS_WITH_SAME_SOURCE = true; + doTest(); + } + + // PY-18792, PY-12926 + public void testJoinFromImportsForSameSourceAndSortNames() { + getPythonCodeStyleSettings().OPTIMIZE_IMPORTS_JOIN_FROM_IMPORTS_WITH_SAME_SOURCE = true; + getPythonCodeStyleSettings().OPTIMIZE_IMPORTS_SORT_NAMES_IN_FROM_IMPORTS = true; + doTest(); + } private void doTest() { myFixture.configureByFile(getTestName(true) + ".py");