From 838d8a7146e435499f2dd676df4702211f2beb40 Mon Sep 17 00:00:00 2001 From: Tagir Valeev Date: Tue, 6 Sep 2016 15:47:54 +0700 Subject: [PATCH] IDEA-160789 Migration to Stream API: replace with mapToInt/Long/Double().sum() when possible --- .../StreamApiMigrationInspection.java | 643 +++++++++++------- .../afterSumArrayLength.java | 10 + .../afterSumArrayLengthNested.java | 11 + .../beforeSumArrayLength.java | 12 + .../beforeSumArrayLengthFloat.java | 12 + .../beforeSumArrayLengthNested.java | 16 + 6 files changed, 443 insertions(+), 261 deletions(-) create mode 100644 java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/afterSumArrayLength.java create mode 100644 java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/afterSumArrayLengthNested.java create mode 100644 java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/beforeSumArrayLength.java create mode 100644 java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/beforeSumArrayLengthFloat.java create mode 100644 java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/beforeSumArrayLengthNested.java 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 c39096d48def..a03379554a4e 100644 --- a/java/java-analysis-impl/src/com/intellij/codeInspection/StreamApiMigrationInspection.java +++ b/java/java-analysis-impl/src/com/intellij/codeInspection/StreamApiMigrationInspection.java @@ -32,6 +32,7 @@ import com.intellij.psi.search.GlobalSearchScope; import com.intellij.psi.search.LocalSearchScope; import com.intellij.psi.search.searches.ReferencesSearch; import com.intellij.psi.util.*; +import com.intellij.util.ArrayUtil; import com.intellij.util.containers.ContainerUtil; import com.intellij.util.containers.IntArrayList; import org.jetbrains.annotations.Contract; @@ -40,10 +41,7 @@ import org.jetbrains.annotations.NotNull; import org.jetbrains.annotations.Nullable; import javax.swing.*; -import java.util.ArrayList; -import java.util.Arrays; -import java.util.Collection; -import java.util.List; +import java.util.*; import java.util.stream.Collectors; /** @@ -132,11 +130,16 @@ public class StreamApiMigrationInspection extends BaseJavaBatchLocalInspectionTo .stream().filter(variable -> !HighlightControlFlowUtil.isEffectivelyFinal(variable, body, null)) .collect(Collectors.toList()); - if(getCounter(statement, tb, operations, nonFinalVariables) != null) { + if(getIncrementedVariable(tb, operations, nonFinalVariables) != null) { holder.registerProblem(iteratedValue, "Can be replaced with count() call", ProblemHighlightType.GENERIC_ERROR_OR_WARNING, new ReplaceWithCountFix()); } + if(getAccumulatedVariable(tb, operations, nonFinalVariables) != null) { + holder.registerProblem(iteratedValue, "Can be replaced with sum() call", + ProblemHighlightType.GENERIC_ERROR_OR_WARNING, + new ReplaceWithSumFix()); + } if(!nonFinalVariables.isEmpty()) { return; } @@ -175,13 +178,59 @@ public class StreamApiMigrationInspection extends BaseJavaBatchLocalInspectionTo }; } + @Contract("null, _ -> false") private static boolean isLiteral(PsiElement element, Object value) { return element instanceof PsiLiteralExpression && value.equals(((PsiLiteralExpression)element).getValue()); } - private static PsiExpression extractIncrementedExpression(PsiStatement statement) { - if(!(statement instanceof PsiExpressionStatement)) return null; - PsiExpression expression = ((PsiExpressionStatement)statement).getExpression(); + @Contract("null -> false") + private static boolean isZero(PsiElement element) { + if(!(element instanceof PsiLiteralExpression)) return false; + Object value = ((PsiLiteralExpression)element).getValue(); + if(!(value instanceof Number)) return false; + return ((Number)value).doubleValue() == 0.0; + } + + @Nullable + private static PsiExpression extractAddend(PsiAssignmentExpression assignment) { + if(JavaTokenType.PLUSEQ.equals(assignment.getOperationTokenType())) { + return assignment.getRExpression(); + } else if(JavaTokenType.EQ.equals(assignment.getOperationTokenType())) { + if (assignment.getRExpression() instanceof PsiBinaryExpression) { + PsiBinaryExpression binOp = (PsiBinaryExpression)assignment.getRExpression(); + if(JavaTokenType.PLUS.equals(binOp.getOperationTokenType()) && binOp.getROperand() != null) { + if(binOp.getLOperand().getText().equals(assignment.getLExpression().getText())) { + return binOp.getROperand(); + } + if(binOp.getROperand().getText().equals(assignment.getLExpression().getText())) { + return binOp.getLOperand(); + } + } + } + } + return null; + } + + @Nullable + private static PsiExpression extractAccumulator(PsiAssignmentExpression assignment) { + if(JavaTokenType.PLUSEQ.equals(assignment.getOperationTokenType())) { + return assignment.getLExpression(); + } else if(JavaTokenType.EQ.equals(assignment.getOperationTokenType())) { + if (assignment.getRExpression() instanceof PsiBinaryExpression) { + PsiBinaryExpression binOp = (PsiBinaryExpression)assignment.getRExpression(); + if(JavaTokenType.PLUS.equals(binOp.getOperationTokenType()) && binOp.getROperand() != null) { + if (binOp.getLOperand().getText().equals(assignment.getLExpression().getText()) || + binOp.getROperand().getText().equals(assignment.getLExpression().getText())) { + return assignment.getLExpression(); + } + } + } + } + return null; + } + + @Contract("null -> null") + private static PsiExpression extractIncrementedLValue(PsiExpression expression) { if(expression instanceof PsiPostfixExpression) { if(JavaTokenType.PLUSPLUS.equals(((PsiPostfixExpression)expression).getOperationTokenType())) { return ((PsiPostfixExpression)expression).getOperand(); @@ -192,34 +241,22 @@ public class StreamApiMigrationInspection extends BaseJavaBatchLocalInspectionTo } } else if(expression instanceof PsiAssignmentExpression) { PsiAssignmentExpression assignment = (PsiAssignmentExpression)expression; - if(JavaTokenType.PLUSEQ.equals(assignment.getOperationTokenType())) { - if (isLiteral(assignment.getRExpression(), 1)) { - return assignment.getLExpression(); - } - } else if(JavaTokenType.EQ.equals(assignment.getOperationTokenType())) { - if (assignment.getRExpression() instanceof PsiBinaryExpression) { - PsiBinaryExpression binOp = (PsiBinaryExpression)assignment.getRExpression(); - if(JavaTokenType.PLUS.equals(binOp.getOperationTokenType()) && binOp.getROperand() != null && ( - isLiteral(binOp.getROperand(), 1) && binOp.getLOperand().getText().equals(assignment.getLExpression().getText()) || - isLiteral(binOp.getLOperand(), 1) && binOp.getROperand().getText().equals(assignment.getLExpression().getText()))) { - return assignment.getLExpression(); - } - } + if(isLiteral(extractAddend(assignment), 1)) { + return assignment.getLExpression(); } } return null; } @Nullable - private static PsiLocalVariable getCounter(PsiForeachStatement foreachStatement, - TerminalBlock tb, - List operations, - List variables) { + private static PsiLocalVariable getIncrementedVariable(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()); + // have single expression which is either ++x or x++ or x+=1 or x=x+1 + PsiExpression operand = extractIncrementedLValue(tb.getSingleExpression(PsiExpression.class)); if(!(operand instanceof PsiReferenceExpression)) return null; PsiElement element = ((PsiReferenceExpression)operand).resolve(); @@ -233,6 +270,34 @@ public class StreamApiMigrationInspection extends BaseJavaBatchLocalInspectionTo return (PsiLocalVariable)element; } + @Nullable + private static PsiLocalVariable getAccumulatedVariable(TerminalBlock tb, + List operations, + List variables) { + // have only one non-final variable + if(variables.size() != 1) return null; + + PsiAssignmentExpression assignment = tb.getSingleExpression(PsiAssignmentExpression.class); + if(assignment == null) return null; + PsiExpression operand = extractAccumulator(assignment); + 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; + PsiLocalVariable var = (PsiLocalVariable)element; + if (!(var.getType() instanceof PsiPrimitiveType) || var.getType().equalsToText("float")) return null; + + // the referred variable is not used in intermediate operations + for(Operation operation : operations) { + if(ReferencesSearch.search(var, new LocalSearchScope(operation.getExpression())).findFirst() != null) return null; + } + PsiExpression addend = extractAddend(assignment); + LOG.assertTrue(addend != null); + if(ReferencesSearch.search(var, new LocalSearchScope(addend)).findFirst() != null) return null; + return var; + } + private static boolean isAddAllCall(TerminalBlock tb) { final PsiVariable variable = tb.getVariable(); final PsiMethodCallExpression methodCallExpression = tb.getSingleMethodCall(); @@ -352,25 +417,145 @@ 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)); + static String compoundLambdaOrMethodReference(PsiVariable variable, + PsiExpression expression, + String samQualifiedName, + PsiType[] samParamTypes) { + String result = ""; + final Project project = variable.getProject(); + final JavaPsiFacade psiFacade = JavaPsiFacade.getInstance(project); + final PsiClass functionClass = psiFacade.findClass(samQualifiedName, GlobalSearchScope.allScope(project)); + for (int i = 0; i < samParamTypes.length; i++) { + if (samParamTypes[i] instanceof PsiPrimitiveType) { + samParamTypes[i] = ((PsiPrimitiveType)samParamTypes[i]).getBoxedType(expression); + } } + final PsiClassType functionalInterfaceType = functionClass != null ? psiFacade.getElementFactory().createType(functionClass, samParamTypes) : null; + final PsiVariable[] parameters = {variable}; + String methodReferenceText = LambdaCanBeMethodReferenceInspection.convertToMethodReference(expression, parameters, functionalInterfaceType, null); + if (methodReferenceText != null) { + LOG.assertTrue(functionalInterfaceType != null); + result += "(" + functionalInterfaceType.getCanonicalText() + ")" + methodReferenceText; + } else { + result += variable.getName() + " -> " + expression.getText(); + } + return result; } - private static class ReplaceWithForeachCallFix implements LocalQuickFix { - private final String myForEachMethodName; - - protected ReplaceWithForeachCallFix(String forEachMethodName) { - myForEachMethodName = forEachMethodName; - } - + private static abstract class MigrateToStreamFix implements LocalQuickFix { @NotNull @Override public String getName() { return getFamilyName(); } + @Override + public void applyFix(@NotNull Project project, @NotNull ProblemDescriptor descriptor) { + final PsiForeachStatement foreachStatement = PsiTreeUtil.getParentOfType(descriptor.getPsiElement(), PsiForeachStatement.class); + if (foreachStatement != null) { + 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); + if (!FileModificationService.getInstance().preparePsiElementForWrite(foreachStatement)) return; + PsiElementFactory factory = JavaPsiFacade.getElementFactory(project); + List replacements = tb.extractOperationReplacements(factory); + migrate(project, descriptor, foreachStatement, iteratedValue, body, tb, replacements); + } + } + } + + abstract void migrate(@NotNull Project project, + @NotNull ProblemDescriptor descriptor, + @NotNull PsiForeachStatement foreachStatement, + @NotNull PsiExpression iteratedValue, + @NotNull PsiStatement body, + @NotNull TerminalBlock tb, + @NotNull List replacements); + + static void replaceWithNumericAddition(@NotNull Project project, + PsiForeachStatement foreachStatement, + PsiLocalVariable var, + StringBuilder builder, + String expressionType) { + PsiElementFactory elementFactory = JavaPsiFacade.getElementFactory(project); + restoreComments(foreachStatement, foreachStatement.getBody()); + if (isDeclarationJustBefore(var, foreachStatement)) { + PsiExpression initializer = var.getInitializer(); + if (isZero(initializer)) { + String typeStr = var.getType().getCanonicalText(); + String replacement = (typeStr.equals(expressionType) ? "" : "(" + typeStr + ") ") + builder; + initializer.replace(elementFactory.createExpressionFromText(replacement, foreachStatement)); + foreachStatement.delete(); + simplifyAndFormat(project, var); + return; + } + } + PsiElement result = + foreachStatement.replace(elementFactory.createStatementFromText(var.getName() + "+=" + builder + ";", foreachStatement)); + simplifyAndFormat(project, result); + } + + static void simplifyAndFormat(@NotNull Project project, PsiElement result) { + if(result == null) return; + simplifyRedundantCast(result); + CodeStyleManager.getInstance(project).reformat(JavaCodeStyleManager.getInstance(project).shortenClassReferences(result)); + } + + static void simplifyRedundantCast(PsiElement result) { + for (PsiMethodReferenceExpression methodReferenceExpression : PsiTreeUtil + .findChildrenOfType(result, PsiMethodReferenceExpression.class)) { + final PsiElement parent = methodReferenceExpression.getParent(); + if (parent instanceof PsiTypeCastExpression) { + if (RedundantCastUtil.isCastRedundant((PsiTypeCastExpression)parent)) { + final PsiExpression operand = ((PsiTypeCastExpression)parent).getOperand(); + LOG.assertTrue(operand != null); + parent.replace(operand); + } + } + } + } + + static void restoreComments(PsiForeachStatement foreachStatement, PsiStatement body) { + final PsiElement parent = foreachStatement.getParent(); + for (PsiElement comment : PsiTreeUtil.findChildrenOfType(body, PsiComment.class)) { + parent.addBefore(comment, foreachStatement); + } + } + + @NotNull + static StringBuilder generateStream(PsiExpression iteratedValue, List intermediateOps) { + StringBuilder buffer = new StringBuilder(); + final PsiType iteratedValueType = iteratedValue.getType(); + if (iteratedValueType instanceof PsiArrayType) { + buffer.append("java.util.Arrays.stream(").append(iteratedValue.getText()).append(")"); + } + else { + buffer.append(getIteratedValueText(iteratedValue)); + if (!intermediateOps.isEmpty()) { + buffer.append(".stream()"); + } + } + intermediateOps.forEach(buffer::append); + return buffer; + } + + static String getIteratedValueText(PsiExpression iteratedValue) { + return iteratedValue instanceof PsiCallExpression || + iteratedValue instanceof PsiReferenceExpression || + iteratedValue instanceof PsiQualifiedExpression || + iteratedValue instanceof PsiParenthesizedExpression ? iteratedValue.getText() : "(" + iteratedValue.getText() + ")"; + } + } + + private static class ReplaceWithForeachCallFix extends MigrateToStreamFix { + private final String myForEachMethodName; + + protected ReplaceWithForeachCallFix(String forEachMethodName) { + myForEachMethodName = forEachMethodName; + } + @NotNull @Override public String getFamilyName() { @@ -378,45 +563,38 @@ public class StreamApiMigrationInspection extends BaseJavaBatchLocalInspectionTo } @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; - PsiStatement body = foreachStatement.getBody(); - final PsiExpression iteratedValue = foreachStatement.getIteratedValue(); - if (body != null && iteratedValue != null) { - restoreComments(foreachStatement, body); + void migrate(@NotNull Project project, + @NotNull ProblemDescriptor descriptor, + @NotNull PsiForeachStatement foreachStatement, + @NotNull PsiExpression iteratedValue, + @NotNull PsiStatement body, + @NotNull TerminalBlock tb, + @NotNull List intermediateOps) { + restoreComments(foreachStatement, body); - final PsiParameter parameter = foreachStatement.getIterationParameter(); - final PsiElementFactory elementFactory = JavaPsiFacade.getElementFactory(project); - TerminalBlock tb = TerminalBlock.from(parameter, body); - List intermediateOps = tb.extractOperationReplacements(elementFactory); + final PsiElementFactory elementFactory = JavaPsiFacade.getElementFactory(project); - StringBuilder buffer = generateStream(iteratedValue, intermediateOps); - PsiElement block = tb.convertToElement(elementFactory); + StringBuilder buffer = generateStream(iteratedValue, intermediateOps); + PsiElement block = tb.convertToElement(elementFactory); - buffer.append(".").append(myForEachMethodName).append("("); + buffer.append(".").append(myForEachMethodName).append("("); - final String functionalExpressionText = createForEachFunctionalExpressionText(project, block, tb.getVariable()); - PsiExpressionStatement callStatement = (PsiExpressionStatement)elementFactory.createStatementFromText(buffer.toString() + functionalExpressionText + ");", foreachStatement); - callStatement = (PsiExpressionStatement)foreachStatement.replace(callStatement); + final String functionalExpressionText = createForEachFunctionalExpressionText(project, block, tb.getVariable()); + PsiExpressionStatement callStatement = (PsiExpressionStatement)elementFactory.createStatementFromText(buffer.toString() + functionalExpressionText + ");", foreachStatement); + callStatement = (PsiExpressionStatement)foreachStatement.replace(callStatement); - final PsiExpressionList argumentList = ((PsiCallExpression)callStatement.getExpression()).getArgumentList(); - LOG.assertTrue(argumentList != null, callStatement.getText()); - final PsiExpression[] expressions = argumentList.getExpressions(); - LOG.assertTrue(expressions.length == 1); + final PsiExpressionList argumentList = ((PsiCallExpression)callStatement.getExpression()).getArgumentList(); + LOG.assertTrue(argumentList != null, callStatement.getText()); + final PsiExpression[] expressions = argumentList.getExpressions(); + LOG.assertTrue(expressions.length == 1); - if (expressions[0] instanceof PsiFunctionalExpression && ((PsiFunctionalExpression)expressions[0]).getFunctionalInterfaceType() == null) { - callStatement = - (PsiExpressionStatement)callStatement.replace(elementFactory.createStatementFromText( - buffer.toString() + "(" + tb.getVariable().getText() + ") -> " + wrapInBlock(block) + ");", callStatement)); - } - - simplifyRedundantCast(callStatement); - - CodeStyleManager.getInstance(project).reformat(JavaCodeStyleManager.getInstance(project).shortenClassReferences(callStatement)); - } + if (expressions[0] instanceof PsiFunctionalExpression && ((PsiFunctionalExpression)expressions[0]).getFunctionalInterfaceType() == null) { + callStatement = + (PsiExpressionStatement)callStatement.replace(elementFactory.createStatementFromText( + buffer.toString() + "(" + tb.getVariable().getText() + ") -> " + wrapInBlock(block) + ");", callStatement)); } + + simplifyAndFormat(project, callStatement); } private static String createForEachFunctionalExpressionText(Project project, PsiElement block, PsiVariable variable) { @@ -464,14 +642,9 @@ public class StreamApiMigrationInspection extends BaseJavaBatchLocalInspectionTo } } - private abstract static class ReplaceWithCollectAbstractFix implements LocalQuickFix { + private abstract static class ReplaceWithCollectAbstractFix extends MigrateToStreamFix { protected abstract String getMethodName(); - @NotNull - @Override - public String getName() { - return getFamilyName(); - } @NotNull @Override @@ -480,76 +653,60 @@ public class StreamApiMigrationInspection extends BaseJavaBatchLocalInspectionTo } @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 PsiType iteratedValueType = iteratedValue.getType(); - final PsiParameter parameter = foreachStatement.getIterationParameter(); - TerminalBlock tb = TerminalBlock.from(parameter, body); - List intermediateOps = tb.extractOperationReplacements(elementFactory); - final PsiMethodCallExpression methodCallExpression = tb.getSingleMethodCall(); + void migrate(@NotNull Project project, + @NotNull ProblemDescriptor descriptor, + @NotNull PsiForeachStatement foreachStatement, + @NotNull PsiExpression iteratedValue, + @NotNull PsiStatement body, + @NotNull TerminalBlock tb, + @NotNull List intermediateOps) { + final PsiElementFactory elementFactory = JavaPsiFacade.getElementFactory(project); + final PsiType iteratedValueType = iteratedValue.getType(); + final PsiMethodCallExpression methodCallExpression = tb.getSingleMethodCall(); - if (methodCallExpression == null) return; + if (methodCallExpression == null) return; - if (intermediateOps.isEmpty() && isAddAllCall(tb)) { - restoreComments(foreachStatement, body); - final PsiExpression qualifierExpression = methodCallExpression.getMethodExpression().getQualifierExpression(); - final String qualifierText = qualifierExpression != null ? qualifierExpression.getText() : ""; - final String collectionText = iteratedValueType instanceof PsiArrayType ? "java.util.Arrays.asList("+iteratedValue.getText()+")" : - getIteratedValueText(iteratedValue); - final String callText = StringUtil.getQualifiedName(qualifierText, "addAll(" + collectionText + ");"); - PsiElement result = foreachStatement.replace(elementFactory.createStatementFromText(callText, foreachStatement)); - reformatWhenNeeded(project, result); - return; - } - intermediateOps.add(createMapperFunctionalExpressionText(tb.getVariable(), methodCallExpression.getArgumentList().getExpressions()[0])); - final StringBuilder builder = generateStream(iteratedValue, intermediateOps); + restoreComments(foreachStatement, body); + if (intermediateOps.isEmpty() && isAddAllCall(tb)) { + final PsiExpression qualifierExpression = methodCallExpression.getMethodExpression().getQualifierExpression(); + final String qualifierText = qualifierExpression != null ? qualifierExpression.getText() : ""; + final String collectionText = + iteratedValueType instanceof PsiArrayType ? "java.util.Arrays.asList(" + iteratedValue.getText() + ")" : + getIteratedValueText(iteratedValue); + final String callText = StringUtil.getQualifiedName(qualifierText, "addAll(" + collectionText + ");"); + PsiElement result = foreachStatement.replace(elementFactory.createStatementFromText(callText, foreachStatement)); + simplifyAndFormat(project, result); + return; + } + intermediateOps + .add(createMapperFunctionalExpressionText(tb.getVariable(), methodCallExpression.getArgumentList().getExpressions()[0])); + final StringBuilder builder = generateStream(iteratedValue, intermediateOps); - builder.append(".collect(java.util.stream.Collectors."); - PsiElement result = null; - try { - final PsiExpression qualifierExpression = methodCallExpression.getMethodExpression().getQualifierExpression(); - if (qualifierExpression instanceof PsiReferenceExpression) { - final PsiElement resolve = ((PsiReferenceExpression)qualifierExpression).resolve(); - if (resolve instanceof PsiLocalVariable) { - PsiLocalVariable var = (PsiLocalVariable)resolve; - PsiElement declaration = var.getParent(); - if (declaration instanceof PsiDeclarationStatement) { - PsiElement[] elements = ((PsiDeclarationStatement)declaration).getDeclaredElements(); - if (elements[elements.length - 1] == resolve && - foreachStatement.equals(PsiTreeUtil.skipSiblingsForward(declaration, PsiWhiteSpace.class, PsiComment.class))) { - final PsiExpression initializer = var.getInitializer(); - if (initializer instanceof PsiNewExpression) { - final PsiExpressionList argumentList = ((PsiNewExpression)initializer).getArgumentList(); - if (argumentList != null && argumentList.getExpressions().length == 0) { - restoreComments(foreachStatement, body); - final String callText = builder.toString() + createInitializerReplacementText(var.getType(), initializer) + ")"; - result = initializer.replace(elementFactory.createExpressionFromText(callText, null)); - simplifyRedundantCast(result); - foreachStatement.delete(); - return; - } - } - } - } + builder.append(".collect(java.util.stream.Collectors."); + final PsiExpression qualifierExpression = methodCallExpression.getMethodExpression().getQualifierExpression(); + if (qualifierExpression instanceof PsiReferenceExpression) { + final PsiElement resolve = ((PsiReferenceExpression)qualifierExpression).resolve(); + if (resolve instanceof PsiLocalVariable) { + PsiLocalVariable var = (PsiLocalVariable)resolve; + if (isDeclarationJustBefore(var, foreachStatement)) { + final PsiExpression initializer = var.getInitializer(); + if (initializer instanceof PsiNewExpression) { + final PsiExpressionList argumentList = ((PsiNewExpression)initializer).getArgumentList(); + if (argumentList != null && argumentList.getExpressions().length == 0) { + final String callText = builder.toString() + createInitializerReplacementText(var.getType(), initializer) + ")"; + PsiElement result = initializer.replace(elementFactory.createExpressionFromText(callText, null)); + simplifyAndFormat(project, result); + foreachStatement.delete(); + return; } } - restoreComments(foreachStatement, body); - final String qualifierText = qualifierExpression != null ? qualifierExpression.getText() : ""; - final String callText = StringUtil.getQualifiedName(qualifierText, "addAll(" + builder.toString() + "toList()));"); - result = foreachStatement.replace(elementFactory.createStatementFromText(callText, foreachStatement)); - simplifyRedundantCast(result); - } - finally { - reformatWhenNeeded(project, result); } } } + final String qualifierText = qualifierExpression != null ? qualifierExpression.getText() : ""; + final String callText = StringUtil.getQualifiedName(qualifierText, "addAll(" + builder.toString() + "toList()));"); + PsiElement result = foreachStatement.replace(elementFactory.createStatementFromText(callText, foreachStatement)); + simplifyAndFormat(project, result); } private static String createInitializerReplacementText(PsiType varType, PsiExpression initializer) { @@ -583,13 +740,7 @@ public class StreamApiMigrationInspection extends BaseJavaBatchLocalInspectionTo } - private static class ReplaceWithCountFix implements LocalQuickFix { - - @NotNull - @Override - public String getName() { - return getFamilyName(); - } + private static class ReplaceWithCountFix extends MigrateToStreamFix { @NotNull @Override @@ -598,48 +749,83 @@ public class StreamApiMigrationInspection extends BaseJavaBatchLocalInspectionTo } @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); - } - } + void migrate(@NotNull Project project, + @NotNull ProblemDescriptor descriptor, + @NotNull PsiForeachStatement foreachStatement, + @NotNull PsiExpression iteratedValue, + @NotNull PsiStatement body, + @NotNull TerminalBlock tb, + @NotNull List intermediateOps) { + PsiExpression operand = extractIncrementedLValue(tb.getSingleExpression(PsiExpression.class)); + 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()"); + replaceWithNumericAddition(project, foreachStatement, var, builder, "long"); + } + } + + private static class ReplaceWithSumFix extends MigrateToStreamFix { + + @NotNull + @Override + public String getFamilyName() { + return "Replace with sum()"; } + @Override + void migrate(@NotNull Project project, + @NotNull ProblemDescriptor descriptor, + @NotNull PsiForeachStatement foreachStatement, + @NotNull PsiExpression iteratedValue, + @NotNull PsiStatement body, + @NotNull TerminalBlock tb, + @NotNull List intermediateOps) { + PsiAssignmentExpression assignment = tb.getSingleExpression(PsiAssignmentExpression.class); + if (assignment == null) return; + PsiExpression operand = extractAccumulator(assignment); + 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); + + PsiExpression addend = extractAddend(assignment); + if (addend == null) return; + PsiType type = var.getType(); + if (!(type instanceof PsiPrimitiveType)) return; + PsiPrimitiveType primitiveType = (PsiPrimitiveType)type; + if (primitiveType.equalsToText("float")) return; + String typeName; + if (primitiveType.equalsToText("double")) { + typeName = "Double"; + } + else if (primitiveType.equalsToText("long")) { + typeName = "Long"; + } + else { + typeName = "Int"; + } + builder.append(".mapTo").append(typeName).append('('); + builder.append(compoundLambdaOrMethodReference(tb.getVariable(), addend, "java.util.function.To" + typeName + "Function", + new PsiType[]{tb.getVariable().getType()})); + builder.append(").sum()"); + replaceWithNumericAddition(project, foreachStatement, var, builder, typeName.toLowerCase(Locale.ENGLISH)); + } + } + + private static boolean isDeclarationJustBefore(PsiLocalVariable var, PsiStatement nextStatement) { + PsiElement declaration = var.getParent(); + if(declaration instanceof PsiDeclarationStatement) { + PsiElement[] elements = ((PsiDeclarationStatement)declaration).getDeclaredElements(); + if (ArrayUtil.getLastElement(elements) == var && nextStatement.equals( + PsiTreeUtil.skipSiblingsForward(declaration, PsiWhiteSpace.class, PsiComment.class))) { + return true; + } + } + return false; } /** @@ -740,18 +926,23 @@ public class StreamApiMigrationInspection extends BaseJavaBatchLocalInspectionTo return myStatements.length == 1 ? myStatements[0] : null; } + @Nullable + T getSingleExpression(Class wantedType) { + PsiStatement statement = getSingleStatement(); + if(statement instanceof PsiExpressionStatement) { + PsiExpression expression = ((PsiExpressionStatement)statement).getExpression(); + if(wantedType.isInstance(expression)) + return wantedType.cast(expression); + } + return null; + } + /** * @return PsiMethodCallExpression if this TerminalBlock contains single method call, null otherwise */ @Nullable PsiMethodCallExpression getSingleMethodCall() { - PsiStatement statement = getSingleStatement(); - if(statement instanceof PsiExpressionStatement) { - PsiExpression expression = ((PsiExpressionStatement)statement).getExpression(); - if(expression instanceof PsiMethodCallExpression) - return (PsiMethodCallExpression)expression; - } - return null; + return getSingleExpression(PsiMethodCallExpression.class); } /** @@ -883,74 +1074,4 @@ public class StreamApiMigrationInspection extends BaseJavaBatchLocalInspectionTo return block; } } - - private static String compoundLambdaOrMethodReference(PsiVariable variable, - PsiExpression expression, - String samQualifiedName, - PsiType[] samParamTypes) { - String result = ""; - final Project project = variable.getProject(); - final JavaPsiFacade psiFacade = JavaPsiFacade.getInstance(project); - final PsiClass functionClass = psiFacade.findClass(samQualifiedName, GlobalSearchScope.allScope(project)); - for (int i = 0; i < samParamTypes.length; i++) { - if (samParamTypes[i] instanceof PsiPrimitiveType) { - samParamTypes[i] = ((PsiPrimitiveType)samParamTypes[i]).getBoxedType(expression); - } - } - final PsiClassType functionalInterfaceType = functionClass != null ? psiFacade.getElementFactory().createType(functionClass, samParamTypes) : null; - final PsiVariable[] parameters = {variable}; - String methodReferenceText = LambdaCanBeMethodReferenceInspection.convertToMethodReference(expression, parameters, functionalInterfaceType, null); - if (methodReferenceText != null) { - LOG.assertTrue(functionalInterfaceType != null); - result += "(" + functionalInterfaceType.getCanonicalText() + ")" + methodReferenceText; - } else { - result += variable.getName() + " -> " + expression.getText(); - } - return result; - } - - private static void simplifyRedundantCast(PsiElement result) { - for (PsiMethodReferenceExpression methodReferenceExpression : PsiTreeUtil - .findChildrenOfType(result, PsiMethodReferenceExpression.class)) { - final PsiElement parent = methodReferenceExpression.getParent(); - if (parent instanceof PsiTypeCastExpression) { - if (RedundantCastUtil.isCastRedundant((PsiTypeCastExpression)parent)) { - final PsiExpression operand = ((PsiTypeCastExpression)parent).getOperand(); - LOG.assertTrue(operand != null); - parent.replace(operand); - } - } - } - } - - private static void restoreComments(PsiForeachStatement foreachStatement, PsiStatement body) { - final PsiElement parent = foreachStatement.getParent(); - for (PsiElement comment : PsiTreeUtil.findChildrenOfType(body, PsiComment.class)) { - parent.addBefore(comment, foreachStatement); - } - } - - @NotNull - private static StringBuilder generateStream(PsiExpression iteratedValue, List intermediateOps) { - StringBuilder buffer = new StringBuilder(); - final PsiType iteratedValueType = iteratedValue.getType(); - if (iteratedValueType instanceof PsiArrayType) { - buffer.append("java.util.Arrays.stream(").append(iteratedValue.getText()).append(")"); - } - else { - buffer.append(getIteratedValueText(iteratedValue)); - if (!intermediateOps.isEmpty()) { - buffer.append(".stream()"); - } - } - intermediateOps.forEach(buffer::append); - return buffer; - } - - private static String getIteratedValueText(PsiExpression iteratedValue) { - return iteratedValue instanceof PsiCallExpression || - iteratedValue instanceof PsiReferenceExpression || - iteratedValue instanceof PsiQualifiedExpression || - iteratedValue instanceof PsiParenthesizedExpression ? iteratedValue.getText() : "(" + iteratedValue.getText() + ")"; - } } diff --git a/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/afterSumArrayLength.java b/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/afterSumArrayLength.java new file mode 100644 index 000000000000..674a7992ffbf --- /dev/null +++ b/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/afterSumArrayLength.java @@ -0,0 +1,10 @@ +// "Replace with sum()" "true" + +import java.util.Arrays; + +public class Main { + public long test(String[] array) { + long i = Arrays.stream(array).filter(a -> a.startsWith("xyz")).mapToLong(String::length).sum(); + return i; + } +} \ No newline at end of file diff --git a/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/afterSumArrayLengthNested.java b/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/afterSumArrayLengthNested.java new file mode 100644 index 000000000000..686fcbaae7fb --- /dev/null +++ b/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/afterSumArrayLengthNested.java @@ -0,0 +1,11 @@ +// "Replace with sum()" "true" + +import java.util.Arrays; + +public class Main { + public double test(String[][] array) { + double d = 10; + d += Arrays.stream(array).filter(arr -> arr != null).flatMap(Arrays::stream).filter(a -> a.startsWith("xyz")).mapToDouble(a -> 1.0 / a.length()).sum(); + return d; + } +} \ No newline at end of file diff --git a/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/beforeSumArrayLength.java b/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/beforeSumArrayLength.java new file mode 100644 index 000000000000..957844fa0504 --- /dev/null +++ b/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/beforeSumArrayLength.java @@ -0,0 +1,12 @@ +// "Replace with sum()" "true" + +public class Main { + public long test(String[] array) { + long i = 0; + for(String a : array) { + if(a.startsWith("xyz")) + i = i + a.length(); + } + return i; + } +} \ No newline at end of file diff --git a/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/beforeSumArrayLengthFloat.java b/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/beforeSumArrayLengthFloat.java new file mode 100644 index 000000000000..517047c531ca --- /dev/null +++ b/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/beforeSumArrayLengthFloat.java @@ -0,0 +1,12 @@ +// "Replace with sum()" "false" + +public class Main { + public long test(String[] array) { + float i = 0; + for(String a : array) { + if(a.startsWith("xyz")) + i = i + a.length(); + } + return i; + } +} \ No newline at end of file diff --git a/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/beforeSumArrayLengthNested.java b/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/beforeSumArrayLengthNested.java new file mode 100644 index 000000000000..830c3c9da3e6 --- /dev/null +++ b/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/beforeSumArrayLengthNested.java @@ -0,0 +1,16 @@ +// "Replace with sum()" "true" + +public class Main { + public double test(String[][] array) { + double d = 10; + for(String[] arr : array) { + if(arr != null) { + for (String a : arr) { + if (a.startsWith("xyz")) + d = d + 1.0/a.length(); + } + } + } + return d; + } +} \ No newline at end of file