diff --git a/java/java-analysis-impl/src/com/intellij/codeInspection/streamMigration/StreamApiMigrationInspection.java b/java/java-analysis-impl/src/com/intellij/codeInspection/streamMigration/StreamApiMigrationInspection.java index 410ac8377e07..0108a0abc26f 100644 --- a/java/java-analysis-impl/src/com/intellij/codeInspection/streamMigration/StreamApiMigrationInspection.java +++ b/java/java-analysis-impl/src/com/intellij/codeInspection/streamMigration/StreamApiMigrationInspection.java @@ -772,10 +772,32 @@ public class StreamApiMigrationInspection extends BaseJavaBatchLocalInspectionTo @Override public String createReplacement(PsiElementFactory factory) { + PsiExpression intermediate = makeIntermediateExpression(factory); PsiExpression expression = - myNegated ? factory.createExpressionFromText(BoolUtils.getNegatedExpressionText(myExpression), myExpression) : myExpression; + myNegated ? factory.createExpressionFromText(BoolUtils.getNegatedExpressionText(intermediate), myExpression) : intermediate; return ".filter(" + LambdaUtil.createLambda(myVariable, expression) + ")"; } + + PsiExpression makeIntermediateExpression(PsiElementFactory factory) { + return myExpression; + } + } + + static class CompoundFilterOp extends FilterOp { + private final FlatMapOp myFlatMapOp; + private final PsiVariable myMatchVariable; + + CompoundFilterOp(FilterOp source, FlatMapOp flatMapOp) { + super(source.getExpression(), flatMapOp.myVariable, source.myNegated); + myMatchVariable = source.myVariable; + myFlatMapOp = flatMapOp; + } + + @Override + PsiExpression makeIntermediateExpression(PsiElementFactory factory) { + return factory.createExpressionFromText(myFlatMapOp.getStreamExpression()+".anyMatch("+ + LambdaUtil.createLambda(myMatchVariable, myExpression)+")", myExpression); + } } static class MapOp extends Operation { @@ -822,25 +844,35 @@ public class StreamApiMigrationInspection extends BaseJavaBatchLocalInspectionTo } static class FlatMapOp extends Operation { - FlatMapOp(PsiExpression expression, PsiVariable variable) { + private final PsiLoopStatement myLoop; + + FlatMapOp(PsiExpression expression, PsiVariable variable, PsiLoopStatement loop) { super(expression, variable); + myLoop = loop; } @Override public String createReplacement(PsiElementFactory factory) { - PsiExpression replacement = factory.createExpressionFromText(myExpression.getText() + ".stream()", myExpression); - return ".flatMap(" + LambdaUtil.createLambda(myVariable, replacement) + ")"; + return ".flatMap(" + myVariable.getName() + " -> " + getStreamExpression() + ")"; + } + + @NotNull + String getStreamExpression() { + return myExpression.getText() + ".stream()"; + } + + boolean breaksMe(PsiBreakStatement statement) { + return statement.findExitedStatement() == myLoop; } } - static class ArrayFlatMapOp extends Operation { - ArrayFlatMapOp(PsiExpression expression, PsiVariable variable) { - super(expression, variable); + static class ArrayFlatMapOp extends FlatMapOp { + ArrayFlatMapOp(PsiExpression expression, PsiVariable variable, PsiLoopStatement loop) { + super(expression, variable, loop); } @Override public String createReplacement(PsiElementFactory factory) { - PsiExpression replacement = factory.createExpressionFromText("java.util.Arrays.stream("+myExpression.getText() + ")", myExpression); String operation = "flatMap"; PsiType type = myExpression.getType(); if(type instanceof PsiArrayType) { @@ -855,7 +887,13 @@ public class StreamApiMigrationInspection extends BaseJavaBatchLocalInspectionTo } } } - return "." + operation + "(" + LambdaUtil.createLambda(myVariable, replacement) + ")"; + return "." + operation + "(" + myVariable.getName() + " -> " + getStreamExpression() + ")"; + } + + @Override + @NotNull + String getStreamExpression() { + return "java.util.Arrays.stream("+myExpression.getText() + ")"; } } @@ -874,8 +912,12 @@ public class StreamApiMigrationInspection extends BaseJavaBatchLocalInspectionTo } private void flatten() { - while(myStatements.length == 1 && myStatements[0] instanceof PsiBlockStatement) { - myStatements = ((PsiBlockStatement)myStatements[0]).getCodeBlock().getStatements(); + while(true) { + if(myStatements.length == 1 && myStatements[0] instanceof PsiBlockStatement) { + myStatements = ((PsiBlockStatement)myStatements[0]).getCodeBlock().getStatements(); + } else if(myStatements.length == 1 && myStatements[0] instanceof PsiLabeledStatement) { + myStatements = new PsiStatement[] {((PsiLabeledStatement)myStatements[0]).getStatement()}; + } else break; } } @@ -914,6 +956,40 @@ public class StreamApiMigrationInspection extends BaseJavaBatchLocalInspectionTo return getSingleExpression(PsiMethodCallExpression.class); } + @Nullable + private FilterOp extractFilter() { + if(getSingleStatement() instanceof PsiIfStatement) { + PsiIfStatement ifStatement = (PsiIfStatement)getSingleStatement(); + if(ifStatement.getElseBranch() == null && ifStatement.getCondition() != null) { + replaceWith(ifStatement.getThenBranch()); + return new FilterOp(ifStatement.getCondition(), myVariable, false); + } + } + if(myStatements.length >= 1) { + PsiStatement first = myStatements[0]; + // extract filter with negation + if(first instanceof PsiIfStatement) { + PsiIfStatement ifStatement = (PsiIfStatement)first; + if(ifStatement.getCondition() == null) return null; + PsiStatement branch = ifStatement.getThenBranch(); + if(branch instanceof PsiBlockStatement) { + PsiStatement[] statements = ((PsiBlockStatement)branch).getCodeBlock().getStatements(); + if(statements.length == 1) + branch = statements[0]; + } + if(!(branch instanceof PsiContinueStatement) || ((PsiContinueStatement)branch).getLabelIdentifier() != null) return null; + if(ifStatement.getElseBranch() != null) { + myStatements[0] = ifStatement.getElseBranch(); + } else { + myStatements = Arrays.copyOfRange(myStatements, 1, myStatements.length); + } + flatten(); + return new FilterOp(ifStatement.getCondition(), myVariable, true); + } + } + return null; + } + /** * If possible, extract single intermediate stream operation from this * {@code TerminalBlock} changing the TerminalBlock itself to exclude this operation @@ -922,14 +998,8 @@ public class StreamApiMigrationInspection extends BaseJavaBatchLocalInspectionTo */ @Nullable Operation extractOperation() { - // extract filter - if(getSingleStatement() instanceof PsiIfStatement) { - PsiIfStatement ifStatement = (PsiIfStatement)getSingleStatement(); - if(ifStatement.getElseBranch() == null && ifStatement.getCondition() != null) { - replaceWith(ifStatement.getThenBranch()); - return new FilterOp(ifStatement.getCondition(), myVariable, false); - } - } + FilterOp filter = extractFilter(); + if(filter != null) return filter; // extract flatMap if(getSingleStatement() instanceof PsiForeachStatement) { // flatMapping of primitive variable is not supported yet @@ -939,23 +1009,40 @@ public class StreamApiMigrationInspection extends BaseJavaBatchLocalInspectionTo final PsiStatement body = foreachStatement.getBody(); if (iteratedValue != null && body != null) { final PsiType iteratedValueType = iteratedValue.getType(); - Operation op = null; + FlatMapOp op = null; if(iteratedValueType instanceof PsiArrayType) { if (!isSupported(((PsiArrayType)iteratedValueType).getComponentType())) return null; - op = new ArrayFlatMapOp(iteratedValue, myVariable); + op = new ArrayFlatMapOp(iteratedValue, myVariable, foreachStatement); } else { final PsiClass iteratorClass = PsiUtil.resolveClassInClassTypeOnly(iteratedValueType); final PsiClass collectionClass = JavaPsiFacade.getInstance(body.getProject()) .findClass(CommonClassNames.JAVA_UTIL_COLLECTION, foreachStatement.getResolveScope()); if (collectionClass != null && InheritanceUtil.isInheritorOrSelf(iteratorClass, collectionClass, true)) { - op = new FlatMapOp(iteratedValue, myVariable); + op = new FlatMapOp(iteratedValue, myVariable, foreachStatement); } } - if(op != null && ReferencesSearch.search(myVariable, new LocalSearchScope(body)).findFirst() == null) { - myVariable = foreachStatement.getIterationParameter(); - replaceWith(body); - return op; + if(op != null) { + if(ReferencesSearch.search(myVariable, new LocalSearchScope(body)).findFirst() == null) { + myVariable = foreachStatement.getIterationParameter(); + replaceWith(body); + return op; + } else { + PsiStatement[] statements = myStatements; + myVariable = foreachStatement.getIterationParameter(); + replaceWith(body); + FilterOp nextFilter = extractFilter(); + myVariable = op.myVariable; + if(nextFilter != null) { + PsiStatement lastStatement = myStatements[myStatements.length - 1]; + if(lastStatement instanceof PsiBreakStatement && op.breaksMe((PsiBreakStatement)lastStatement) && + ReferencesSearch.search(nextFilter.myVariable, new LocalSearchScope(myStatements)).findFirst() == null) { + myStatements = Arrays.copyOfRange(myStatements, 0, myStatements.length-1); + return new CompoundFilterOp(nextFilter, op); + } + } + myStatements = statements; + } } } } @@ -985,25 +1072,6 @@ public class StreamApiMigrationInspection extends BaseJavaBatchLocalInspectionTo } } } - // extract filter with negation - if(first instanceof PsiIfStatement) { - PsiIfStatement ifStatement = (PsiIfStatement)first; - if(ifStatement.getCondition() == null) return null; - PsiStatement branch = ifStatement.getThenBranch(); - if(branch instanceof PsiBlockStatement) { - PsiStatement[] statements = ((PsiBlockStatement)branch).getCodeBlock().getStatements(); - if(statements.length == 1) - branch = statements[0]; - } - if(!(branch instanceof PsiContinueStatement) || ((PsiContinueStatement)branch).getLabelIdentifier() != null) return null; - if(ifStatement.getElseBranch() != null) { - myStatements[0] = ifStatement.getElseBranch(); - } else { - myStatements = Arrays.copyOfRange(myStatements, 1, myStatements.length); - } - flatten(); - return new FilterOp(ifStatement.getCondition(), myVariable, true); - } } return null; } diff --git a/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/afterNestedAnyMatch.java b/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/afterNestedAnyMatch.java new file mode 100644 index 000000000000..a86881f91374 --- /dev/null +++ b/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/afterNestedAnyMatch.java @@ -0,0 +1,9 @@ +// "Replace with toArray" "true" + +import java.util.*; + +public class Main { + public Integer[] testNestedAnyMatch(List> data) { + return data.stream().filter(list -> list.stream().anyMatch(str -> !str.isEmpty())).map(List::size).toArray(Integer[]::new); + } +} \ No newline at end of file diff --git a/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/beforeNestedAnyMatch.java b/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/beforeNestedAnyMatch.java new file mode 100644 index 000000000000..da40ea0faf87 --- /dev/null +++ b/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/beforeNestedAnyMatch.java @@ -0,0 +1,18 @@ +// "Replace with toArray" "true" + +import java.util.*; + +public class Main { + public Integer[] testNestedAnyMatch(List> data) { + List result = new ArrayList<>(); + for (List list : data) { + for (String str : list) { + if (!str.isEmpty()) { + result.add(list.size()); + break; + } + } + } + return result.toArray(new Integer[result.size()]); + } +} \ No newline at end of file diff --git a/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/beforeNestedAnyMatchFailed.java b/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/beforeNestedAnyMatchFailed.java new file mode 100644 index 000000000000..8d2988730443 --- /dev/null +++ b/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/quickFix/streamApiMigration/beforeNestedAnyMatchFailed.java @@ -0,0 +1,18 @@ +// "Replace with toArray" "false" + +import java.util.*; + +public class Main { + public Integer[] testNestedAnyMatch(List> data) { + List result = new ArrayList<>(); + for (List list : data) { + for (String str : list) { + if (!str.isEmpty()) { + result.add(str.length()); + break; + } + } + } + return result.toArray(new Integer[result.size()]); + } +} \ No newline at end of file