diff --git a/java/java-analysis-impl/src/com/intellij/codeInspection/StreamApiMigrationInspection.java b/java/java-analysis-impl/src/com/intellij/codeInspection/StreamApiMigrationInspection.java index be69482d6e6f..93cbb2792411 100644 --- a/java/java-analysis-impl/src/com/intellij/codeInspection/StreamApiMigrationInspection.java +++ b/java/java-analysis-impl/src/com/intellij/codeInspection/StreamApiMigrationInspection.java @@ -44,6 +44,7 @@ import java.util.ArrayList; import java.util.Arrays; import java.util.Collection; import java.util.List; +import java.util.stream.Collectors; /** * User: anna @@ -123,17 +124,22 @@ public class StreamApiMigrationInspection extends BaseJavaBatchLocalInspectionTo PsiBreakStatement.class, PsiReturnStatement.class, PsiThrowStatement.class); if (exitPoints.isEmpty()) { - final List usedVariables = ControlFlowUtil.getUsedVariables(controlFlow, startOffset, endOffset); - for (PsiVariable variable : usedVariables) { - if (!HighlightControlFlowUtil.isEffectivelyFinal(variable, body, null)) { - return; - } - } - if (ExceptionUtil.getThrownCheckedExceptions(new PsiElement[]{body}).isEmpty()) { TerminalBlock tb = TerminalBlock.from(statement.getIterationParameter(), body); List operations = tb.extractOperations(); + final List nonFinalVariables = ControlFlowUtil.getUsedVariables(controlFlow, startOffset, endOffset) + .stream().filter(variable -> !HighlightControlFlowUtil.isEffectivelyFinal(variable, body, null)) + .collect(Collectors.toList()); + + if(getCounter(statement, tb, operations, nonFinalVariables) != null) { + holder.registerProblem(iteratedValue, "Can be replaced with count() call", + ProblemHighlightType.GENERIC_ERROR_OR_WARNING, + new ReplaceWithCountFix()); + } + if(!nonFinalVariables.isEmpty()) { + return; + } if ((isArray || !isRawSubstitution(iteratedValueType, collectionClass)) && isCollectCall(tb, operations)) { boolean addAll = operations.isEmpty() && isAddAllCall(tb); holder.registerProblem(iteratedValue, "Can be replaced with " + (addAll ? "addAll call" : "collect call"), @@ -169,6 +175,43 @@ public class StreamApiMigrationInspection extends BaseJavaBatchLocalInspectionTo }; } + private static PsiExpression extractIncrementedExpression(PsiStatement statement) { + if(!(statement instanceof PsiExpressionStatement)) return null; + PsiExpression expression = ((PsiExpressionStatement)statement).getExpression(); + PsiExpression operand; + if(expression instanceof PsiPostfixExpression) { + if(!JavaTokenType.PLUSPLUS.equals(((PsiPostfixExpression)expression).getOperationTokenType())) return null; + operand = ((PsiPostfixExpression)expression).getOperand(); + } else if(expression instanceof PsiPrefixExpression) { + if(!JavaTokenType.PLUSPLUS.equals(((PsiPrefixExpression)expression).getOperationTokenType())) return null; + operand = ((PsiPrefixExpression)expression).getOperand(); + } else return null; // TODO: support i = i+1; + return operand; + } + + @Nullable + private static PsiLocalVariable getCounter(PsiForeachStatement foreachStatement, + TerminalBlock tb, + List operations, + List variables) { + // have only one non-final variable + if(variables.size() != 1) return null; + + // have single expression which is either ++x or x++ + PsiExpression operand = extractIncrementedExpression(tb.getSingleStatement()); + if(!(operand instanceof PsiReferenceExpression)) return null; + PsiElement element = ((PsiReferenceExpression)operand).resolve(); + + // the referred variable is the same as non-final variable + if(!(element instanceof PsiLocalVariable) || !variables.contains(element)) return null; + + // the referred variable is not used in intermediate operations + for(Operation operation : operations) { + if(ReferencesSearch.search(element, new LocalSearchScope(operation.getExpression())).findFirst() != null) return null; + } + return (PsiLocalVariable)element; + } + private static boolean isAddAllCall(TerminalBlock tb) { final PsiVariable variable = tb.getVariable(); final PsiMethodCallExpression methodCallExpression = tb.getSingleMethodCall(); @@ -288,6 +331,12 @@ public class StreamApiMigrationInspection extends BaseJavaBatchLocalInspectionTo return mapperCall instanceof PsiReferenceExpression && ((PsiReferenceExpression)mapperCall).resolve() == variable; } + private static void reformatWhenNeeded(@NotNull Project project, PsiElement result) { + if (result != null) { + CodeStyleManager.getInstance(project).reformat(JavaCodeStyleManager.getInstance(project).shortenClassReferences(result)); + } + } + private static class ReplaceWithForeachCallFix implements LocalQuickFix { private final String myForEachMethodName; @@ -476,12 +525,6 @@ public class StreamApiMigrationInspection extends BaseJavaBatchLocalInspectionTo } } - private static void reformatWhenNeeded(@NotNull Project project, PsiElement result) { - if (result != null) { - CodeStyleManager.getInstance(project).reformat(JavaCodeStyleManager.getInstance(project).shortenClassReferences(result)); - } - } - private static String createInitializerReplacementText(PsiType varType, PsiExpression initializer) { final PsiType initializerType = initializer.getType(); final PsiClassType rawType = initializerType instanceof PsiClassType ? ((PsiClassType)initializerType).rawType() : null; @@ -513,6 +556,65 @@ public class StreamApiMigrationInspection extends BaseJavaBatchLocalInspectionTo } + private static class ReplaceWithCountFix implements LocalQuickFix { + + @NotNull + @Override + public String getName() { + return getFamilyName(); + } + + @NotNull + @Override + public String getFamilyName() { + return "Replace with count()"; + } + + @Override + public void applyFix(@NotNull Project project, @NotNull ProblemDescriptor descriptor) { + final PsiForeachStatement foreachStatement = PsiTreeUtil.getParentOfType(descriptor.getPsiElement(), PsiForeachStatement.class); + if (foreachStatement != null) { + if (!FileModificationService.getInstance().preparePsiElementForWrite(foreachStatement)) return; + final PsiElementFactory elementFactory = JavaPsiFacade.getElementFactory(project); + PsiStatement body = foreachStatement.getBody(); + final PsiExpression iteratedValue = foreachStatement.getIteratedValue(); + if (body != null && iteratedValue != null) { + final PsiParameter parameter = foreachStatement.getIterationParameter(); + TerminalBlock tb = TerminalBlock.from(parameter, body); + List intermediateOps = tb.extractOperationReplacements(elementFactory); + PsiExpression operand = extractIncrementedExpression(tb.getSingleStatement()); + if(!(operand instanceof PsiReferenceExpression)) return; + PsiElement element = ((PsiReferenceExpression)operand).resolve(); + if(!(element instanceof PsiLocalVariable)) return; + PsiLocalVariable var = (PsiLocalVariable)element; + final StringBuilder builder = generateStream(iteratedValue, intermediateOps); + builder.append(".count()"); + PsiElement declaration = var.getParent(); + if(declaration instanceof PsiDeclarationStatement) { + PsiElement[] elements = ((PsiDeclarationStatement)declaration).getDeclaredElements(); + if(elements[elements.length-1] == var && foreachStatement.equals( + PsiTreeUtil.skipSiblingsForward(declaration, PsiWhiteSpace.class, PsiComment.class))) { + PsiExpression initializer = var.getInitializer(); + if(initializer != null && initializer.getText().equals("0")) { + String typeStr = var.getType().getCanonicalText(); + String replacement = (typeStr.equals("long") ? "" : "(" + typeStr + ") ") + builder; + initializer.replace(elementFactory.createExpressionFromText(replacement, foreachStatement)); + simplifyRedundantCast(var); + foreachStatement.delete(); + reformatWhenNeeded(project, var); + return; + } + } + } + PsiElement result = foreachStatement.replace(elementFactory.createStatementFromText(var.getName()+"+="+builder+";", foreachStatement)); + simplifyRedundantCast(result); + reformatWhenNeeded(project, result); + } + } + } + + } + /** * Intermediate stream operation representation */ diff --git a/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/afterCountArray.java b/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/afterCountArray.java new file mode 100644 index 000000000000..a05d937ac819 --- /dev/null +++ b/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/afterCountArray.java @@ -0,0 +1,10 @@ +// "Replace with count()" "true" + +import java.util.Arrays; + +public class Main { + public long test(String[] array) { + long longStrings = Arrays.stream(array).map(String::trim).filter(trimmed -> trimmed.length() > 10).count(); + return longStrings; + } +} \ No newline at end of file diff --git a/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/afterCountInner.java b/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/afterCountInner.java new file mode 100644 index 000000000000..12c45a103a95 --- /dev/null +++ b/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/afterCountInner.java @@ -0,0 +1,14 @@ +// "Replace with count()" "true" +import java.util.List; +import java.util.Set; + +public class Main { + public void test(List> nested) { + int count = 0; + for(Set element : nested) { + if(element != null) { + count += element.stream().filter(str -> str.startsWith("xyz")).count(); + } + } + } +} \ No newline at end of file diff --git a/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/afterCountOuter.java b/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/afterCountOuter.java new file mode 100644 index 000000000000..5d2889738b9f --- /dev/null +++ b/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/afterCountOuter.java @@ -0,0 +1,10 @@ +// "Replace with count()" "true" +import java.util.Collection; +import java.util.List; +import java.util.Set; + +public class Main { + public void test(List> nested) { + int count = (int) nested.stream().filter(element -> element != null).flatMap(Collection::stream).filter(str -> str.startsWith("xyz")).count(); + } +} \ No newline at end of file diff --git a/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/beforeCountArray.java b/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/beforeCountArray.java new file mode 100644 index 000000000000..5d1ed4169250 --- /dev/null +++ b/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/beforeCountArray.java @@ -0,0 +1,14 @@ +// "Replace with count()" "true" + +public class Main { + public long test(String[] array) { + long longStrings = 0; + for(String str : array) { + String trimmed = str.trim(); + if(trimmed.length() > 10) { + longStrings++; + } + } + return longStrings; + } +} \ No newline at end of file diff --git a/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/beforeCountInner.java b/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/beforeCountInner.java new file mode 100644 index 000000000000..825262c52ac5 --- /dev/null +++ b/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/beforeCountInner.java @@ -0,0 +1,18 @@ +// "Replace with count()" "true" +import java.util.List; +import java.util.Set; + +public class Main { + public void test(List> nested) { + int count = 0; + for(Set element : nested) { + if(element != null) { + for(String str : element) { + if(str.startsWith("xyz")) { + count++; + } + } + } + } + } +} \ No newline at end of file diff --git a/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/beforeCountOuter.java b/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/beforeCountOuter.java new file mode 100644 index 000000000000..4da30a7d6f1f --- /dev/null +++ b/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/beforeCountOuter.java @@ -0,0 +1,18 @@ +// "Replace with count()" "true" +import java.util.List; +import java.util.Set; + +public class Main { + public void test(List> nested) { + int count = 0; + for(Set element : nested) { + if(element != null) { + for(String str : element) { + if(str.startsWith("xyz")) { + count++; + } + } + } + } + } +} \ No newline at end of file