From 57d38b610fb8d20ea816d42bfe7345c928aa421c Mon Sep 17 00:00:00 2001
From: Tagir Valeev
Date: Wed, 19 Oct 2016 15:46:58 +0700
Subject: [PATCH] IDEA-162662 inspection: simplify obvious stream collect
transformation to direct java.util equivalent
---
...SimplifyStreamApiCallChainsInspection.java | 132 +++++++++++++++++-
.../afterStreamToCollection.java | 10 ++
.../afterStreamToCollectionGeneric.java | 10 ++
.../afterStreamToCollectionMyTypeAddAll.java | 18 +++
.../afterStreamToCollectionMyTypeGeneric.java | 18 +++
.../afterStreamToCollectionOtherType.java | 10 ++
.../afterStreamToList.java | 11 ++
.../afterStreamToListOtherType.java | 11 ++
.../streamApiCallChains/afterStreamToSet.java | 11 ++
.../beforeStreamToCollection.java | 10 ++
.../beforeStreamToCollectionGeneric.java | 10 ++
.../beforeStreamToCollectionInvalid.java | 10 ++
.../beforeStreamToCollectionMyType.java | 14 ++
.../beforeStreamToCollectionMyTypeAddAll.java | 18 +++
...StreamToCollectionMyTypeAddAllPrivate.java | 18 +++
...beforeStreamToCollectionMyTypeGeneric.java | 18 +++
.../beforeStreamToCollectionOtherType.java | 10 ++
.../beforeStreamToList.java | 10 ++
.../beforeStreamToListOtherType.java | 10 ++
.../beforeStreamToSet.java | 10 ++
.../SimplifyStreamApiCallChains.html | 39 +++---
21 files changed, 384 insertions(+), 24 deletions(-)
create mode 100644 java/java-tests/testData/inspection/streamApiCallChains/afterStreamToCollection.java
create mode 100644 java/java-tests/testData/inspection/streamApiCallChains/afterStreamToCollectionGeneric.java
create mode 100644 java/java-tests/testData/inspection/streamApiCallChains/afterStreamToCollectionMyTypeAddAll.java
create mode 100644 java/java-tests/testData/inspection/streamApiCallChains/afterStreamToCollectionMyTypeGeneric.java
create mode 100644 java/java-tests/testData/inspection/streamApiCallChains/afterStreamToCollectionOtherType.java
create mode 100644 java/java-tests/testData/inspection/streamApiCallChains/afterStreamToList.java
create mode 100644 java/java-tests/testData/inspection/streamApiCallChains/afterStreamToListOtherType.java
create mode 100644 java/java-tests/testData/inspection/streamApiCallChains/afterStreamToSet.java
create mode 100644 java/java-tests/testData/inspection/streamApiCallChains/beforeStreamToCollection.java
create mode 100644 java/java-tests/testData/inspection/streamApiCallChains/beforeStreamToCollectionGeneric.java
create mode 100644 java/java-tests/testData/inspection/streamApiCallChains/beforeStreamToCollectionInvalid.java
create mode 100644 java/java-tests/testData/inspection/streamApiCallChains/beforeStreamToCollectionMyType.java
create mode 100644 java/java-tests/testData/inspection/streamApiCallChains/beforeStreamToCollectionMyTypeAddAll.java
create mode 100644 java/java-tests/testData/inspection/streamApiCallChains/beforeStreamToCollectionMyTypeAddAllPrivate.java
create mode 100644 java/java-tests/testData/inspection/streamApiCallChains/beforeStreamToCollectionMyTypeGeneric.java
create mode 100644 java/java-tests/testData/inspection/streamApiCallChains/beforeStreamToCollectionOtherType.java
create mode 100644 java/java-tests/testData/inspection/streamApiCallChains/beforeStreamToList.java
create mode 100644 java/java-tests/testData/inspection/streamApiCallChains/beforeStreamToListOtherType.java
create mode 100644 java/java-tests/testData/inspection/streamApiCallChains/beforeStreamToSet.java
diff --git a/java/java-analysis-impl/src/com/intellij/codeInspection/SimplifyStreamApiCallChainsInspection.java b/java/java-analysis-impl/src/com/intellij/codeInspection/SimplifyStreamApiCallChainsInspection.java
index e29355b85b9c..30c38eae0a3e 100644
--- a/java/java-analysis-impl/src/com/intellij/codeInspection/SimplifyStreamApiCallChainsInspection.java
+++ b/java/java-analysis-impl/src/com/intellij/codeInspection/SimplifyStreamApiCallChainsInspection.java
@@ -20,7 +20,9 @@ import com.intellij.openapi.diagnostic.Logger;
import com.intellij.openapi.project.Project;
import com.intellij.openapi.util.TextRange;
import com.intellij.psi.*;
+import com.intellij.psi.codeStyle.CodeStyleManager;
import com.intellij.psi.codeStyle.JavaCodeStyleManager;
+import com.intellij.psi.impl.PsiDiamondTypeUtil;
import com.intellij.psi.util.*;
import com.siyeh.ig.psiutils.BoolUtils;
import org.jetbrains.annotations.Contract;
@@ -30,6 +32,7 @@ import org.jetbrains.annotations.Nullable;
import java.text.MessageFormat;
import java.util.Arrays;
+import java.util.stream.Stream;
/**
* @author Pavel.Dolgov
@@ -57,6 +60,9 @@ public class SimplifyStreamApiCallChainsInspection extends BaseJavaBatchLocalIns
private static final String ALL_MATCH_METHOD = "allMatch";
private static final String COUNTING_COLLECTOR = "counting";
+ private static final String TO_LIST_COLLECTOR = "toList";
+ private static final String TO_SET_COLLECTOR = "toSet";
+ private static final String TO_COLLECTION_COLLECTOR = "toCollection";
private static final String MIN_BY_COLLECTOR = "minBy";
private static final String MAX_BY_COLLECTOR = "maxBy";
private static final String MAPPING_COLLECTOR = "mapping";
@@ -150,6 +156,13 @@ public class SimplifyStreamApiCallChainsInspection extends BaseJavaBatchLocalIns
}
}
+ @Contract("null -> false")
+ private boolean isCollectionStream(PsiMethodCallExpression qualifierCall) {
+ if (qualifierCall == null) return false;
+ PsiMethod qualifier = qualifierCall.resolveMethod();
+ return isCallOf(qualifier, CommonClassNames.JAVA_UTIL_COLLECTION, STREAM_METHOD, 0);
+ }
+
private void handleStreamForEach(PsiMethodCallExpression methodCall, PsiMethod method) {
final String name;
if (isCallOf(method, CommonClassNames.JAVA_UTIL_STREAM_STREAM, FOR_EACH_METHOD, 1)) {
@@ -162,9 +175,7 @@ public class SimplifyStreamApiCallChainsInspection extends BaseJavaBatchLocalIns
return;
}
final PsiMethodCallExpression qualifierCall = getQualifierMethodCall(methodCall);
- if (qualifierCall == null) return;
- final PsiMethod qualifier = qualifierCall.resolveMethod();
- if (isCallOf(qualifier, CommonClassNames.JAVA_UTIL_COLLECTION, STREAM_METHOD, 0)) {
+ if (isCollectionStream(qualifierCall)) {
final ReplaceStreamMethodFix fix = new ReplaceStreamMethodFix(name, FOR_EACH_METHOD, true);
holder
.registerProblem(methodCall, getCallChainRange(methodCall, qualifierCall), fix.getMessage(), new SimplifyCallChainFix(fix));
@@ -176,7 +187,7 @@ public class SimplifyStreamApiCallChainsInspection extends BaseJavaBatchLocalIns
if(parameter instanceof PsiMethodCallExpression) {
PsiMethodCallExpression collectorCall = (PsiMethodCallExpression)parameter;
PsiMethod collectorMethod = collectorCall.resolveMethod();
- ReplaceCollectorFix fix = null;
+ ReplaceCollectorFix fix;
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)) {
@@ -197,9 +208,26 @@ public class SimplifyStreamApiCallChainsInspection extends BaseJavaBatchLocalIns
fix = new ReplaceCollectorFix(SUMMING_LONG_COLLECTOR, "mapToLong({0}).sum()", false);
} else if(isCallOf(collectorMethod, CommonClassNames.JAVA_UTIL_STREAM_COLLECTORS, SUMMING_DOUBLE_COLLECTOR, 1)) {
fix = new ReplaceCollectorFix(SUMMING_DOUBLE_COLLECTOR, "mapToDouble({0}).sum()", false);
+ } else {
+ PsiType type = methodCall.getType();
+ if(type instanceof PsiClassType && !(((PsiClassType)type).resolve() instanceof PsiTypeParameter)) {
+ String replacement = collectorToCollection(collectorCall);
+ if (replacement != null) {
+ PsiMethodCallExpression qualifier = getQualifierMethodCall(methodCall);
+ if (isCollectionStream(qualifier)) {
+ PsiElement startElement = qualifier.getMethodExpression().getReferenceNameElement();
+ if (startElement != null) {
+ holder.registerProblem(methodCall, new TextRange(startElement.getTextOffset() - methodCall.getTextOffset(),
+ methodCall.getTextLength()),
+ "Can be replaced with '" + replacement + "' constructor",
+ new SimplifyCallChainFix(new SimplifyCollectionCreationFix(replacement)));
+ }
+ }
+ }
+ }
+ return;
}
- if (fix != null &&
- collectorCall.getArgumentList().getExpressions().length == collectorMethod.getParameterList().getParametersCount()) {
+ if (collectorCall.getArgumentList().getExpressions().length == collectorMethod.getParameterList().getParametersCount()) {
TextRange range = methodCall.getTextRange();
PsiElement nameElement = methodCall.getMethodExpression().getReferenceNameElement();
if(nameElement != null) {
@@ -247,6 +275,51 @@ public class SimplifyStreamApiCallChainsInspection extends BaseJavaBatchLocalIns
};
}
+ private static boolean isCollectionConstructor(PsiMethod ctor) {
+ if(!ctor.getModifierList().hasExplicitModifier(PsiModifier.PUBLIC)) return false;
+ PsiParameterList list = ctor.getParameterList();
+ if(list.getParametersCount() != 1) return false;
+ PsiParameter parameter = list.getParameters()[0];
+ PsiTypeElement typeElement = parameter.getTypeElement();
+ if(typeElement == null) return false;
+ PsiType type = typeElement.getType();
+ if(!(type instanceof PsiClassType)) return false;
+ PsiClass aClass = ((PsiClassType)type).resolve();
+ if(aClass == null) return false;
+ return CommonClassNames.JAVA_UTIL_COLLECTION.equals(aClass.getQualifiedName());
+ }
+
+ @Nullable
+ private static String collectorToCollection(PsiMethodCallExpression call) {
+ PsiMethod method = call.resolveMethod();
+ if(isCallOf(method, CommonClassNames.JAVA_UTIL_STREAM_COLLECTORS, TO_LIST_COLLECTOR, 0)) {
+ return CommonClassNames.JAVA_UTIL_ARRAY_LIST;
+ }
+ if(isCallOf(method, CommonClassNames.JAVA_UTIL_STREAM_COLLECTORS, TO_SET_COLLECTOR, 0)) {
+ return CommonClassNames.JAVA_UTIL_HASH_SET;
+ }
+ if(isCallOf(method, CommonClassNames.JAVA_UTIL_STREAM_COLLECTORS, TO_COLLECTION_COLLECTOR, 1)) {
+ PsiExpression[] expressions = call.getArgumentList().getExpressions();
+ if(expressions.length == 1 && expressions[0] instanceof PsiMethodReferenceExpression) {
+ PsiMethodReferenceExpression methodRef = (PsiMethodReferenceExpression)expressions[0];
+ if(methodRef.isConstructor()) {
+ PsiElement element = methodRef.resolve();
+ if(element instanceof PsiMethod) {
+ PsiMethod ctor = (PsiMethod)element;
+ if(ctor.getParameterList().getParametersCount() == 0) {
+ PsiClass aClass = ctor.getContainingClass();
+ if (aClass != null &&
+ Stream.of(aClass.getConstructors()).anyMatch(SimplifyStreamApiCallChainsInspection::isCollectionConstructor)) {
+ return aClass.getQualifiedName();
+ }
+ }
+ }
+ }
+ }
+ }
+ return null;
+ }
+
static boolean isParentNegated(PsiMethodCallExpression methodCall) {
PsiElement parent = PsiUtil.skipParenthesizedExprUp(methodCall.getParent());
return parent instanceof PsiExpression && BoolUtils.isNegation((PsiExpression)parent);
@@ -689,4 +762,51 @@ public class SimplifyStreamApiCallChainsInspection extends BaseJavaBatchLocalIns
}
}
}
+
+ private static class SimplifyCollectionCreationFix implements CallChainFix {
+ private String myReplacement;
+
+ public SimplifyCollectionCreationFix(String replacement) {
+ myReplacement = replacement;
+ }
+
+ @Override
+ public String getName() {
+ return "Replace with '"+myReplacement+"' constructor";
+ }
+
+ @Override
+ public void applyFix(@NotNull Project project, @NotNull ProblemDescriptor descriptor) {
+ PsiElement element = descriptor.getStartElement();
+ if(!(element instanceof PsiMethodCallExpression)) return;
+ PsiMethodCallExpression collectCall = (PsiMethodCallExpression)element;
+ PsiType type = collectCall.getType();
+ if(!(type instanceof PsiClassType)) return;
+ PsiClass resolvedType = ((PsiClassType)type).resolve();
+ if(resolvedType == null || resolvedType instanceof PsiTypeParameter) return;
+ PsiMethodCallExpression streamCall = getQualifierMethodCall(collectCall);
+ if(streamCall == null) return;
+ PsiExpression collectionExpression = streamCall.getMethodExpression().getQualifierExpression();
+ if(collectionExpression == null) return;
+ String typeText = type.getCanonicalText();
+ if(CommonClassNames.JAVA_UTIL_LIST.equals(resolvedType.getQualifiedName()) ||
+ CommonClassNames.JAVA_UTIL_SET.equals(resolvedType.getQualifiedName())) {
+ PsiType[] parameters = ((PsiClassType)type).getParameters();
+ if(parameters.length != 1) return;
+ typeText = myReplacement + "<" + parameters[0].getCanonicalText() + ">";
+ }
+ if (!FileModificationService.getInstance().preparePsiElementForWrite(element)) return;
+ PsiElementFactory factory = JavaPsiFacade.getElementFactory(project);
+ PsiExpression result = factory
+ .createExpressionFromText("new " + typeText + "(" + collectionExpression.getText() + ")", element);
+ PsiNewExpression newExpression = (PsiNewExpression)element.replace(result);
+ PsiJavaCodeReferenceElement classReference = newExpression.getClassOrAnonymousClassReference();
+ LOG.assertTrue(classReference != null);
+ JavaCodeStyleManager.getInstance(project).shortenClassReferences(classReference);
+ if (PsiDiamondTypeUtil.canCollapseToDiamond(newExpression, newExpression, null)) {
+ PsiDiamondTypeUtil.replaceExplicitWithDiamond(classReference.getParameterList());
+ }
+ CodeStyleManager.getInstance(project).reformat(newExpression);
+ }
+ }
}
diff --git a/java/java-tests/testData/inspection/streamApiCallChains/afterStreamToCollection.java b/java/java-tests/testData/inspection/streamApiCallChains/afterStreamToCollection.java
new file mode 100644
index 000000000000..dccf7083d015
--- /dev/null
+++ b/java/java-tests/testData/inspection/streamApiCallChains/afterStreamToCollection.java
@@ -0,0 +1,10 @@
+// "Replace with 'java.util.TreeSet' constructor" "true"
+
+import java.util.*;
+import java.util.stream.*;
+
+class Test {
+ public static void test(List s) {
+ new TreeSet<>(s).contains("abc");
+ }
+}
\ No newline at end of file
diff --git a/java/java-tests/testData/inspection/streamApiCallChains/afterStreamToCollectionGeneric.java b/java/java-tests/testData/inspection/streamApiCallChains/afterStreamToCollectionGeneric.java
new file mode 100644
index 000000000000..4a40cf0dc51e
--- /dev/null
+++ b/java/java-tests/testData/inspection/streamApiCallChains/afterStreamToCollectionGeneric.java
@@ -0,0 +1,10 @@
+// "Replace with 'java.util.TreeSet' constructor" "true"
+
+import java.util.*;
+import java.util.stream.*;
+
+class Test {
+ public static void test(List s) {
+ new TreeSet(s).contains("abc");
+ }
+}
\ No newline at end of file
diff --git a/java/java-tests/testData/inspection/streamApiCallChains/afterStreamToCollectionMyTypeAddAll.java b/java/java-tests/testData/inspection/streamApiCallChains/afterStreamToCollectionMyTypeAddAll.java
new file mode 100644
index 000000000000..4fb25baaf003
--- /dev/null
+++ b/java/java-tests/testData/inspection/streamApiCallChains/afterStreamToCollectionMyTypeAddAll.java
@@ -0,0 +1,18 @@
+// "Replace with 'Test.MyType' constructor" "true"
+
+import java.util.*;
+import java.util.stream.*;
+
+class Test {
+ static class MyType extends ArrayList {
+ public MyType() {}
+
+ public MyType(Collection coll) {
+ super(coll);
+ }
+ }
+
+ public static void test(List s) {
+ new MyType(s).contains("abc");
+ }
+}
\ No newline at end of file
diff --git a/java/java-tests/testData/inspection/streamApiCallChains/afterStreamToCollectionMyTypeGeneric.java b/java/java-tests/testData/inspection/streamApiCallChains/afterStreamToCollectionMyTypeGeneric.java
new file mode 100644
index 000000000000..8b73ffd5cef4
--- /dev/null
+++ b/java/java-tests/testData/inspection/streamApiCallChains/afterStreamToCollectionMyTypeGeneric.java
@@ -0,0 +1,18 @@
+// "Replace with 'Test.MyType' constructor" "true"
+
+import java.util.*;
+import java.util.stream.*;
+
+class Test {
+ static class MyType extends ArrayList {
+ public MyType() {}
+
+ public MyType(Collection coll) {
+ super(coll);
+ }
+ }
+
+ public static void testMy(List s) {
+ new MyType(s).contains("abc");
+ }
+}
\ No newline at end of file
diff --git a/java/java-tests/testData/inspection/streamApiCallChains/afterStreamToCollectionOtherType.java b/java/java-tests/testData/inspection/streamApiCallChains/afterStreamToCollectionOtherType.java
new file mode 100644
index 000000000000..6b1baeccdcaa
--- /dev/null
+++ b/java/java-tests/testData/inspection/streamApiCallChains/afterStreamToCollectionOtherType.java
@@ -0,0 +1,10 @@
+// "Replace with 'java.util.TreeSet' constructor" "true"
+
+import java.util.*;
+import java.util.stream.*;
+
+class Test {
+ public static void test(List s) {
+ new TreeSet
- Collection.stream().forEach() → Collection.forEach()
- Collection.stream().forEachOrdered() → Collection.forEach()
+ collection.stream().forEach() → collection.forEach()
+ collection.stream().forEachOrdered() → collection.forEach()
+ collection.stream().collect(Collectors.toList()) → new ArrayList<>(collection)
+ collection.stream().collect(Collectors.toSet()) → new HashSet<>(collection)
+ collection.stream().collect(Collectors.toCollection(CollectionType::new)) → new CollectionType<>(collection)
Arrays.asList().stream() → Arrays.stream() or Stream.of()
Collections.singleton().stream() → Stream.of()
Collections.singletonList().stream() → Stream.of()
Collections.emptyList().stream() → Stream.empty()
Collections.emptySet().stream() → Stream.empty()
- Stream.filter().findFirst().isPresent() → Stream.anyMatch()
- Stream.filter().findAny().isPresent() → Stream.anyMatch()
- Stream.collect(Collectors.counting()) → Stream.count()
- Stream.collect(Collectors.maxBy()) → Stream.max()
- Stream.collect(Collectors.minBy()) → Stream.min()
- Stream.collect(Collectors.mapping()) → Stream.map().collect()
- Stream.collect(Collectors.reducing()) → Stream.reduce() or Stream.map().reduce()
- Stream.collect(Collectors.summingInt()) → Stream.mapToInt().sum()
- Stream.collect(Collectors.summingLong()) → Stream.mapToLong().sum()
- Stream.collect(Collectors.summingDouble()) → Stream.mapToDouble().sum()
- !Stream.anyMatch() → Stream.noneMatch()
- !Stream.anyMatch(x -> !(...)) → Stream.allMatch()
- !Stream.noneMatch() → Stream.anyMatch()
- Stream.noneMatch(x -> !(...)) → Stream.allMatch()
- Stream.allMatch(x -> !(...)) → Stream.noneMatch()
- !Stream.allMatch(x -> !(...)) → Stream.anyMatch()
+ stream.filter().findFirst().isPresent() → stream.anyMatch()
+ stream.filter().findAny().isPresent() → stream.anyMatch()
+ stream.collect(Collectors.counting()) → stream.count()
+ stream.collect(Collectors.maxBy()) → stream.max()
+ stream.collect(Collectors.minBy()) → stream.min()
+ stream.collect(Collectors.mapping()) → stream.map().collect()
+ stream.collect(Collectors.reducing()) → stream.reduce() or Stream.map().reduce()
+ stream.collect(Collectors.summingInt()) → stream.mapToInt().sum()
+ stream.collect(Collectors.summingLong()) → stream.mapToLong().sum()
+ stream.collect(Collectors.summingDouble()) → stream.mapToDouble().sum()
+ !stream.anyMatch() → stream.noneMatch()
+ !stream.anyMatch(x -> !(...)) → stream.allMatch()
+ !stream.noneMatch() → stream.anyMatch()
+ stream.noneMatch(x -> !(...)) → stream.allMatch()
+ stream.allMatch(x -> !(...)) → stream.noneMatch()
+ !stream.allMatch(x -> !(...)) → stream.anyMatch()
Note that the replacements semantic may have minor difference in some cases.