IDEA-160802 Add inspection to optimize ...flatMap(Collection::stream).count() to ...mapToLong(Collection::size).sum();

New inspection along with Collection.stream().count() -> Collection.size() moved to separate file ReplaceInefficientStreamCountInspection
This commit is contained in:
Tagir Valeev
2016-09-07 15:20:43 +07:00
parent 70c9a795e4
commit bef996aa8b
13 changed files with 323 additions and 40 deletions
@@ -0,0 +1,189 @@
/*
* Copyright 2000-2016 JetBrains s.r.o.
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* http://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package com.intellij.codeInspection;
import com.intellij.codeInsight.FileModificationService;
import com.intellij.openapi.project.Project;
import com.intellij.psi.*;
import com.intellij.psi.util.PsiUtil;
import com.intellij.psi.util.RedundantCastUtil;
import org.jetbrains.annotations.Nls;
import org.jetbrains.annotations.NotNull;
import org.jetbrains.annotations.Nullable;
import static com.intellij.codeInspection.SimplifyStreamApiCallChainsInspection.*;
/**
* @author Tagir Valeev
*/
public class ReplaceInefficientStreamCountInspection extends BaseJavaBatchLocalInspectionTool {
private static final String COUNT_METHOD = "count";
private static final String SIZE_METHOD = "size";
private static final String STREAM_METHOD = "stream";
private static final String FLAT_MAP_METHOD = "flatMap";
@Override
public boolean isEnabledByDefault() {
return true;
}
@NotNull
@Override
public PsiElementVisitor buildVisitor(@NotNull ProblemsHolder holder, boolean isOnTheFly) {
if (!PsiUtil.isLanguageLevel8OrHigher(holder.getFile())) {
return PsiElementVisitor.EMPTY_VISITOR;
}
return new JavaElementVisitor() {
@Override
public void visitMethodCallExpression(PsiMethodCallExpression methodCall) {
final PsiMethod method = methodCall.resolveMethod();
if (isCallOf(method, CommonClassNames.JAVA_UTIL_STREAM_STREAM, COUNT_METHOD, 0)) {
final PsiMethodCallExpression qualifierCall = getQualifierMethodCall(methodCall);
if (qualifierCall == null) return;
final PsiMethod qualifier = qualifierCall.resolveMethod();
if (isCallOf(qualifier, CommonClassNames.JAVA_UTIL_COLLECTION, STREAM_METHOD, 0)) {
final StreamCountFix fix = new StreamCountFix();
holder.registerProblem(methodCall, getCallChainRange(methodCall, qualifierCall), fix.getMessage(), fix);
}
else if (isCallOf(qualifier, CommonClassNames.JAVA_UTIL_STREAM_STREAM, FLAT_MAP_METHOD, 1) &&
doesFlatMapCallCollectionStream(qualifierCall)) {
FlatMapCountFix fix = new FlatMapCountFix();
holder.registerProblem(methodCall, getCallChainRange(methodCall, qualifierCall), fix.getMessage(), fix);
}
}
}
};
}
boolean doesFlatMapCallCollectionStream(PsiMethodCallExpression flatMapCall) {
PsiElement parameter = flatMapCall.getArgumentList().getExpressions()[0];
if (parameter instanceof PsiMethodReferenceExpression) {
PsiMethodReferenceExpression methodRef = (PsiMethodReferenceExpression)parameter;
PsiElement resolvedMethodRef = methodRef.resolve();
if (resolvedMethodRef instanceof PsiMethod && isCallOf((PsiMethod)resolvedMethodRef,
CommonClassNames.JAVA_UTIL_COLLECTION, STREAM_METHOD, 0)) {
return true;
}
}
else if (parameter instanceof PsiLambdaExpression) {
PsiExpression expression = extractLambdaReturnExpression((PsiLambdaExpression)parameter);
if (expression instanceof PsiMethodCallExpression) {
PsiMethod lambdaMethod = ((PsiMethodCallExpression)expression).resolveMethod();
if (isCallOf(lambdaMethod, CommonClassNames.JAVA_UTIL_COLLECTION, STREAM_METHOD, 0)) {
return true;
}
}
}
return false;
}
@Nullable
private static PsiExpression extractLambdaReturnExpression(PsiLambdaExpression lambda) {
PsiElement lambdaBody = lambda.getBody();
PsiExpression expression = null;
if (lambdaBody instanceof PsiExpression) {
expression = (PsiExpression)lambdaBody;
}
else if (lambdaBody instanceof PsiCodeBlock) {
PsiStatement[] statements = ((PsiCodeBlock)lambdaBody).getStatements();
if (statements.length == 1 && statements[0] instanceof PsiReturnStatement) {
expression = ((PsiReturnStatement)statements[0]).getReturnValue();
}
}
return PsiUtil.skipParenthesizedExprDown(expression);
}
private static class StreamCountFix extends ReplaceStreamMethodFix {
public StreamCountFix() {
super(COUNT_METHOD, SIZE_METHOD, false);
}
@Override
protected void replaceMethodCall(@NotNull PsiMethodCallExpression methodCall,
@NotNull PsiMethodCallExpression qualifierCall,
@Nullable PsiExpression qualifierExpression) {
super.replaceMethodCall(methodCall, qualifierCall, qualifierExpression);
PsiElement parent = methodCall.getParent();
if (parent != null && !(parent instanceof PsiExpressionStatement)) {
Project project = methodCall.getProject();
PsiExpression expression =
JavaPsiFacade.getElementFactory(project).createExpressionFromText("(long) " + methodCall.getText(), methodCall);
PsiElement replacement = methodCall.replace(expression);
if (replacement instanceof PsiTypeCastExpression && RedundantCastUtil.isCastRedundant((PsiTypeCastExpression)replacement)) {
RedundantCastUtil.removeCast((PsiTypeCastExpression)replacement);
}
}
}
}
private static class FlatMapCountFix implements LocalQuickFix {
@Nls
@NotNull
@Override
public String getName() {
return getFamilyName();
}
@Nls
@NotNull
@Override
public String getFamilyName() {
return "Replace Stream.flatMap().count() with Stream.mapToLong().sum()";
}
@Override
public void applyFix(@NotNull Project project, @NotNull ProblemDescriptor descriptor) {
PsiElement element = descriptor.getStartElement();
if (!(element instanceof PsiMethodCallExpression)) return;
PsiElement countName = ((PsiMethodCallExpression)element).getMethodExpression().getReferenceNameElement();
if (countName == null) return;
PsiMethodCallExpression qualifierCall = getQualifierMethodCall((PsiMethodCallExpression)element);
if (qualifierCall == null) return;
PsiMethod qualifier = qualifierCall.resolveMethod();
if (!isCallOf(qualifier, CommonClassNames.JAVA_UTIL_STREAM_STREAM, FLAT_MAP_METHOD, 1)) return;
PsiElement flatMapName = qualifierCall.getMethodExpression().getReferenceNameElement();
if (flatMapName == null) return;
PsiElement parameter = qualifierCall.getArgumentList().getExpressions()[0];
if (!FileModificationService.getInstance().preparePsiElementForWrite(element.getContainingFile())) return;
PsiElementFactory factory = JavaPsiFacade.getElementFactory(project);
PsiElement streamCallName = null;
if (parameter instanceof PsiMethodReferenceExpression) {
PsiMethodReferenceExpression methodRef = (PsiMethodReferenceExpression)parameter;
streamCallName = methodRef.getReferenceNameElement();
}
else if (parameter instanceof PsiLambdaExpression) {
PsiExpression expression = extractLambdaReturnExpression((PsiLambdaExpression)parameter);
if (expression instanceof PsiMethodCallExpression) {
streamCallName = ((PsiMethodCallExpression)expression).getMethodExpression().getReferenceNameElement();
}
}
if (streamCallName == null || !streamCallName.getText().equals("stream")) return;
streamCallName.replace(factory.createIdentifier("size"));
flatMapName.replace(factory.createIdentifier("mapToLong"));
countName.replace(factory.createIdentifier("sum"));
PsiReferenceParameterList parameterList = qualifierCall.getMethodExpression().getParameterList();
if(parameterList != null) {
parameterList.delete();
}
}
public String getMessage() {
return "Stream.flatMap().count() can be replaced with Stream.mapToLong().sum()";
}
}
}
@@ -38,8 +38,6 @@ import java.util.Arrays;
public class SimplifyStreamApiCallChainsInspection extends BaseJavaBatchLocalInspectionTool {
private static final String FOR_EACH_METHOD = "forEach";
private static final String FOR_EACH_ORDERED_METHOD = "forEachOrdered";
private static final String COUNT_METHOD = "count";
private static final String SIZE_METHOD = "size";
private static final String STREAM_METHOD = "stream";
private static final String EMPTY_METHOD = "empty";
private static final String AS_LIST_METHOD = "asList";
@@ -154,9 +152,6 @@ public class SimplifyStreamApiCallChainsInspection extends BaseJavaBatchLocalIns
else if (isCallOf(method, CommonClassNames.JAVA_UTIL_STREAM_STREAM, FOR_EACH_ORDERED_METHOD, 1)) {
name = FOR_EACH_ORDERED_METHOD;
}
else if (isCallOf(method, CommonClassNames.JAVA_UTIL_STREAM_STREAM, COUNT_METHOD, 0)) {
name = COUNT_METHOD;
}
else {
return;
}
@@ -164,13 +159,7 @@ public class SimplifyStreamApiCallChainsInspection extends BaseJavaBatchLocalIns
if (qualifierCall == null) return;
final PsiMethod qualifier = qualifierCall.resolveMethod();
if (isCallOf(qualifier, CommonClassNames.JAVA_UTIL_COLLECTION, STREAM_METHOD, 0)) {
final ReplaceStreamMethodFix fix;
if (COUNT_METHOD.equals(name)) {
fix = new StreamCountFix();
}
else {
fix = new ReplaceStreamMethodFix(name, FOR_EACH_METHOD, true);
}
final ReplaceStreamMethodFix fix = new ReplaceStreamMethodFix(name, FOR_EACH_METHOD, true);
holder.registerProblem(methodCall, getCallChainRange(methodCall, qualifierCall), fix.getMessage(), fix);
}
}
@@ -199,7 +188,7 @@ public class SimplifyStreamApiCallChainsInspection extends BaseJavaBatchLocalIns
}
@Nullable
private static PsiMethodCallExpression getQualifierMethodCall(PsiMethodCallExpression methodCall) {
static PsiMethodCallExpression getQualifierMethodCall(PsiMethodCallExpression methodCall) {
final PsiExpression qualifierExpression = methodCall.getMethodExpression().getQualifierExpression();
if (qualifierExpression instanceof PsiMethodCallExpression) {
return (PsiMethodCallExpression)qualifierExpression;
@@ -208,8 +197,8 @@ public class SimplifyStreamApiCallChainsInspection extends BaseJavaBatchLocalIns
}
@NotNull
protected TextRange getCallChainRange(@NotNull PsiMethodCallExpression expression,
@NotNull PsiMethodCallExpression qualifierExpression) {
protected static TextRange getCallChainRange(@NotNull PsiMethodCallExpression expression,
@NotNull PsiMethodCallExpression qualifierExpression) {
final PsiReferenceExpression qualifierMethodExpression = qualifierExpression.getMethodExpression();
final PsiElement qualifierNameElement = qualifierMethodExpression.getReferenceNameElement();
final int startOffset = (qualifierNameElement != null ? qualifierNameElement : qualifierMethodExpression).getTextOffset();
@@ -371,7 +360,7 @@ public class SimplifyStreamApiCallChainsInspection extends BaseJavaBatchLocalIns
}
}
private static class ReplaceStreamMethodFix extends CallChainFixBase {
static class ReplaceStreamMethodFix extends CallChainFixBase {
private final String myStreamMethod;
private final String myCollectionMethod;
private final boolean myChangeSemantics;
@@ -416,29 +405,6 @@ public class SimplifyStreamApiCallChainsInspection extends BaseJavaBatchLocalIns
}
}
private static class StreamCountFix extends ReplaceStreamMethodFix {
public StreamCountFix() {
super(COUNT_METHOD, SIZE_METHOD, false);
}
@Override
protected void replaceMethodCall(@NotNull PsiMethodCallExpression methodCall,
@NotNull PsiMethodCallExpression qualifierCall,
@Nullable PsiExpression qualifierExpression) {
super.replaceMethodCall(methodCall, qualifierCall, qualifierExpression);
PsiElement parent = methodCall.getParent();
if(parent != null && !(parent instanceof PsiExpressionStatement)) {
Project project = methodCall.getProject();
PsiExpression expression =
JavaPsiFacade.getElementFactory(project).createExpressionFromText("(long) " + methodCall.getText(), methodCall);
PsiElement replacement = methodCall.replace(expression);
if (replacement instanceof PsiTypeCastExpression && RedundantCastUtil.isCastRedundant((PsiTypeCastExpression)replacement)) {
RedundantCastUtil.removeCast((PsiTypeCastExpression)replacement);
}
}
}
}
private static class ReplaceCollectorFix implements LocalQuickFix {
private final String myCollector;
private final String myStreamSequence;
@@ -0,0 +1,12 @@
// "Replace Stream.flatMap().count() with Stream.mapToLong().sum()" "true"
import java.util.List;
import java.util.Map;
public class Main {
private String s;
public Main(List<Map<String, String>> s) {
long count = s.stream().mapToLong(map -> map.values().size()).sum();
}
}
@@ -0,0 +1,13 @@
// "Replace Stream.flatMap().count() with Stream.mapToLong().sum()" "true"
import java.util.List;
public class Main {
private String s;
public Main(List<List<String>> s) {
long count = s.stream().mapToLong((strings) -> {
return (strings.size());
}).sum();
}
}
@@ -0,0 +1,12 @@
// "Replace Stream.flatMap().count() with Stream.mapToLong().sum()" "true"
import java.util.Collection;
import java.util.List;
public class Main {
private String s;
public Main(List<List<String>> s) {
long count = s.stream().mapToLong(Collection::size).sum();
}
}
@@ -0,0 +1,12 @@
// "Replace Stream.flatMap().count() with Stream.mapToLong().sum()" "true"
import java.util.List;
import java.util.Map;
public class Main {
private String s;
public Main(List<Map<String, String>> s) {
long count = s.stream().<String>f<caret>latMap(map -> map.values().stream()).count();
}
}
@@ -0,0 +1,13 @@
// "Replace Stream.flatMap().count() with Stream.mapToLong().sum()" "true"
import java.util.List;
public class Main {
private String s;
public Main(List<List<String>> s) {
long count = s.stream().flatMap((str<caret>ings) -> {
return (strings.stream());
}).count();
}
}
@@ -0,0 +1,12 @@
// "Replace Stream.flatMap().count() with Stream.mapToLong().sum()" "true"
import java.util.Collection;
import java.util.List;
public class Main {
private String s;
public Main(List<List<String>> s) {
long count = s.stream().flatMap(Collection::stream).cou<caret>nt();
}
}
@@ -0,0 +1,40 @@
/*
* Copyright 2000-2016 JetBrains s.r.o.
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* http://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package com.intellij.codeInspection;
import com.intellij.codeInsight.daemon.quickFix.LightQuickFixParameterizedTestCase;
import org.jetbrains.annotations.NotNull;
/**
* @author Tagir Valeev
*/
public class ReplaceInefficientStreamCountInspectionTest extends LightQuickFixParameterizedTestCase {
@NotNull
@Override
protected LocalInspectionTool[] configureLocalInspectionTools() {
return new LocalInspectionTool[]{new ReplaceInefficientStreamCountInspection()};
}
public void test() throws Exception {
doAllTests();
}
@Override
protected String getBasePath() {
return "/inspection/inefficientStreamCount";
}
}
@@ -0,0 +1,15 @@
<html>
<body>
This inspection reports stream API call chains ending with count() operation which
could be optimized.
<p>
The following call chains are replaced by this inspection:
</p>
<ul>
<li><code>Collection.stream().count()</code> &rarr; <code>Collection.size()</code>. In Java 8 Collection.stream().count()
actually iterates over collection elements to count them while Collection.size() is much faster for most of collections.</li>
<li><code>Stream.flatMap(Collection::stream).count()</code> &rarr; <code>Stream.mapToLong(Collection::size).sum()</code>. Similarly
there's no need to iterate all the nested collections. Instead, their sizes could be summed up.</li>
</ul>
</body>
</html>
@@ -8,7 +8,6 @@ It allows to avoid creating redundant temporary objects when traversing a collec
<ul>
<li><code>Collection.stream().forEach()</code> &rarr; <code>Collection.forEach()</code></li>
<li><code>Collection.stream().forEachOrdered()</code> &rarr; <code>Collection.forEach()</code></li>
<li><code>Collection.stream().count()</code> &rarr; <code>Collection.size()</code></li>
<li><code>Arrays.asList().stream()</code> &rarr; <code>Arrays.stream()</code> or <code>Stream.of()</code></li>
<li><code>Collections.singleton().stream()</code> &rarr; <code>Stream.of()</code></li>
<li><code>Collections.singletonList().stream()</code> &rarr; <code>Stream.of()</code></li>