diff --git a/java/java-impl/src/com/intellij/codeInspection/streamToLoop/Operation.java b/java/java-impl/src/com/intellij/codeInspection/streamToLoop/Operation.java index 57b990739c59..8c7d8a1d5dbb 100644 --- a/java/java-impl/src/com/intellij/codeInspection/streamToLoop/Operation.java +++ b/java/java-impl/src/com/intellij/codeInspection/streamToLoop/Operation.java @@ -18,6 +18,8 @@ package com.intellij.codeInspection.streamToLoop; import com.intellij.codeInspection.streamToLoop.StreamToLoopInspection.StreamToLoopReplacementContext; import com.intellij.psi.*; import com.intellij.psi.util.PsiTypesUtil; +import com.siyeh.ig.psiutils.BoolUtils; +import com.siyeh.ig.psiutils.StreamApiUtil; import one.util.streamex.StreamEx; import org.jetbrains.annotations.NotNull; import org.jetbrains.annotations.Nullable; @@ -188,11 +190,18 @@ abstract class Operation { private String myVarName; private final FunctionHelper myFn; private final List myRecords; + private PsiExpression myCondition; + private final boolean myInverted; - private FlatMapOperation(String varName, FunctionHelper fn, List records) { + private FlatMapOperation(String varName, + FunctionHelper fn, + List records, + PsiExpression condition, boolean inverted) { myVarName = varName; myFn = fn; myRecords = records; + myCondition = condition; + myInverted = inverted; } @Override @@ -213,11 +222,17 @@ abstract class Operation { @Override public void registerReusedElements(Consumer consumer) { myRecords.forEach(or -> or.myOperation.registerReusedElements(consumer)); + if(myCondition != null) { + consumer.accept(myCondition); + } } @Override void rename(String oldName, String newName, StreamToLoopReplacementContext context) { myRecords.forEach(or -> or.myOperation.rename(oldName, newName, context)); + if (myCondition != null) { + myCondition = FunctionHelper.replaceVarReference(myCondition, oldName, newName, context); + } } @Override @@ -230,6 +245,13 @@ abstract class Operation { for(StreamToLoopInspection.OperationRecord or : StreamEx.ofReversed(myRecords)) { replacement = or.myOperation.wrap(or.myInVar, or.myOutVar, replacement, innerContext); } + if (myCondition != null) { + String conditionText = myCondition.getText(); + if (myInverted) { + conditionText = BoolUtils.getNegatedExpressionText(context.createExpression(conditionText)); + } + return "if(" + conditionText + "){\n" + replacement + "}\n"; + } return replacement; } @@ -240,11 +262,27 @@ abstract class Operation { String varName = fn.tryLightTransform(inType); if(varName == null) return null; PsiExpression body = fn.getExpression(); + PsiExpression condition = null; + boolean inverted = false; + if(body instanceof PsiConditionalExpression) { + PsiConditionalExpression ternary = (PsiConditionalExpression)body; + condition = ternary.getCondition(); + PsiExpression thenExpression = ternary.getThenExpression(); + PsiExpression elseExpression = ternary.getElseExpression(); + if(StreamApiUtil.isNullOrEmptyStream(thenExpression)) { + body = elseExpression; + inverted = true; + } + else if(StreamApiUtil.isNullOrEmptyStream(elseExpression)) { + body = thenExpression; + } + else return null; + } if(!(body instanceof PsiMethodCallExpression)) return null; PsiMethodCallExpression terminalCall = (PsiMethodCallExpression)body; List records = StreamToLoopInspection.extractOperations(outVar, terminalCall); if(records == null || StreamToLoopInspection.getTerminal(records) != null) return null; - return new FlatMapOperation(varName, fn, records); + return new FlatMapOperation(varName, fn, records, condition, inverted); } } diff --git a/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamToLoop/afterFlatMapConditional.java b/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamToLoop/afterFlatMapConditional.java new file mode 100644 index 000000000000..86e03a4f78da --- /dev/null +++ b/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamToLoop/afterFlatMapConditional.java @@ -0,0 +1,16 @@ +// "Replace Stream API chain with loop" "true" + +import java.util.*; +import java.util.stream.*; + +public class Main { + public void test(List> list) { + for (List lst : list) { + if (lst != null) { + for (String s : lst) { + System.out.println(s); + } + } + } + } +} \ No newline at end of file diff --git a/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamToLoop/beforeFlatMapConditional.java b/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamToLoop/beforeFlatMapConditional.java new file mode 100644 index 000000000000..0289cfac69ae --- /dev/null +++ b/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamToLoop/beforeFlatMapConditional.java @@ -0,0 +1,10 @@ +// "Replace Stream API chain with loop" "true" + +import java.util.*; +import java.util.stream.*; + +public class Main { + public void test(List> list) { + list.stream().flatMap(lst -> lst == null ? Stream.empty() : lst.stream()).forEach(System.out::println); + } +} \ No newline at end of file diff --git a/plugins/InspectionGadgets/InspectionGadgetsAnalysis/src/com/siyeh/ig/psiutils/StreamApiUtil.java b/plugins/InspectionGadgets/InspectionGadgetsAnalysis/src/com/siyeh/ig/psiutils/StreamApiUtil.java index cfb6b1c98279..4ac268baa255 100644 --- a/plugins/InspectionGadgets/InspectionGadgetsAnalysis/src/com/siyeh/ig/psiutils/StreamApiUtil.java +++ b/plugins/InspectionGadgets/InspectionGadgetsAnalysis/src/com/siyeh/ig/psiutils/StreamApiUtil.java @@ -46,4 +46,22 @@ public class StreamApiUtil { } return streamType; } + + public static boolean isNullOrEmptyStream(PsiExpression expression) { + if(ExpressionUtils.isNullLiteral(expression)) { + return true; + } + if (!(expression instanceof PsiMethodCallExpression)) return false; + PsiMethodCallExpression call = (PsiMethodCallExpression)expression; + String name = call.getMethodExpression().getReferenceName(); + if ((!"empty".equals(name) && !"of".equals(name)) || !(call.getArgumentList().getExpressions().length == 0)) { + return false; + } + PsiMethod method = call.resolveMethod(); + if (method == null || !method.hasModifierProperty(PsiModifier.STATIC)) return false; + PsiClass aClass = method.getContainingClass(); + if(aClass == null) return false; + String qualifiedName = aClass.getQualifiedName(); + return qualifiedName != null && qualifiedName.startsWith("java.util.stream."); + } }