IDEA-160442 Warn about excessive use of collectors

This commit is contained in:
Tagir Valeev
2016-08-30 19:05:30 +03:00
parent 42764d22ff
commit 88614739d8
18 changed files with 301 additions and 0 deletions
@@ -45,6 +45,16 @@ public class SimplifyStreamApiCallChainsInspection extends BaseJavaBatchLocalIns
private static final String EMPTY_SET_METHOD = "emptySet";
private static final String SINGLETON_LIST_METHOD = "singletonList";
private static final String SINGLETON_METHOD = "singleton";
private static final String COLLECT_METHOD = "collect";
private static final String COUNTING_COLLECTOR = "counting";
private static final String MIN_BY_COLLECTOR = "minBy";
private static final String MAX_BY_COLLECTOR = "maxBy";
private static final String MAPPING_COLLECTOR = "mapping";
private static final String REDUCING_COLLECTOR = "reducing";
private static final String SUMMING_INT_COLLECTOR = "summingInt";
private static final String SUMMING_LONG_COLLECTOR = "summingLong";
private static final String SUMMING_DOUBLE_COLLECTOR = "summingDouble";
@Override
public boolean isEnabledByDefault() {
@@ -95,6 +105,44 @@ public class SimplifyStreamApiCallChainsInspection extends BaseJavaBatchLocalIns
holder.registerProblem(methodCall, null, fix.getMessage(), fix);
}
}
else if (isCallOf(method, CommonClassNames.JAVA_UTIL_STREAM_STREAM, COLLECT_METHOD, 1)) {
PsiElement parameter = methodCall.getArgumentList().getExpressions()[0];
if(parameter instanceof PsiMethodCallExpression) {
PsiMethodCallExpression collectorCall = (PsiMethodCallExpression)parameter;
PsiMethod collectorMethod = collectorCall.resolveMethod();
ReplaceCollectorFix fix = null;
if(isCallOf(collectorMethod, CommonClassNames.JAVA_UTIL_STREAM_COLLECTORS, COUNTING_COLLECTOR, 0)) {
fix = new ReplaceCollectorFix(COUNTING_COLLECTOR, "count()", false);
} else if(isCallOf(collectorMethod, CommonClassNames.JAVA_UTIL_STREAM_COLLECTORS, MIN_BY_COLLECTOR, 1)) {
fix = new ReplaceCollectorFix(MIN_BY_COLLECTOR, "min({1})", true);
} else if(isCallOf(collectorMethod, CommonClassNames.JAVA_UTIL_STREAM_COLLECTORS, MAX_BY_COLLECTOR, 1)) {
fix = new ReplaceCollectorFix(MAX_BY_COLLECTOR, "max({1})", true);
} else if(isCallOf(collectorMethod, CommonClassNames.JAVA_UTIL_STREAM_COLLECTORS, MAPPING_COLLECTOR, 2)) {
fix = new ReplaceCollectorFix(MAPPING_COLLECTOR, "map({1}).collect({2})", false);
} else if(isCallOf(collectorMethod, CommonClassNames.JAVA_UTIL_STREAM_COLLECTORS, REDUCING_COLLECTOR, 1)) {
fix = new ReplaceCollectorFix(REDUCING_COLLECTOR, "reduce({1})", true);
} else if(isCallOf(collectorMethod, CommonClassNames.JAVA_UTIL_STREAM_COLLECTORS, REDUCING_COLLECTOR, 2)) {
fix = new ReplaceCollectorFix(REDUCING_COLLECTOR, "reduce({1}, {2})", false);
} else if(isCallOf(collectorMethod, CommonClassNames.JAVA_UTIL_STREAM_COLLECTORS, REDUCING_COLLECTOR, 3)) {
fix = new ReplaceCollectorFix(REDUCING_COLLECTOR, "map({2}).reduce({1}, {3})", false);
} else if(isCallOf(collectorMethod, CommonClassNames.JAVA_UTIL_STREAM_COLLECTORS, SUMMING_INT_COLLECTOR, 1)) {
fix = new ReplaceCollectorFix(SUMMING_INT_COLLECTOR, "mapToInt({1}).sum()", false);
} else if(isCallOf(collectorMethod, CommonClassNames.JAVA_UTIL_STREAM_COLLECTORS, SUMMING_LONG_COLLECTOR, 1)) {
fix = new ReplaceCollectorFix(SUMMING_LONG_COLLECTOR, "mapToLong({1}).sum()", false);
} else if(isCallOf(collectorMethod, CommonClassNames.JAVA_UTIL_STREAM_COLLECTORS, SUMMING_DOUBLE_COLLECTOR, 1)) {
fix = new ReplaceCollectorFix(SUMMING_DOUBLE_COLLECTOR, "mapToDouble({1}).sum()", false);
}
if (fix != null &&
collectorCall.getArgumentList().getExpressions().length == collectorMethod.getParameterList().getParametersCount()) {
TextRange range = methodCall.getTextRange();
PsiElement nameElement = methodCall.getMethodExpression().getReferenceNameElement();
if(nameElement != null) {
range = new TextRange(nameElement.getTextOffset(), range.getEndOffset());
}
holder.registerProblem(methodCall, range.shiftRight(-methodCall.getTextOffset()), fix.getMessage(), fix);
}
}
}
else {
final String name;
if (isCallOf(method, CommonClassNames.JAVA_UTIL_STREAM_STREAM, FOR_EACH_METHOD, 1)) {
@@ -390,4 +438,87 @@ public class SimplifyStreamApiCallChainsInspection extends BaseJavaBatchLocalIns
}
}
}
private static class ReplaceCollectorFix implements LocalQuickFix {
private final String myCollector;
private final String myStreamSequence;
private final String myStreamSequenceStripped;
private final boolean myChangeSemantics;
public ReplaceCollectorFix(String collector, String streamSequence, boolean changeSemantics) {
myCollector = collector;
myStreamSequence = streamSequence;
myStreamSequenceStripped = streamSequence.replaceAll("\\([^)]+\\)", "()");
myChangeSemantics = changeSemantics;
}
@Nls
@NotNull
@Override
public String getName() {
return getFamilyName();
}
@Nls
@NotNull
@Override
public String getFamilyName() {
return "Replace Stream.collect(" + myCollector +
"()) with Stream." + myStreamSequenceStripped +
(myChangeSemantics ? " (may change semantics when result is null)" : "");
}
@Override
public void applyFix(@NotNull Project project, @NotNull ProblemDescriptor descriptor) {
PsiElement element = descriptor.getStartElement();
if (element instanceof PsiMethodCallExpression) {
PsiMethodCallExpression collectCall = (PsiMethodCallExpression)element;
PsiExpression qualifierExpression = collectCall.getMethodExpression().getQualifierExpression();
if (qualifierExpression != null) {
PsiElement parameter = collectCall.getArgumentList().getExpressions()[0];
if (parameter instanceof PsiMethodCallExpression) {
PsiMethodCallExpression collectorCall = (PsiMethodCallExpression)parameter;
PsiExpression[] collectorArgs = collectorCall.getArgumentList().getExpressions();
String result = myStreamSequence;
for(int i=0; i<collectorArgs.length; i++) {
result = result.replace("{"+(i+1)+"}", collectorArgs[i].getText());
}
if (!FileModificationService.getInstance().preparePsiElementForWrite(element.getContainingFile())) return;
PsiElementFactory factory = JavaPsiFacade.getElementFactory(project);
PsiExpression replacement = factory.createExpressionFromText(
qualifierExpression.getText() + "." + result, collectCall);
PsiElement expression = collectCall.replace(replacement);
// Replacements like .collect(counting()) -> .count() change the result type from boxed to primitive
// In rare cases it's necessary to add cast to return back to boxed type
// example:
// List<Integer> intList; List<String> stringList;
// intList.remove(stringList.stream().collect(summingInt(String::length)) -- remove given element
// intList.remove(stringList.stream().mapToInt(String::length).sum()) -- remove element by index
if(expression instanceof PsiExpression) {
PsiType type = ((PsiExpression)expression).getType();
if(type instanceof PsiPrimitiveType) {
PsiClassType boxedType = ((PsiPrimitiveType)type).getBoxedType(expression);
if(boxedType != null) {
PsiExpression castExpression =
factory.createExpressionFromText("(" + boxedType.getCanonicalText() + ") " + expression.getText(), expression);
PsiElement cast = expression.replace(castExpression);
if (cast instanceof PsiTypeCastExpression && RedundantCastUtil.isCastRedundant((PsiTypeCastExpression)cast)) {
RedundantCastUtil.removeCast((PsiTypeCastExpression)cast);
}
}
}
}
}
}
}
}
@NotNull
String getMessage() {
return "Stream.collect(" + myCollector +
"()) can be replaced with Stream." + myStreamSequenceStripped + "()" +
(myChangeSemantics ? " (may change semantics when result is null)" : "");
}
}
}
@@ -0,0 +1,10 @@
// "Replace Stream.collect(counting()) with Stream.count()" "true"
import java.util.List;
import java.util.stream.Collectors;
public class Main {
public long count(List<String> data) {
return data.stream().filter(x -> x.startsWith("xyz")).count();
}
}
@@ -0,0 +1,10 @@
// "Replace Stream.collect(mapping()) with Stream.map().collect()" "true"
import java.util.List;
import java.util.stream.Collectors;
public class Main {
public List<Integer> sizes(List<String> data) {
return data.stream().filter(x -> x.startsWith("xyz")).map(String::length).collect(Collectors.toList());
}
}
@@ -0,0 +1,10 @@
// "Replace Stream.collect(maxBy()) with Stream.max() (may change semantics when result is null)" "true"
import java.util.List;
import java.util.stream.Collectors;
public class Main {
public String max(List<String> data) {
return data.stream().filter(x -> x.startsWith("xyz")).max(String.CASE_INSENSITIVE_ORDER).orElse("");
}
}
@@ -0,0 +1,10 @@
// "Replace Stream.collect(minBy()) with Stream.min() (may change semantics when result is null)" "true"
import java.util.List;
import java.util.stream.Collectors;
public class Main {
public Optional<String> min(List<String> data) {
return data.stream().filter(x -> x.startsWith("xyz")).min(String.CASE_INSENSITIVE_ORDER);
}
}
@@ -0,0 +1,10 @@
// "Replace Stream.collect(reducing()) with Stream.reduce() (may change semantics when result is null)" "true"
import java.util.List;
import java.util.stream.Collectors;
public class Main {
public Optional<String> concat(List<String> data) {
return data.stream().filter(x -> x.startsWith("xyz")).reduce(String::concat);
}
}
@@ -0,0 +1,10 @@
// "Replace Stream.collect(reducing()) with Stream.reduce()" "true"
import java.util.List;
import java.util.stream.Collectors;
public class Main {
public String concat(List<String> data) {
return data.stream().filter(x -> x.startsWith("xyz")).reduce("", String::concat);
}
}
@@ -0,0 +1,10 @@
// "Replace Stream.collect(reducing()) with Stream.map().reduce()" "true"
import java.util.List;
import java.util.stream.Collectors;
public class Main {
public int sum(List<String> data) {
return data.stream().filter(x -> x.startsWith("xyz")).map(String::length).reduce(0, Integer::sum);
}
}
@@ -0,0 +1,10 @@
// "Replace Stream.collect(summingInt()) with Stream.mapToInt().sum()" "true"
import java.util.List;
import java.util.stream.Collectors;
public class Main {
public void remove(List<Integer> ints, List<String> data) {
ints.remove((Integer) data.stream().mapToInt(String::length).sum());
}
}
@@ -0,0 +1,10 @@
// "Replace Stream.collect(summingLong()) with Stream.mapToLong().sum()" "true"
import java.util.List;
import java.util.stream.Collectors;
public class Main {
public void remove(List<Integer> ints, List<String> data) {
ints.remove(data.stream().mapToLong(String::length).sum());
}
}
@@ -0,0 +1,10 @@
// "Replace Stream.collect(counting()) with Stream.count()" "true"
import java.util.List;
import java.util.stream.Collectors;
public class Main {
public long count(List<String> data) {
return data.stream().filter(x -> x.startsWith("xyz")).collect(Collectors.<caret>counting());
}
}
@@ -0,0 +1,10 @@
// "Replace Stream.collect(mapping()) with Stream.map().collect()" "true"
import java.util.List;
import java.util.stream.Collectors;
public class Main {
public List<Integer> sizes(List<String> data) {
return data.stream().filter(x -> x.startsWith("xyz")).collect(Collectors.mapping(Str<caret>ing::length, Collectors.toList()));
}
}
@@ -0,0 +1,10 @@
// "Replace Stream.collect(maxBy()) with Stream.max() (may change semantics when result is null)" "true"
import java.util.List;
import java.util.stream.Collectors;
public class Main {
public String max(List<String> data) {
return data.stream().filter(x -> x.startsWith("xyz")).collect(Collectors.maxBy(String.CASE_INSENSIT<caret>IVE_ORDER)).orElse("");
}
}
@@ -0,0 +1,10 @@
// "Replace Stream.collect(minBy()) with Stream.min() (may change semantics when result is null)" "true"
import java.util.List;
import java.util.stream.Collectors;
public class Main {
public Optional<String> min(List<String> data) {
return data.stream().filter(x -> x.startsWith("xyz")).collect(Collectors.minBy(Str<caret>ing.CASE_INSENSITIVE_ORDER));
}
}
@@ -0,0 +1,10 @@
// "Replace Stream.collect(reducing()) with Stream.reduce() (may change semantics when result is null)" "true"
import java.util.List;
import java.util.stream.Collectors;
public class Main {
public Optional<String> concat(List<String> data) {
return data.stream().filter(x -> x.startsWith("xyz")).colle<caret>ct(Collectors.reducing(String::concat));
}
}
@@ -0,0 +1,10 @@
// "Replace Stream.collect(reducing()) with Stream.map().reduce()" "true"
import java.util.List;
import java.util.stream.Collectors;
public class Main {
public int sum(List<String> data) {
return data.stream().filter(x -> x.startsWith("xyz")).colle<caret>ct(Collectors.reducing(0, String::length, Integer::sum));
}
}
@@ -0,0 +1,10 @@
// "Replace Stream.collect(summingInt()) with Stream.mapToInt().sum()" "true"
import java.util.List;
import java.util.stream.Collectors;
public class Main {
public void remove(List<Integer> ints, List<String> data) {
ints.remove(data.stream().collect(Collectors<caret>.summingInt(String::length)));
}
}
@@ -0,0 +1,10 @@
// "Replace Stream.collect(summingLong()) with Stream.mapToLong().sum()" "true"
import java.util.List;
import java.util.stream.Collectors;
public class Main {
public void remove(List<Integer> ints, List<String> data) {
ints.remove(data.stream().collect(Collectors<caret>.summingLong(String::length)));
}
}