PY-18792 Add several new options for Optimize Imports in Python

Namely allow to:
* disable alphabetical ordering of imports
* order individual imported names inside "from" import (PY-19292)
* combine multiple "from" imports with the same source (PY-14176)
This commit is contained in:
Mikhail Golubev
2016-06-15 19:34:11 +03:00
parent 26f2e06037
commit 84381c7273
10 changed files with 187 additions and 16 deletions
@@ -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.<String>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<PyImportElement> IMPORT_ELEMENT_COMPARATOR = (o1, o2) -> Comparing.compare(o1.getImportedQName(),
o2.getImportedQName());
private final PyFile myFile;
private final List<PyImportStatementBase> myImportBlock;
private final Map<ImportPriority, List<PyImportStatementBase>> 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<PyImportStatementBase> original = myGroups.get(priority);
final List<PyImportStatementBase> transformed = transformImportStatements(original);
hasTransformedImports |= !original.equals(transformed);
myGroups.put(priority, transformed);
}
if (hasTransformedImports || needBlankLinesBetweenGroups() || groupsNotSorted()) {
applyResults();
}
}
@NotNull
private List<PyImportStatementBase> transformImportStatements(@NotNull List<PyImportStatementBase> imports) {
final List<PyImportStatementBase> result = new ArrayList<>();
final PyElementGenerator generator = PyElementGenerator.getInstance(myFile.getProject());
final LanguageLevel langLevel = LanguageLevel.forElement(myFile);
final MultiMap<QualifiedName, PyFromImportStatement> 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<PyFromImportStatement> sameSourceImports = fromImportSources.get(source);
if (!sameSourceImports.isEmpty()) {
final List<PyImportElement> 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<PyImportStatementBase> importOrdering = Ordering.from(AddImportHelper.IMPORT_TYPE_THEN_NAME_COMPARATOR);
return ContainerUtil.exists(myGroups.values(), imports -> !importOrdering.isOrdered(imports));
}
@@ -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)
@@ -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)
@@ -0,0 +1,3 @@
from module import B as Z, A, C
print(A, C, Z)
@@ -0,0 +1,5 @@
from module import B as Z
from module import A
from module import C
print(A, C, Z)
@@ -0,0 +1,3 @@
from module import A, B as Z, C
print(A, C, Z)
@@ -0,0 +1,5 @@
from module import B as Z
from module import A
from module import C
print(A, C, Z)
@@ -0,0 +1,3 @@
from module import A as Z, C, a, b, c
print(C, Z, a, b, c)
@@ -0,0 +1,3 @@
from module import C, A as Z, a, c, b
print(C, Z, a, b, c)
@@ -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");