IDEA-160789 Migration to Stream API: replace with mapToInt/Long/Double().sum() when possible

This commit is contained in:
Tagir Valeev
2016-09-06 15:50:03 +07:00
parent 6679bf239a
commit 838d8a7146
6 changed files with 443 additions and 261 deletions
@@ -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<Operation> operations,
List<PsiVariable> variables) {
private static PsiLocalVariable getIncrementedVariable(TerminalBlock tb,
List<Operation> operations,
List<PsiVariable> 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<Operation> operations,
List<PsiVariable> 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<String> 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<String> 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<String> 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<String> intermediateOps) {
restoreComments(foreachStatement, body);
final PsiParameter parameter = foreachStatement.getIterationParameter();
final PsiElementFactory elementFactory = JavaPsiFacade.getElementFactory(project);
TerminalBlock tb = TerminalBlock.from(parameter, body);
List<String> 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<String> 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<String> 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<String> 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<String> 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<String> 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 extends PsiExpression> T getSingleExpression(Class<T> 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<String> 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() + ")";
}
}
@@ -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;
}
}
@@ -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;
}
}
@@ -0,0 +1,12 @@
// "Replace with sum()" "true"
public class Main {
public long test(String[] array) {
long i = 0;
for(String a : ar<caret>ray) {
if(a.startsWith("xyz"))
i = i + a.length();
}
return i;
}
}
@@ -0,0 +1,12 @@
// "Replace with sum()" "false"
public class Main {
public long test(String[] array) {
float i = 0;
for(String a : ar<caret>ray) {
if(a.startsWith("xyz"))
i = i + a.length();
}
return i;
}
}
@@ -0,0 +1,16 @@
// "Replace with sum()" "true"
public class Main {
public double test(String[][] array) {
double d = 10;
for(String[] arr : arra<caret>y) {
if(arr != null) {
for (String a : arr) {
if (a.startsWith("xyz"))
d = d + 1.0/a.length();
}
}
}
return d;
}
}