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 11430747f69d..f012f3eb8b5a 100644 --- a/java/java-analysis-impl/src/com/intellij/codeInspection/SimplifyStreamApiCallChainsInspection.java +++ b/java/java-analysis-impl/src/com/intellij/codeInspection/SimplifyStreamApiCallChainsInspection.java @@ -242,13 +242,6 @@ 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)) { @@ -325,34 +318,7 @@ public class SimplifyStreamApiCallChainsInspection extends BaseJavaBatchLocalIns } private void handleCollectionStream(PsiMethodCallExpression methodCall) { - final PsiMethodCallExpression qualifierCall = getQualifierMethodCall(methodCall); - if (qualifierCall == null) return; - final PsiMethod qualifier = qualifierCall.resolveMethod(); - ReplaceCollectionStreamFix fix = null; - if (isCallOf(qualifier, CommonClassNames.JAVA_UTIL_ARRAYS, AS_LIST_METHOD, 1)) { - if (hasSingleArrayArgument(qualifierCall)) { - fix = new ArraysAsListSingleArrayFix(); - } - else { - fix = new ReplaceWithStreamOfFix("Arrays.asList()"); - } - } - else if (isCallOf(qualifier, CommonClassNames.JAVA_UTIL_COLLECTIONS, SINGLETON_LIST_METHOD, 1)) { - if (!hasSingleArrayArgument(qualifierCall)) { - fix = new ReplaceSingletonWithStreamOfFix("Collections.singletonList()"); - } - } - else if (isCallOf(qualifier, CommonClassNames.JAVA_UTIL_COLLECTIONS, SINGLETON_METHOD, 1)) { - if (!hasSingleArrayArgument(qualifierCall)) { - fix = new ReplaceSingletonWithStreamOfFix("Collections.singleton()"); - } - } - else if (isCallOf(qualifier, CommonClassNames.JAVA_UTIL_COLLECTIONS, EMPTY_LIST_METHOD, 0)) { - fix = new ReplaceWithStreamEmptyFix(EMPTY_LIST_METHOD); - } - else if (isCallOf(qualifier, CommonClassNames.JAVA_UTIL_COLLECTIONS, EMPTY_SET_METHOD, 0)) { - fix = new ReplaceWithStreamEmptyFix(EMPTY_SET_METHOD); - } + ReplaceCollectionStreamFix fix = findCollectionStreamFix(methodCall); if (fix != null) { holder.registerProblem(methodCall, null, fix.getMessage(), new SimplifyCallChainFix(fix)); } @@ -360,6 +326,64 @@ public class SimplifyStreamApiCallChainsInspection extends BaseJavaBatchLocalIns }; } + @Nullable + private static ReplaceCollectionStreamFix findCollectionStreamFix(PsiMethodCallExpression methodCall) { + PsiMethodCallExpression qualifierCall = getQualifierMethodCall(methodCall); + if (qualifierCall == null) return null; + PsiMethod qualifier = qualifierCall.resolveMethod(); + if (isCallOf(qualifier, CommonClassNames.JAVA_UTIL_ARRAYS, AS_LIST_METHOD, 1)) { + return hasSingleArrayArgument(qualifierCall) ? new ArraysAsListSingleArrayFix() : new ReplaceWithStreamOfFix("Arrays.asList()"); + } + else if (isCallOf(qualifier, CommonClassNames.JAVA_UTIL_COLLECTIONS, SINGLETON_LIST_METHOD, 1)) { + if (!hasSingleArrayArgument(qualifierCall)) { + return new ReplaceSingletonWithStreamOfFix("Collections.singletonList()"); + } + } + else if (isCallOf(qualifier, CommonClassNames.JAVA_UTIL_COLLECTIONS, SINGLETON_METHOD, 1)) { + if (!hasSingleArrayArgument(qualifierCall)) { + return new ReplaceSingletonWithStreamOfFix("Collections.singleton()"); + } + } + else if (isCallOf(qualifier, CommonClassNames.JAVA_UTIL_COLLECTIONS, EMPTY_LIST_METHOD, 0)) { + return new ReplaceWithStreamEmptyFix(EMPTY_LIST_METHOD); + } + else if (isCallOf(qualifier, CommonClassNames.JAVA_UTIL_COLLECTIONS, EMPTY_SET_METHOD, 0)) { + return new ReplaceWithStreamEmptyFix(EMPTY_SET_METHOD); + } + return null; + } + + @Contract("null -> false") + private static boolean isCollectionStream(PsiMethodCallExpression qualifierCall) { + if (qualifierCall == null) return false; + PsiMethod qualifier = qualifierCall.resolveMethod(); + return isCallOf(qualifier, CommonClassNames.JAVA_UTIL_COLLECTION, STREAM_METHOD, 0); + } + + public static PsiElement simplifyCollectionStreamCalls(PsiElement element) { + boolean replaced = true; + while(replaced) { + replaced = false; + PsiElement[] streamCalls = + PsiTreeUtil.collectElements(element, e -> e instanceof PsiMethodCallExpression && isCollectionStream((PsiMethodCallExpression)e)); + for (PsiElement streamCall : streamCalls) { + if (streamCall.isValid()) { + ReplaceCollectionStreamFix fix = findCollectionStreamFix((PsiMethodCallExpression)streamCall); + if (fix != null) { + PsiElement replacement = fix.simplify((PsiMethodCallExpression)streamCall); + if(replacement != null) { + replaced = true; + if(element == streamCall) { + element = replacement; + } + } + } + } + } + } + return element; + } + @Nullable private static PsiArrayType getArrayType(PsiMethodCallExpression call) { PsiType type = call.getType(); @@ -521,7 +545,7 @@ public class SimplifyStreamApiCallChainsInspection extends BaseJavaBatchLocalIns interface CallChainFix { String getName(); - void applyFix(@NotNull Project project, @NotNull ProblemDescriptor descriptor); + void applyFix(@NotNull Project project, PsiElement element); } private static class SimplifyCallChainFix implements LocalQuickFix { @@ -547,31 +571,11 @@ public class SimplifyStreamApiCallChainsInspection extends BaseJavaBatchLocalIns @Override public void applyFix(@NotNull Project project, @NotNull ProblemDescriptor descriptor) { - myFix.applyFix(project, descriptor); + myFix.applyFix(project, descriptor.getStartElement()); } } - private static abstract class CallChainFixBase implements CallChainFix { - @Override - public void applyFix(@NotNull Project project, @NotNull ProblemDescriptor descriptor) { - final PsiElement element = descriptor.getStartElement(); - if (element instanceof PsiMethodCallExpression) { - final PsiMethodCallExpression expression = (PsiMethodCallExpression)element; - final PsiExpression forEachMethodQualifier = expression.getMethodExpression().getQualifierExpression(); - if (forEachMethodQualifier instanceof PsiMethodCallExpression) { - final PsiMethodCallExpression previousExpression = (PsiMethodCallExpression)forEachMethodQualifier; - final PsiExpression qualifierExpression = previousExpression.getMethodExpression().getQualifierExpression(); - replaceMethodCall(expression, previousExpression, qualifierExpression); - } - } - } - - protected abstract void replaceMethodCall(@NotNull PsiMethodCallExpression methodCall, - @NotNull PsiMethodCallExpression qualifierCall, - @Nullable PsiExpression qualifierExpression); - } - - private static abstract class ReplaceCollectionStreamFix extends CallChainFixBase { + private static abstract class ReplaceCollectionStreamFix implements CallChainFix { private final String myClassName; private final String myMethodName; private final String myQualifierCall; @@ -601,13 +605,17 @@ public class SimplifyStreamApiCallChainsInspection extends BaseJavaBatchLocalIns } @Override - protected void replaceMethodCall(@NotNull PsiMethodCallExpression methodCall, - @NotNull PsiMethodCallExpression qualifierCall, - @Nullable PsiExpression qualifierExpression) { - methodCall.getArgumentList().replace(qualifierCall.getArgumentList()); + public void applyFix(@NotNull Project project, PsiElement element) { + if (!(element instanceof PsiMethodCallExpression)) return; + simplify((PsiMethodCallExpression)element); + } - final Project project = methodCall.getProject(); - String typeParameter = getTypeParameter(qualifierCall); + @Nullable + private PsiElement simplify(PsiMethodCallExpression streamCall) { + PsiMethodCallExpression collectionCall = getQualifierMethodCall(streamCall); + if (collectionCall == null) return null; + streamCall.getArgumentList().replace(collectionCall.getArgumentList()); + String typeParameter = getTypeParameter(collectionCall); String replacement; if (typeParameter != null) { replacement = myClassName + ".<" + typeParameter + ">" + myMethodName; @@ -615,8 +623,9 @@ public class SimplifyStreamApiCallChainsInspection extends BaseJavaBatchLocalIns else { replacement = myClassName + "." + myMethodName; } - final PsiExpression newMethodExpression = JavaPsiFacade.getElementFactory(project).createExpressionFromText(replacement, methodCall); - JavaCodeStyleManager.getInstance(project).shortenClassReferences(methodCall.getMethodExpression().replace(newMethodExpression)); + Project project = streamCall.getProject(); + PsiExpression newMethodExpression = JavaPsiFacade.getElementFactory(project).createExpressionFromText(replacement, streamCall); + return JavaCodeStyleManager.getInstance(project).shortenClassReferences(streamCall.getMethodExpression().replace(newMethodExpression)); } } @@ -661,7 +670,7 @@ public class SimplifyStreamApiCallChainsInspection extends BaseJavaBatchLocalIns } } - static class ReplaceStreamMethodFix extends CallChainFixBase { + static class ReplaceStreamMethodFix implements CallChainFix { private final String myStreamMethod; private final String myCollectionMethod; private final boolean myChangeSemantics; @@ -689,19 +698,16 @@ public class SimplifyStreamApiCallChainsInspection extends BaseJavaBatchLocalIns } @Override - protected void replaceMethodCall(@NotNull PsiMethodCallExpression methodCall, - @NotNull PsiMethodCallExpression qualifierCall, - @Nullable PsiExpression qualifierExpression) { - if (qualifierExpression != null) { - final PsiElement nameElement = methodCall.getMethodExpression().getReferenceNameElement(); - if (nameElement != null) { - qualifierCall.replace(qualifierExpression); - if (!myStreamMethod.equals(myCollectionMethod)) { - final Project project = methodCall.getProject(); - PsiIdentifier collectionMethod = JavaPsiFacade.getElementFactory(project).createIdentifier(myCollectionMethod); - nameElement.replace(collectionMethod); - } - } + public void applyFix(@NotNull Project project, PsiElement element) { + if (!(element instanceof PsiMethodCallExpression)) return; + PsiMethodCallExpression streamMethodCall = (PsiMethodCallExpression)element; + PsiMethodCallExpression collectionStreamCall = getQualifierMethodCall(streamMethodCall); + if (collectionStreamCall == null) return; + PsiExpression collectionExpression = collectionStreamCall.getMethodExpression().getQualifierExpression(); + if (collectionExpression == null) return; + collectionStreamCall.replace(collectionExpression); + if (!myStreamMethod.equals(myCollectionMethod)) { + streamMethodCall.getMethodExpression().handleElementRename(myCollectionMethod); } } } @@ -729,8 +735,7 @@ public class SimplifyStreamApiCallChainsInspection extends BaseJavaBatchLocalIns } @Override - public void applyFix(@NotNull Project project, @NotNull ProblemDescriptor descriptor) { - PsiElement element = descriptor.getStartElement(); + public void applyFix(@NotNull Project project, PsiElement element) { if (element instanceof PsiMethodCallExpression) { PsiMethodCallExpression collectCall = (PsiMethodCallExpression)element; PsiExpression qualifierExpression = collectCall.getMethodExpression().getQualifierExpression(); @@ -796,8 +801,7 @@ public class SimplifyStreamApiCallChainsInspection extends BaseJavaBatchLocalIns } @Override - public void applyFix(@NotNull Project project, @NotNull ProblemDescriptor descriptor) { - PsiElement element = descriptor.getStartElement(); + public void applyFix(@NotNull Project project, PsiElement element) { if (element instanceof PsiMethodCallExpression) { PsiMethodCallExpression isPresentCall = (PsiMethodCallExpression)element; PsiExpression isPresentQualifier = isPresentCall.getMethodExpression().getQualifierExpression(); @@ -840,8 +844,7 @@ public class SimplifyStreamApiCallChainsInspection extends BaseJavaBatchLocalIns } @Override - public void applyFix(@NotNull Project project, @NotNull ProblemDescriptor descriptor) { - PsiElement element = descriptor.getStartElement(); + public void applyFix(@NotNull Project project, PsiElement element) { if(element instanceof PsiIdentifier) { String from = element.getText(); boolean removeParentNegation; @@ -895,8 +898,7 @@ public class SimplifyStreamApiCallChainsInspection extends BaseJavaBatchLocalIns } @Override - public void applyFix(@NotNull Project project, @NotNull ProblemDescriptor descriptor) { - PsiElement element = descriptor.getStartElement(); + public void applyFix(@NotNull Project project, PsiElement element) { if(!(element instanceof PsiMethodCallExpression)) return; PsiMethodCallExpression collectCall = (PsiMethodCallExpression)element; PsiType type = collectCall.getType(); @@ -934,8 +936,7 @@ public class SimplifyStreamApiCallChainsInspection extends BaseJavaBatchLocalIns } @Override - public void applyFix(@NotNull Project project, @NotNull ProblemDescriptor descriptor) { - PsiElement element = descriptor.getStartElement(); + public void applyFix(@NotNull Project project, PsiElement element) { if(!(element instanceof PsiIdentifier)) return; PsiElement parent = element.getParent(); if(!(parent instanceof PsiReferenceExpression)) return; @@ -963,8 +964,8 @@ public class SimplifyStreamApiCallChainsInspection extends BaseJavaBatchLocalIns } @Override - public void applyFix(@NotNull Project project, @NotNull ProblemDescriptor descriptor) { - PsiMethodCallExpression toArrayCall = PsiTreeUtil.getParentOfType(descriptor.getStartElement(), PsiMethodCallExpression.class); + public void applyFix(@NotNull Project project, PsiElement element) { + PsiMethodCallExpression toArrayCall = PsiTreeUtil.getParentOfType(element, PsiMethodCallExpression.class); if (toArrayCall == null) return; PsiExpression qualifier = toArrayCall.getMethodExpression().getQualifierExpression(); if(!(qualifier instanceof PsiMethodCallExpression)) return; diff --git a/java/java-impl/src/com/intellij/codeInspection/streamMigration/MigrateToStreamFix.java b/java/java-impl/src/com/intellij/codeInspection/streamMigration/MigrateToStreamFix.java index 629e4a319b9d..b63c976f4d43 100644 --- a/java/java-impl/src/com/intellij/codeInspection/streamMigration/MigrateToStreamFix.java +++ b/java/java-impl/src/com/intellij/codeInspection/streamMigration/MigrateToStreamFix.java @@ -18,6 +18,7 @@ package com.intellij.codeInspection.streamMigration; import com.intellij.codeInspection.LambdaCanBeMethodReferenceInspection; import com.intellij.codeInspection.LocalQuickFix; import com.intellij.codeInspection.ProblemDescriptor; +import com.intellij.codeInspection.SimplifyStreamApiCallChainsInspection; import com.intellij.codeInspection.streamMigration.StreamApiMigrationInspection.*; import com.intellij.openapi.project.Project; import com.intellij.psi.*; @@ -109,6 +110,7 @@ abstract class MigrateToStreamFix implements LocalQuickFix { if (result == null) return; LambdaCanBeMethodReferenceInspection.replaceAllLambdasWithMethodReferences(result); PsiDiamondTypeUtil.removeRedundantTypeArguments(result); + result = SimplifyStreamApiCallChainsInspection.simplifyCollectionStreamCalls(result); CodeStyleManager.getInstance(project).reformat(JavaCodeStyleManager.getInstance(project).shortenClassReferences(result)); } diff --git a/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/afterFromAsList.java b/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/afterFromAsList.java new file mode 100644 index 000000000000..8053495e25a1 --- /dev/null +++ b/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/afterFromAsList.java @@ -0,0 +1,11 @@ +// "Replace with collect" "true" + +import java.util.*; +import java.util.stream.Collectors; +import java.util.stream.Stream; + +public class Main { + private void collect() { + Set res = Stream.of("a", "b", "c").filter(s -> !s.isEmpty()).collect(Collectors.toSet()); + } +} diff --git a/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/afterFromEmptyList.java b/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/afterFromEmptyList.java new file mode 100644 index 000000000000..774c9ed41217 --- /dev/null +++ b/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/afterFromEmptyList.java @@ -0,0 +1,11 @@ +// "Replace with collect" "true" + +import java.util.*; +import java.util.stream.Collectors; +import java.util.stream.Stream; + +public class Main { + private void collect() { + Set res = Stream.empty().filter(s -> !s.isEmpty()).collect(Collectors.toSet()); + } +} diff --git a/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/beforeFromAsList.java b/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/beforeFromAsList.java new file mode 100644 index 000000000000..a6258eb8d94a --- /dev/null +++ b/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/beforeFromAsList.java @@ -0,0 +1,13 @@ +// "Replace with collect" "true" + +import java.util.*; + +public class Main { + private void collect() { + Set res = new HashSet<>(); + for (String s : Arrays.asList("a", "b", "c")) { + if (!s.isEmpty()) + res.add(s); + } + } +} diff --git a/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/beforeFromEmptyList.java b/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/beforeFromEmptyList.java new file mode 100644 index 000000000000..5c599f6cdb56 --- /dev/null +++ b/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/beforeFromEmptyList.java @@ -0,0 +1,13 @@ +// "Replace with collect" "true" + +import java.util.*; + +public class Main { + private void collect() { + Set res = new HashSet<>(); + for (String s : Collections.emptyList()) { + if (!s.isEmpty()) + res.add(s); + } + } +}