Stream API migration: support nested anyMatch

This commit is contained in:
Tagir Valeev
2016-10-03 12:43:44 +07:00
parent 086ac8f4f2
commit b53646340b
4 changed files with 158 additions and 45 deletions
@@ -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;
}
@@ -0,0 +1,9 @@
// "Replace with toArray" "true"
import java.util.*;
public class Main {
public Integer[] testNestedAnyMatch(List<List<String>> data) {
return data.stream().filter(list -> list.stream().anyMatch(str -> !str.isEmpty())).map(List::size).toArray(Integer[]::new);
}
}
@@ -0,0 +1,18 @@
// "Replace with toArray" "true"
import java.util.*;
public class Main {
public Integer[] testNestedAnyMatch(List<List<String>> data) {
List<Integer> result = new ArrayList<>();
for (List<String> list : d<caret>ata) {
for (String str : list) {
if (!str.isEmpty()) {
result.add(list.size());
break;
}
}
}
return result.toArray(new Integer[result.size()]);
}
}
@@ -0,0 +1,18 @@
// "Replace with toArray" "false"
import java.util.*;
public class Main {
public Integer[] testNestedAnyMatch(List<List<String>> data) {
List<Integer> result = new ArrayList<>();
for (List<String> list : d<caret>ata) {
for (String str : list) {
if (!str.isEmpty()) {
result.add(str.length());
break;
}
}
}
return result.toArray(new Integer[result.size()]);
}
}