diff --git a/java/java-analysis-impl/src/com/intellij/codeInspection/ReplaceInefficientStreamCountInspection.java b/java/java-analysis-impl/src/com/intellij/codeInspection/ReplaceInefficientStreamCountInspection.java index 7eb56143f766..dd90cf57c762 100644 --- a/java/java-analysis-impl/src/com/intellij/codeInspection/ReplaceInefficientStreamCountInspection.java +++ b/java/java-analysis-impl/src/com/intellij/codeInspection/ReplaceInefficientStreamCountInspection.java @@ -56,12 +56,12 @@ public class ReplaceInefficientStreamCountInspection extends BaseJavaBatchLocalI if (qualifierCall == null) return; final PsiMethod qualifier = qualifierCall.resolveMethod(); if (isCallOf(qualifier, CommonClassNames.JAVA_UTIL_COLLECTION, STREAM_METHOD, 0)) { - final StreamCountFix fix = new StreamCountFix(); + final CountFix fix = new CountFix(false); 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(); + final CountFix fix = new CountFix(true); holder.registerProblem(methodCall, getCallChainRange(methodCall, qualifierCall), fix.getMessage(), fix); } } @@ -70,7 +70,11 @@ public class ReplaceInefficientStreamCountInspection extends BaseJavaBatchLocalI } boolean doesFlatMapCallCollectionStream(PsiMethodCallExpression flatMapCall) { - PsiElement parameter = flatMapCall.getArgumentList().getExpressions()[0]; + PsiExpression[] parameters = flatMapCall.getArgumentList().getExpressions(); + if(parameters.length != 1) { + return false; + } + PsiElement parameter = parameters[0]; if (parameter instanceof PsiMethodReferenceExpression) { PsiMethodReferenceExpression methodRef = (PsiMethodReferenceExpression)parameter; PsiElement resolvedMethodRef = methodRef.resolve(); @@ -107,60 +111,71 @@ public class ReplaceInefficientStreamCountInspection extends BaseJavaBatchLocalI return PsiUtil.skipParenthesizedExprDown(expression); } - private static class StreamCountFix extends ReplaceStreamMethodFix { - public StreamCountFix() { - super(COUNT_METHOD, SIZE_METHOD, false); - } + private static class CountFix implements LocalQuickFix { + private final boolean myFlatMapMode; - @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); - } - } + CountFix(boolean flatMapMode) { + myFlatMapMode = flatMapMode; } - } - - private static class FlatMapCountFix implements LocalQuickFix { @Nls @NotNull @Override public String getName() { - return getFamilyName(); + return myFlatMapMode + ? "Replace Stream.flatMap().count() with Stream.mapToLong().sum()" + : "Replace Collection.stream().count() with Collection.size()"; } @Nls @NotNull @Override public String getFamilyName() { - return "Replace Stream.flatMap().count() with Stream.mapToLong().sum()"; + return "Replace inefficient Stream.count()"; } @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(); + PsiMethodCallExpression countCall = (PsiMethodCallExpression)element; + PsiElement countName = countCall.getMethodExpression().getReferenceNameElement(); if (countName == null) return; - PsiMethodCallExpression qualifierCall = getQualifierMethodCall((PsiMethodCallExpression)element); + PsiMethodCallExpression qualifierCall = getQualifierMethodCall(countCall); if (qualifierCall == null) return; PsiMethod qualifier = qualifierCall.resolveMethod(); + if (!FileModificationService.getInstance().preparePsiElementForWrite(element.getContainingFile())) return; + PsiElementFactory factory = JavaPsiFacade.getElementFactory(project); + if(myFlatMapMode) { + replaceFlatMap(countName, qualifierCall, qualifier, factory); + } + else { + replaceSimpleCount(countCall, qualifierCall, qualifier, factory); + } + } + + private static void replaceSimpleCount(PsiMethodCallExpression countCall, + PsiMethodCallExpression qualifierCall, + PsiMethod qualifier, + PsiElementFactory factory) { + if (!isCallOf(qualifier, CommonClassNames.JAVA_UTIL_COLLECTION, STREAM_METHOD, 0)) return; + PsiExpression qualifierExpression = qualifierCall.getMethodExpression().getQualifierExpression(); + if(qualifierExpression == null) return; + String replacementText = "(long) "+qualifierExpression.getText()+"."+SIZE_METHOD+"()"; + PsiElement replacement = countCall.replace(factory.createExpressionFromText(replacementText, countCall)); + if (replacement instanceof PsiTypeCastExpression && RedundantCastUtil.isCastRedundant((PsiTypeCastExpression)replacement)) { + RedundantCastUtil.removeCast((PsiTypeCastExpression)replacement); + } + } + + private static void replaceFlatMap(PsiElement countName, + PsiMethodCallExpression qualifierCall, + PsiMethod qualifier, + PsiElementFactory factory) { 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; @@ -183,7 +198,8 @@ public class ReplaceInefficientStreamCountInspection extends BaseJavaBatchLocalI } public String getMessage() { - return "Stream.flatMap().count() can be replaced with Stream.mapToLong().sum()"; + return myFlatMapMode ? "Stream.flatMap().count() can be replaced with Stream.mapToLong().sum()" : + "Collection.stream().count() can be replaced with Collection.size()"; } } } 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 06b6d62529e2..4cee8ae3c475 100644 --- a/java/java-analysis-impl/src/com/intellij/codeInspection/SimplifyStreamApiCallChainsInspection.java +++ b/java/java-analysis-impl/src/com/intellij/codeInspection/SimplifyStreamApiCallChainsInspection.java @@ -103,7 +103,7 @@ public class SimplifyStreamApiCallChainsInspection extends BaseJavaBatchLocalIns fix = new ReplaceWithStreamEmptyFix(EMPTY_SET_METHOD); } if (fix != null) { - holder.registerProblem(methodCall, null, fix.getMessage(), fix); + holder.registerProblem(methodCall, null, fix.getMessage(), new SimplifyCallChainFix(fix)); } } else if (isCallOf(method, CommonClassNames.JAVA_UTIL_STREAM_STREAM, COLLECT_METHOD, 1)) { @@ -140,7 +140,8 @@ public class SimplifyStreamApiCallChainsInspection extends BaseJavaBatchLocalIns if(nameElement != null) { range = new TextRange(nameElement.getTextOffset(), range.getEndOffset()); } - holder.registerProblem(methodCall, range.shiftRight(-methodCall.getTextOffset()), fix.getMessage(), fix); + holder.registerProblem(methodCall, range.shiftRight(-methodCall.getTextOffset()), fix.getMessage(), + new SimplifyCallChainFix(fix)); } } } @@ -160,7 +161,8 @@ public class SimplifyStreamApiCallChainsInspection extends BaseJavaBatchLocalIns final PsiMethod qualifier = qualifierCall.resolveMethod(); if (isCallOf(qualifier, CommonClassNames.JAVA_UTIL_COLLECTION, STREAM_METHOD, 0)) { final ReplaceStreamMethodFix fix = new ReplaceStreamMethodFix(name, FOR_EACH_METHOD, true); - holder.registerProblem(methodCall, getCallChainRange(methodCall, qualifierCall), fix.getMessage(), fix); + holder + .registerProblem(methodCall, getCallChainRange(methodCall, qualifierCall), fix.getMessage(), new SimplifyCallChainFix(fix)); } } } @@ -221,14 +223,39 @@ public class SimplifyStreamApiCallChainsInspection extends BaseJavaBatchLocalIns return false; } - private static abstract class CallChainFixBase implements LocalQuickFix { + interface CallChainFix { + String getName(); + void applyFix(@NotNull Project project, @NotNull ProblemDescriptor descriptor); + } + + private static class SimplifyCallChainFix implements LocalQuickFix { + private final CallChainFix myFix; + + SimplifyCallChainFix(CallChainFix fix) { + myFix = fix; + } + @Nls @NotNull @Override public String getName() { - return getFamilyName(); + return myFix.getName(); } + @Nls + @NotNull + @Override + public String getFamilyName() { + return "Simplify stream call chain"; + } + + @Override + public void applyFix(@NotNull Project project, @NotNull ProblemDescriptor descriptor) { + myFix.applyFix(project, descriptor); + } + } + + private static abstract class CallChainFixBase implements CallChainFix { @Override public void applyFix(@NotNull Project project, @NotNull ProblemDescriptor descriptor) { final PsiElement element = descriptor.getStartElement(); @@ -261,7 +288,7 @@ public class SimplifyStreamApiCallChainsInspection extends BaseJavaBatchLocalIns } @NotNull - String getMessage() { + public String getMessage() { return myQualifierCall + ".stream() can be replaced with " + ClassUtil.extractClassName(myClassName) + "." + myMethodName + "()"; } @@ -302,13 +329,6 @@ public class SimplifyStreamApiCallChainsInspection extends BaseJavaBatchLocalIns private ReplaceWithStreamOfFix(String qualifierCall) { super(qualifierCall, CommonClassNames.JAVA_UTIL_STREAM_STREAM, OF_METHOD); } - - @Nls - @NotNull - @Override - public String getFamilyName() { - return "Replace with Stream.of()"; - } } private static class ReplaceSingletonWithStreamOfFix extends ReplaceWithStreamOfFix { @@ -338,26 +358,12 @@ public class SimplifyStreamApiCallChainsInspection extends BaseJavaBatchLocalIns private ArraysAsListSingleArrayFix() { super("Arrays.asList()", CommonClassNames.JAVA_UTIL_ARRAYS, STREAM_METHOD); } - - @Nls - @NotNull - @Override - public String getFamilyName() { - return "Replace Arrays.asList().stream() with Arrays.stream()"; - } } private static class ReplaceWithStreamEmptyFix extends ReplaceCollectionStreamFix { private ReplaceWithStreamEmptyFix(String qualifierMethodName) { super("Collections." + qualifierMethodName + "()", CommonClassNames.JAVA_UTIL_STREAM_STREAM, EMPTY_METHOD); } - - @Nls - @NotNull - @Override - public String getFamilyName() { - return "Replace with Stream.empty()"; - } } static class ReplaceStreamMethodFix extends CallChainFixBase { @@ -374,14 +380,14 @@ public class SimplifyStreamApiCallChainsInspection extends BaseJavaBatchLocalIns @Nls @NotNull @Override - public String getFamilyName() { + public String getName() { return "Replace Collection.stream()." + myStreamMethod + "() with Collection." + myCollectionMethod + "()" + (myChangeSemantics ? " (may change semantics)" : ""); } @NotNull - String getMessage() { + public String getMessage() { return "Collection.stream()." + myStreamMethod + "() can be replaced with Collection." + myCollectionMethod + "()" + (myChangeSemantics ? " (may change semantics)" : ""); @@ -405,7 +411,7 @@ public class SimplifyStreamApiCallChainsInspection extends BaseJavaBatchLocalIns } } - private static class ReplaceCollectorFix implements LocalQuickFix { + private static class ReplaceCollectorFix implements CallChainFix { private final String myCollector; private final String myStreamSequence; private final String myStreamSequenceStripped; @@ -422,13 +428,6 @@ public class SimplifyStreamApiCallChainsInspection extends BaseJavaBatchLocalIns @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)" : ""); @@ -481,9 +480,9 @@ public class SimplifyStreamApiCallChainsInspection extends BaseJavaBatchLocalIns } @NotNull - String getMessage() { + public String getMessage() { return "Stream.collect(" + myCollector + - "()) can be replaced with Stream." + myStreamSequenceStripped + "()" + + "()) can be replaced with Stream." + myStreamSequenceStripped + (myChangeSemantics ? " (may change semantics when result is null)" : ""); } }