StreamToLoopInspection: support block void lambdas (and returns in forEach not inside nested loop)

This commit is contained in:
Tagir Valeev
2017-01-31 16:29:03 +03:00
parent c15f721206
commit d8836c315d
6 changed files with 222 additions and 27 deletions
@@ -24,6 +24,7 @@ import com.intellij.psi.codeStyle.SuggestedNameInfo;
import com.intellij.psi.codeStyle.VariableKind;
import com.intellij.psi.search.LocalSearchScope;
import com.intellij.psi.search.searches.ReferencesSearch;
import com.intellij.psi.util.PsiTreeUtil;
import com.intellij.refactoring.util.LambdaRefactoringUtil;
import com.intellij.util.ArrayUtil;
import com.siyeh.ig.psiutils.ExpressionUtils;
@@ -35,10 +36,7 @@ import org.jetbrains.annotations.NotNull;
import org.jetbrains.annotations.Nullable;
import java.text.MessageFormat;
import java.util.Arrays;
import java.util.Collections;
import java.util.List;
import java.util.Locale;
import java.util.*;
import java.util.function.Consumer;
/**
@@ -63,6 +61,10 @@ abstract class FunctionHelper {
return getExpression().getText();
}
String getStatementText() {
return getText() + ";\n";
}
abstract PsiExpression getExpression();
/**
@@ -133,6 +135,12 @@ abstract class FunctionHelper {
@Contract("null, _ -> null")
@Nullable
static FunctionHelper create(PsiExpression expression, int paramCount) {
return create(expression, paramCount, false);
}
@Contract("null, _, _ -> null")
@Nullable
static FunctionHelper create(PsiExpression expression, int paramCount, boolean allowReturns) {
if(expression == null) return null;
PsiType type = expression instanceof PsiFunctionalExpression
? ((PsiFunctionalExpression)expression).getFunctionalInterfaceType()
@@ -149,9 +157,23 @@ abstract class FunctionHelper {
PsiParameterList list = lambda.getParameterList();
if (list.getParametersCount() != paramCount) return null;
String[] parameters = StreamEx.of(list.getParameters()).map(PsiVariable::getName).toArray(String[]::new);
PsiExpression body = LambdaUtil.extractSingleExpressionFromBody(lambda.getBody());
if (body == null) return null;
return new LambdaFunctionHelper(returnType, body, parameters);
PsiElement body = lambda.getBody();
PsiExpression lambdaExpression = LambdaUtil.extractSingleExpressionFromBody(body);
if (lambdaExpression == null) {
if (PsiType.VOID.equals(returnType) && body instanceof PsiCodeBlock) {
List<PsiReturnStatement> returns = getReturns(body);
if (!allowReturns && !returns.isEmpty()) return null;
// Return inside loop is not supported yet
for (PsiReturnStatement ret : returns) {
if (PsiTreeUtil.getParentOfType(ret, PsiLoopStatement.class, true, PsiLambdaExpression.class) != null) {
return null;
}
}
return new VoidBlockLambdaFunctionHelper((PsiCodeBlock)body, parameters);
}
return null;
}
return new LambdaFunctionHelper(returnType, lambdaExpression, parameters);
}
if (expression instanceof PsiMethodReferenceExpression) {
PsiMethodReferenceExpression methodRef = (PsiMethodReferenceExpression)expression;
@@ -262,8 +284,8 @@ abstract class FunctionHelper {
};
}
static boolean hasVarReference(PsiExpression expression, String name, StreamToLoopReplacementContext context) {
PsiLambdaExpression lambda = (PsiLambdaExpression)context.createExpression(name+"->"+expression.getText());
static boolean hasVarReference(PsiElement expressionOrCodeBlock, String name, StreamToLoopReplacementContext context) {
PsiLambdaExpression lambda = (PsiLambdaExpression)context.createExpression(name + "->" + expressionOrCodeBlock.getText());
PsiParameter var = lambda.getParameterList().getParameters()[0];
PsiElement body = lambda.getBody();
LOG.assertTrue(body != null);
@@ -277,19 +299,19 @@ abstract class FunctionHelper {
* If the replacement is a new name to the variable, the caller must take care that this new name was not used before.
* </p>
*
* @param expression an expression to search-and-replace references inside
* @param expressionOrCodeBlock an expression or code block to search-and-replace references inside
* @param name a reference name to replace
* @param replacement a replacement expression (new name or literal)
* @param context context
* @return resulting expression (might be the same as input expression)
*/
@NotNull
static PsiExpression replaceVarReference(@NotNull PsiExpression expression,
String name,
String replacement,
StreamToLoopReplacementContext context) {
if(name.equals(replacement)) return expression;
PsiLambdaExpression lambda = (PsiLambdaExpression)context.createExpression(name+"->"+expression.getText());
static <T extends PsiElement> T replaceVarReference(@NotNull T expressionOrCodeBlock,
String name,
String replacement,
StreamToLoopReplacementContext context) {
if (name.equals(replacement)) return expressionOrCodeBlock;
PsiLambdaExpression lambda = (PsiLambdaExpression)context.createExpression(name + "->" + expressionOrCodeBlock.getText());
PsiParameter var = lambda.getParameterList().getParameters()[0];
PsiElement body = lambda.getBody();
LOG.assertTrue(body != null);
@@ -297,7 +319,27 @@ abstract class FunctionHelper {
for (PsiReference ref : ReferencesSearch.search(var, new LocalSearchScope(body)).findAll()) {
ref.getElement().replace(replacementExpression);
}
return (PsiExpression)lambda.getBody();
//noinspection unchecked
return (T)lambda.getBody();
}
@NotNull
private static List<PsiReturnStatement> getReturns(PsiElement body) {
List<PsiReturnStatement> returns = new ArrayList<>();
body.accept(new JavaRecursiveElementWalkingVisitor() {
@Override
public void visitClass(@NotNull PsiClass psiClass) { }
@Override
public void visitLambdaExpression(PsiLambdaExpression expression) { }
@Override
public void visitReturnStatement(@NotNull PsiReturnStatement returnStatement) {
super.visitReturnStatement(returnStatement);
returns.add(returnStatement);
}
});
return returns;
}
private static class MethodReferenceFunctionHelper extends FunctionHelper {
@@ -535,10 +577,10 @@ abstract class FunctionHelper {
}
private static class LambdaFunctionHelper extends FunctionHelper {
private String[] myParameters;
private PsiExpression myBody;
String[] myParameters;
PsiElement myBody;
LambdaFunctionHelper(PsiType returnType, PsiExpression body, String[] parameters) {
LambdaFunctionHelper(PsiType returnType, PsiElement body, String[] parameters) {
super(returnType);
myParameters = parameters;
myBody = body;
@@ -551,7 +593,8 @@ abstract class FunctionHelper {
}
PsiExpression getExpression() {
return myBody;
// Usage logic presume that this method is called only if myBody is PsiExpression
return (PsiExpression)myBody;
}
void transform(StreamToLoopReplacementContext context, String... argumentValues) {
@@ -592,4 +635,24 @@ abstract class FunctionHelper {
suggestFromExpression(var, project, expr);
}
}
private static class VoidBlockLambdaFunctionHelper extends LambdaFunctionHelper {
VoidBlockLambdaFunctionHelper(PsiCodeBlock body, String[] parameters) {
super(PsiType.VOID, body, parameters);
}
@Override
String getStatementText() {
PsiElement[] children = myBody.getChildren();
// Keep everything except braces
return StreamEx.of(children, 1, children.length - 1).map(PsiElement::getText).joining().trim();
}
void transform(StreamToLoopReplacementContext context, String... argumentValues) {
super.transform(context, argumentValues);
List<PsiReturnStatement> returns = getReturns(myBody);
String continueStatement = "continue;";
returns.forEach(ret -> ret.replace(context.createStatement(continueStatement)));
}
}
}
@@ -145,7 +145,7 @@ abstract class Operation {
@Override
String wrap(StreamVariable outVar, String code, StreamToLoopReplacementContext context) {
return myFn.getText() + ";\n" + code;
return myFn.getStatementText() + code;
}
}
@@ -194,7 +194,7 @@ public class StreamToLoopInspection extends BaseJavaBatchLocalInspectionTool {
if (InheritanceUtil.isInheritor(type, CommonClassNames.JAVA_UTIL_STREAM_BASE_STREAM)) return null;
PsiExpression[] args = terminalCall.getArgumentList().getExpressions();
if (args.length != 1) return null;
FunctionHelper fn = FunctionHelper.create(args[0], 1);
FunctionHelper fn = FunctionHelper.create(args[0], 1, true);
if (fn == null) return null;
PsiType elementType = PsiUtil.substituteTypeParameter(type, CommonClassNames.JAVA_LANG_ITERABLE, 0, false);
if(elementType == null) return null;
@@ -606,7 +606,7 @@ public class StreamToLoopInspection extends BaseJavaBatchLocalInspectionTool {
if(fn != null) {
fn.transform(this, ((ConditionalExpression.Optional)conditionalExpression).unwrap("").getTrueBranch());
myPlaceholder = call.getParent();
return fn.getText() + ";\n" + getBreakStatement();
return fn.getStatementText() + getBreakStatement();
}
}
}
@@ -701,6 +701,10 @@ public class StreamToLoopInspection extends BaseJavaBatchLocalInspectionTool {
return myFactory.createExpressionFromText(text, myStatement);
}
public PsiStatement createStatement(String text) {
return myFactory.createStatementFromText(text, myStatement);
}
public PsiType createType(String text) {
return myFactory.createTypeFromText(text, myStatement);
}
@@ -68,7 +68,7 @@ abstract class TerminalOperation extends Operation {
@NotNull PsiType elementType, @NotNull PsiType resultType, boolean isVoid) {
if(isVoid) {
if ((name.equals("forEach") || name.equals("forEachOrdered")) && args.length == 1) {
FunctionHelper fn = FunctionHelper.create(args[0], 1);
FunctionHelper fn = FunctionHelper.create(args[0], 1, true);
return fn == null ? null : new ForEachTerminalOperation(fn);
}
return null;
@@ -425,7 +425,7 @@ abstract class TerminalOperation extends Operation {
String candidate = mySupplier.suggestFinalOutputNames(context, myAccumulator.getParameterName(0), "acc").get(0);
String acc = context.declareResult(candidate, mySupplier.getResultType(), mySupplier.getText(), ResultKind.FINAL);
myAccumulator.transform(context, acc, inVar.getName());
return myAccumulator.getText()+";\n";
return myAccumulator.getStatementText();
}
}
@@ -1067,7 +1067,7 @@ abstract class TerminalOperation extends Operation {
@Override
String generate(StreamVariable inVar, StreamToLoopReplacementContext context) {
myFn.transform(context, inVar.getName());
return myFn.getText()+";\n";
return myFn.getStatementText();
}
}
@@ -0,0 +1,77 @@
// "Fix all 'Stream API call chain can be replaced with loop' problems in file" "true"
import java.util.Arrays;
import java.util.Collections;
import java.util.List;
import java.util.Optional;
import java.util.function.Function;
import java.util.stream.Collectors;
import java.util.stream.IntStream;
import java.util.stream.Stream;
public class Main {
public static void main(String[] args) {
List<String> list = Arrays.asList("a", "b", "c", "d");
for (String o : list) {
if (!o.isEmpty()) {
System.out.println("Peek: " + o);
System.out.println("Peek2: " + o);
String n = "";
System.out.println(o);
if (o.equals("c")) continue;
System.out.println(o + "!!!" + n);
new Runnable() {
String e = "x";
public void run() {
System.out.println(e);
for (int i = 0; i < 10; i++) {
if (e.length() == i) return;
}
}
}.run();
}
}
for (String s : list) {
if ("b".equals(s)) {
System.out.println("Found:");
System.out.println(s);
break;
}
}
Optional<String> found = Optional.empty();
for (String qq : list) {
if ("b".equals(qq)) {
found = Optional.of(qq);
break;
}
}
found.ifPresent(str -> {
System.out.println("Found:");
if(str.isEmpty()) return; // return inside ifPresent is not supported
System.out.println(str);
});
StringBuilder res = new StringBuilder();
for (String str : list) {
if (str != null) {
str = "[" + str + "]";
res.append(str);
}
}
System.out.println(res);
long count = 0L;
for (String n : list) {
if (!"a".equals(n)) {
for (int i = 0; i < 3; i++) {
System.out.println("In flatmap idx: " + i);
System.out.println("In flatmap: " + n);
count++;
}
}
}
System.out.println(count);
}
}
@@ -0,0 +1,51 @@
// "Fix all 'Stream API call chain can be replaced with loop' problems in file" "true"
import java.util.Arrays;
import java.util.Collections;
import java.util.List;
import java.util.function.Function;
import java.util.stream.Collectors;
import java.util.stream.IntStream;
import java.util.stream.Stream;
public class Main {
public static void main(String[] args) {
List<String> list = Arrays.asList("a", "b", "c", "d");
list.stream().filter(n -> !n.isEmpty()).peek(o -> {
System.out.println("Peek: "+o);
System.out.println("Peek2: "+o);
}).for<caret>Each(e -> {
String n = "";
System.out.println(e);
if (e.equals("c")) return;
System.out.println(e + "!!!" + n);
new Runnable() {String e = "x"; public void run() {
System.out.println(e);
for(int i=0; i<10; i++) {
if(e.length() == i) return;
}
}}.run();
});
list.stream().filter("b"::equals).findFirst().ifPresent(str -> {
System.out.println("Found:");
System.out.println(str);
});
list.stream().filter(qq -> "b".equals(qq)).findFirst().ifPresent(str -> {
System.out.println("Found:");
if(str.isEmpty()) return; // return inside ifPresent is not supported
System.out.println(str);
});
StringBuilder res = list.stream().filter(str -> str != null).collect(StringBuilder::new, (sb, s) -> {
s = "[" + s + "]";
sb.append(s);
}, StringBuilder::append);
System.out.println(res);
System.out.println(list.stream().filter(n -> !"a".equals(n)).flatMapToInt(l -> IntStream.range(0, 3).peek(n -> {
System.out.println("In flatmap idx: "+n);
System.out.println("In flatmap: "+l);
})).count());
}
}