diff --git a/python/src/com/jetbrains/python/codeInsight/imports/PyImportOptimizer.java b/python/src/com/jetbrains/python/codeInsight/imports/PyImportOptimizer.java index 8c2aa27bfaf1..467a3591275c 100644 --- a/python/src/com/jetbrains/python/codeInsight/imports/PyImportOptimizer.java +++ b/python/src/com/jetbrains/python/codeInsight/imports/PyImportOptimizer.java @@ -18,19 +18,25 @@ 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.project.Project; import com.intellij.openapi.util.Comparing; import com.intellij.openapi.util.text.StringUtil; +import com.intellij.psi.PsiComment; import com.intellij.psi.PsiElement; import com.intellij.psi.PsiFile; +import com.intellij.psi.PsiWhiteSpace; import com.intellij.psi.codeStyle.CodeStyleManager; import com.intellij.psi.codeStyle.CodeStyleSettingsManager; +import com.intellij.psi.util.PsiTreeUtil; +import com.intellij.util.ObjectUtils; 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; import com.jetbrains.python.inspections.unresolvedReference.PyUnresolvedReferencesInspection; import com.jetbrains.python.psi.*; +import com.jetbrains.python.psi.impl.PyPsiUtils; +import one.util.streamex.StreamEx; import org.jetbrains.annotations.NotNull; import java.util.*; @@ -70,7 +76,6 @@ public class PyImportOptimizer implements ImportOptimizer { } private static class ImportSorter { - private static final Comparator IMPORT_ELEMENT_COMPARATOR = (o1, o2) -> { final int byImportedName = Comparing.compare(o1.getImportedQName(), o2.getImportedQName()); if (byImportedName != 0) { @@ -80,14 +85,16 @@ public class PyImportOptimizer implements ImportOptimizer { }; private final PyFile myFile; + private final PyCodeStyleSettings myPySettings; private final List myImportBlock; private final Map> myGroups; - private final PyCodeStyleSettings myPySettings; + private final MultiMap myImportToComments; private ImportSorter(@NotNull PyFile file) { myFile = file; myPySettings = CodeStyleSettingsManager.getSettings(myFile.getProject()).getCustomSettings(PyCodeStyleSettings.class); myImportBlock = myFile.getImportBlock(); + myImportToComments = MultiMap.create(); myGroups = new EnumMap<>(ImportPriority.class); for (ImportPriority priority : ImportPriority.values()) { myGroups.put(priority, new ArrayList<>()); @@ -120,34 +127,43 @@ public class PyImportOptimizer implements ImportOptimizer { @NotNull private List transformImportStatements(@NotNull List imports) { final List result = new ArrayList<>(); - - final PyElementGenerator generator = PyElementGenerator.getInstance(myFile.getProject()); + + final Project project = myFile.getProject(); + final PyElementGenerator generator = PyElementGenerator.getInstance(project); final LanguageLevel langLevel = LanguageLevel.forElement(myFile); + // Used to combine "from" imports with the same sources final MultiMap fromImportSources = MultiMap.create(); + // Preserve line comments if any + final MultiMap precedingComments = MultiMap.create(); + for (PyImportStatementBase statement : imports) { final PyFromImportStatement fromImport = as(statement, PyFromImportStatement.class); - if (fromImport != null) { - if (fromImport.isStarImport()) { - continue; - } + if (fromImport != null && !fromImport.isStarImport()) { fromImportSources.putValue(getNormalizedFromImportSource(fromImport), fromImport); } + precedingComments.putValues(statement, collectPrecedingLineComments(statement)); } - + 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); + final List newImports = ContainerUtil.map(importElements, e -> generator.createImportStatement(langLevel, e.getText(), null)); + final PyImportStatement topmostImport; + if (myPySettings.OPTIMIZE_IMPORTS_SORT_IMPORTS) { + topmostImport = Collections.min(newImports, AddImportHelper.getSameGroupImportsComparator(project)); } + else { + topmostImport = newImports.get(0); + } + myImportToComments.putValues(topmostImport, precedingComments.get(statement)); + result.addAll(newImports); } else { + myImportToComments.putValues(statement, precedingComments.get(statement)); result.add(importStatement); } } @@ -156,7 +172,7 @@ public class PyImportOptimizer implements ImportOptimizer { final String source = getNormalizedFromImportSource(fromImport); final List newStatementElements = new ArrayList<>(); - // We cannot neither sort, not combine star imports + // We can neither sort, nor combine star imports if (!fromImport.isStarImport()) { final Collection sameSourceImports = fromImportSources.get(source); if (sameSourceImports.isEmpty()) { @@ -184,18 +200,43 @@ public class PyImportOptimizer implements ImportOptimizer { Collections.sort(newStatementElements, IMPORT_ELEMENT_COMPARATOR); } final String importedNames = StringUtil.join(newStatementElements, PsiElement::getText, ", "); - result.add(generator.createFromImportStatement(langLevel, source, importedNames, null)); + final PyFromImportStatement combinedImport = generator.createFromImportStatement(langLevel, source, importedNames, null); + ContainerUtil.map2LinkedSet(newStatementElements, e -> (PyImportStatementBase)e.getParent()).forEach(affected -> { + myImportToComments.putValues(combinedImport, precedingComments.get(affected)); + }); + result.add(combinedImport); } else { + myImportToComments.putValues(fromImport, precedingComments.get(fromImport)); result.add(fromImport); } } } - - return result; } + @NotNull + private static List collectPrecedingLineComments(@NotNull PyImportStatementBase statement) { + final List result = new ArrayList<>(); + PsiElement prev = PyPsiUtils.getPrevNonWhitespaceSibling(statement); + while ((prev instanceof PsiComment) && onItsOwnLine(prev) && !isShebangComment(((PsiComment)prev))) { + result.add((PsiComment)prev); + prev = PyPsiUtils.getPrevNonWhitespaceSibling(prev); + } + Collections.reverse(result); + return result; + } + + private static boolean isShebangComment(@NotNull PsiComment comment) { + return comment.getTextRange().getStartOffset() == 0 && comment.getText().startsWith("#!"); + } + + private static boolean onItsOwnLine(@NotNull PsiElement element) { + if (element.getTextRange().getStartOffset() == 0) return true; + final PsiWhiteSpace sibling = as(PsiTreeUtil.prevLeaf(element), PsiWhiteSpace.class); + return sibling != null && (sibling.textContains('\n') || sibling.getTextRange().getStartOffset() == 0); + } + @NotNull public static String getNormalizedFromImportSource(@NotNull PyFromImportStatement statement) { return StringUtil.repeatSymbol('.', statement.getRelativeLevel()) + Objects.toString(statement.getImportSourceQName(), ""); @@ -210,13 +251,7 @@ public class PyImportOptimizer implements ImportOptimizer { } private boolean needBlankLinesBetweenGroups() { - int nonEmptyGroups = 0; - for (List bases : myGroups.values()) { - if (!bases.isEmpty()) { - nonEmptyGroups++; - } - } - return nonEmptyGroups > 1; + return StreamEx.of(myGroups.values()).remove(List::isEmpty).count() > 1; } private void applyResults() { @@ -227,47 +262,40 @@ public class PyImportOptimizer implements ImportOptimizer { myGroups.put(priority, imports); } } - prepareNewImports(); - markGroupStarts(); - addImports(myImportBlock.get(0)); - - myFile.deleteChildRange(myImportBlock.get(0), myImportBlock.get(myImportBlock.size() - 1)); + final PyImportStatementBase firstImport = myImportBlock.get(0); + final List comments = collectPrecedingLineComments(firstImport); + final PsiElement topmostAnchor = ObjectUtils.notNull(ContainerUtil.getFirstItem(comments), firstImport); + addImportsBefore(topmostAnchor); + myFile.deleteChildRange(topmostAnchor, ContainerUtil.getLastItem(myImportBlock)); } - private void prepareNewImports() { + private void addImportsBefore(@NotNull PsiElement anchor) { + final StringBuilder content = new StringBuilder(); + for (List imports : myGroups.values()) { - for (int i = 0; i < imports.size(); i++) { - final PyImportStatementBase newImport = imports.get(i); - final CodeStyleManager styleManager = CodeStyleManager.getInstance(newImport.getProject()); - // Some of imports were copied as is and they're still present in the original PSI file - final PyImportStatementBase formatted = (PyImportStatementBase)styleManager.reformat(newImport.copy()); - imports.set(i, formatted); + if (content.length() > 0) { + // one extra blank line between import groups according to PEP 8 + content.append("\n"); } - } - } - - private void markGroupStarts() { - for (List group : myGroups.values()) { - boolean firstImportInGroup = true; - for (PyImportStatementBase statement : group) { - if (firstImportInGroup) { - statement.putCopyableUserData(PyBlock.IMPORT_GROUP_BEGIN, true); - firstImportInGroup = false; - } - else { - statement.putCopyableUserData(PyBlock.IMPORT_GROUP_BEGIN, null); + for (PyImportStatementBase statement : imports) { + //StringUtil.join(myImportToComments.get(statement), PsiElement::getText, "\n", content); + for (PsiComment comment : myImportToComments.get(statement)) { + content.append(comment.getText()).append("\n"); } + content.append(statement.getText()).append("\n"); } } - } + + final Project project = anchor.getProject(); + final PyElementGenerator generator = PyElementGenerator.getInstance(project); + PyFile file = (PyFile)generator.createDummyFile(LanguageLevel.forElement(anchor), content.toString()); + file = (PyFile)CodeStyleManager.getInstance(project).reformat(file); + final List newImportBlock = file.getImportBlock(); + assert newImportBlock != null; - private void addImports(@NotNull PyImportStatementBase anchor) { - // EnumMap returns values in key order, i.e. according to import groups priority - for (List imports : myGroups.values()) { - for (PyImportStatementBase newImport : imports) { - myFile.addBefore(newImport, anchor); - } - } + final PyImportStatementBase lastImport = ContainerUtil.getLastItem(newImportBlock); + assert lastImport != null; + myFile.addRangeBefore(file.getFirstChild(), lastImport, anchor); } } } diff --git a/python/testData/optimizeImports/commentsHandling.after.py b/python/testData/optimizeImports/commentsHandling.after.py new file mode 100644 index 000000000000..c6685a8a6eeb --- /dev/null +++ b/python/testData/optimizeImports/commentsHandling.after.py @@ -0,0 +1,15 @@ +#!/usr/bin/python +# comment for a +import a # trailing comment for normal import +# comment for b +import b +# comment for c, d +import c +import d +# comment for name1 and name3 +# comment for name2 +from mod import name1, name2, name3 +# comment for star import +from mod2 import * + +print(a, b, c, d, name1, name2, name3) diff --git a/python/testData/optimizeImports/commentsHandling.py b/python/testData/optimizeImports/commentsHandling.py new file mode 100644 index 000000000000..fd92a238d680 --- /dev/null +++ b/python/testData/optimizeImports/commentsHandling.py @@ -0,0 +1,19 @@ +#!/usr/bin/python + +# comment for b +import b +# comment for a +import a # trailing comment for normal import + +# comment for c, d +import d, c + +# comment for name2 +from mod import name2 +# comment for name1 and name3 +from mod import name1, name3 # trailing comment for "from" import + +# comment for star import +from mod2 import * + +print(a, b, c, d, name1, name2, name3) diff --git a/python/testData/refactoring/move/relativeImportsInsideMovedModule/after/src/subpkg1/mod1.py b/python/testData/refactoring/move/relativeImportsInsideMovedModule/after/src/subpkg1/mod1.py index 2b8de1e2ee84..dfeddbb5c665 100644 --- a/python/testData/refactoring/move/relativeImportsInsideMovedModule/after/src/subpkg1/mod1.py +++ b/python/testData/refactoring/move/relativeImportsInsideMovedModule/after/src/subpkg1/mod1.py @@ -1,7 +1,9 @@ import +# malformed imports from from +# absolute imports import pkg1.subpkg2 as foo from pkg1 import subpkg2 from pkg1 import subpkg2 as bar diff --git a/python/testSrc/com/jetbrains/python/PyOptimizeImportsTest.java b/python/testSrc/com/jetbrains/python/PyOptimizeImportsTest.java index 6124cfce4738..eb9aef2dde9e 100644 --- a/python/testSrc/com/jetbrains/python/PyOptimizeImportsTest.java +++ b/python/testSrc/com/jetbrains/python/PyOptimizeImportsTest.java @@ -242,6 +242,13 @@ public class PyOptimizeImportsTest extends PyTestCase { doTest(); } + // PY-19837 + public void testCommentsHandling() { + 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"); OptimizeImportsAction.actionPerformedImpl(DataManager.getInstance().getDataContext(myFixture.getEditor().getContentComponent()));