diff --git a/java/java-analysis-impl/src/com/intellij/codeInspection/java18api/Java8ReplaceMapGetInspection.java b/java/java-analysis-impl/src/com/intellij/codeInspection/java18api/Java8ReplaceMapGetInspection.java index ef8f4986aa16..6bb987846c29 100644 --- a/java/java-analysis-impl/src/com/intellij/codeInspection/java18api/Java8ReplaceMapGetInspection.java +++ b/java/java-analysis-impl/src/com/intellij/codeInspection/java18api/Java8ReplaceMapGetInspection.java @@ -15,15 +15,14 @@ */ package com.intellij.codeInspection.java18api; -import com.intellij.codeInsight.ExceptionUtil; import com.intellij.codeInsight.FileModificationService; import com.intellij.codeInsight.daemon.QuickFixBundle; -import com.intellij.codeInsight.daemon.impl.analysis.HighlightControlFlowUtil; import com.intellij.codeInspection.BaseJavaBatchLocalInspectionTool; import com.intellij.codeInspection.LocalQuickFix; import com.intellij.codeInspection.ProblemDescriptor; import com.intellij.codeInspection.ProblemsHolder; import com.intellij.codeInspection.ui.MultipleCheckboxOptionsPanel; +import com.intellij.codeInspection.util.LambdaGenerationUtil; import com.intellij.openapi.project.Project; import com.intellij.psi.*; import com.intellij.psi.codeStyle.CodeStyleManager; @@ -93,7 +92,7 @@ public class Java8ReplaceMapGetInspection extends BaseJavaBatchLocalInspectionTo holder.registerProblem(condition, QuickFixBundle.message("java.8.replace.map.get.inspection.description"), new ReplaceGetNullCheck("getOrDefault")); } - } else if(thenBranch instanceof PsiBlockStatement) { + } else { /* value = map.get(key); if(value == null) { @@ -103,37 +102,9 @@ public class Java8ReplaceMapGetInspection extends BaseJavaBatchLocalInspectionTo */ if (!mySuggestMapComputeIfAbsent) return; PsiExpression key = getArguments[0]; - PsiStatement[] statements = ((PsiBlockStatement)thenBranch).getCodeBlock().getStatements(); - if(statements.length != 2) return; - assignment = ExpressionUtils.getAssignment(statements[0]); - if(assignment == null) return; - PsiExpression lambdaCandidate = assignment.getRExpression(); - if (lambdaCandidate == null || - !equivalence.expressionsAreEquivalent(assignment.getLExpression(), value) || - !(statements[1] instanceof PsiExpressionStatement)) { - return; - } - PsiExpression expression = ((PsiExpressionStatement)statements[1]).getExpression(); - if(!(expression instanceof PsiMethodCallExpression)) return; - PsiMethodCallExpression putCall = (PsiMethodCallExpression)expression; - if(!Java8CollectionsApiInspection.isJavaUtilMapMethodWithName(putCall, "put")) return; - PsiExpression[] putArguments = putCall.getArgumentList().getExpressions(); - if (putArguments.length != 2 || - !equivalence.expressionsAreEquivalent(putCall.getMethodExpression().getQualifierExpression(), - getCall.getMethodExpression().getQualifierExpression()) || - !equivalence.expressionsAreEquivalent(key, putArguments[0]) || - !equivalence.expressionsAreEquivalent(value, putArguments[1])) { - return; - } - if(!ExceptionUtil.getThrownCheckedExceptions(lambdaCandidate).isEmpty()) return; - if(!PsiTreeUtil.processElements(lambdaCandidate, e -> { - if(!(e instanceof PsiReferenceExpression)) return true; - PsiElement element = ((PsiReferenceExpression)e).resolve(); - if(!(element instanceof PsiVariable)) return true; - return HighlightControlFlowUtil.isEffectivelyFinal((PsiVariable)element, lambdaCandidate, null); - })) { - return; - } + PsiExpression mapExpression = getCall.getMethodExpression().getQualifierExpression(); + PsiExpression lambdaCandidate = extractLambdaCandidate(thenBranch, mapExpression, key, value); + if (lambdaCandidate == null) return; holder.registerProblem(condition, QuickFixBundle.message("java.8.replace.map.get.inspection.description"), new ReplaceGetNullCheck("computeIfAbsent")); } @@ -141,6 +112,55 @@ public class Java8ReplaceMapGetInspection extends BaseJavaBatchLocalInspectionTo }; } + @Nullable + static PsiExpression extractLambdaCandidate(PsiStatement statement, PsiExpression mapExpression, + PsiExpression keyExpression, PsiReferenceExpression valueExpression) { + EquivalenceChecker equivalence = EquivalenceChecker.getCanonicalPsiEquivalence(); + PsiAssignmentExpression assignment; + PsiMethodCallExpression putCall = extractPutCall(statement); + if(putCall != null) { + // like map.put(key, val = new ArrayList<>()); + PsiExpression[] putArguments = putCall.getArgumentList().getExpressions(); + if (putArguments.length != 2 || + !equivalence.expressionsAreEquivalent(putCall.getMethodExpression().getQualifierExpression(), mapExpression) || + !equivalence.expressionsAreEquivalent(keyExpression, putArguments[0])) { + return null; + } + assignment = ExpressionUtils.getAssignment(putArguments[1]); + } + else { + if (!(statement instanceof PsiBlockStatement)) return null; + // like val = new ArrayList<>(); map.put(key, val); + PsiStatement[] statements = ((PsiBlockStatement)statement).getCodeBlock().getStatements(); + putCall = extractPutCall(statements[1]); + if (putCall == null) return null; + PsiExpression[] putArguments = putCall.getArgumentList().getExpressions(); + if (putArguments.length != 2 || + !equivalence.expressionsAreEquivalent(putCall.getMethodExpression().getQualifierExpression(), mapExpression) || + !equivalence.expressionsAreEquivalent(keyExpression, putArguments[0]) || + !equivalence.expressionsAreEquivalent(valueExpression, putArguments[1])) { + return null; + } + assignment = ExpressionUtils.getAssignment(statements[0]); + } + if (assignment == null) return null; + PsiExpression lambdaCandidate = assignment.getRExpression(); + if (lambdaCandidate == null || !equivalence.expressionsAreEquivalent(assignment.getLExpression(), valueExpression)) return null; + if (!LambdaGenerationUtil.canBeUncheckedLambda(lambdaCandidate)) return null; + return lambdaCandidate; + } + + @Contract("null -> null") + @Nullable + private static PsiMethodCallExpression extractPutCall(PsiStatement statement) { + if(!(statement instanceof PsiExpressionStatement)) return null; + PsiExpression expression = ((PsiExpressionStatement)statement).getExpression(); + if (!(expression instanceof PsiMethodCallExpression)) return null; + PsiMethodCallExpression putCall = (PsiMethodCallExpression)expression; + if (!Java8CollectionsApiInspection.isJavaUtilMapMethodWithName(putCall, "put")) return null; + return putCall; + } + @Nullable private static PsiReferenceExpression getReferenceComparedWithNull(PsiExpression condition) { if(!(condition instanceof PsiBinaryExpression)) return null; @@ -223,27 +243,22 @@ public class Java8ReplaceMapGetInspection extends BaseJavaBatchLocalInspectionTo Collection comments = ContainerUtil.map(PsiTreeUtil.findChildrenOfType(ifStatement, PsiComment.class), comment -> (PsiComment)comment.copy()); PsiElementFactory factory = JavaPsiFacade.getElementFactory(project); - if(thenBranch instanceof PsiExpressionStatement) { - PsiExpression expression = ((PsiExpressionStatement)thenBranch).getExpression(); - if (!(expression instanceof PsiAssignmentExpression)) return; - PsiExpression defaultValue = ((PsiAssignmentExpression)expression).getRExpression(); + PsiAssignmentExpression assignment = ExpressionUtils.getAssignment(thenBranch); + if (!FileModificationService.getInstance().preparePsiElementForWrite(element)) return; + if(assignment != null) { + PsiExpression defaultValue = assignment.getRExpression(); if (!ExpressionUtils.isSimpleExpression(defaultValue)) return; - if (!FileModificationService.getInstance().preparePsiElementForWrite(element)) return; nameElement.replace(factory.createIdentifier("getOrDefault")); getCall.getArgumentList().add(defaultValue); - } else if(thenBranch instanceof PsiBlockStatement) { - PsiStatement[] statements = ((PsiBlockStatement)thenBranch).getCodeBlock().getStatements(); - if(statements.length != 2) return; - PsiAssignmentExpression assignment = ExpressionUtils.getAssignment(statements[0]); - if(assignment == null) return; - PsiExpression lambdaCandidate = assignment.getRExpression(); - if(lambdaCandidate == null) return; - if (!FileModificationService.getInstance().preparePsiElementForWrite(element)) return; + } else { + PsiExpression lambdaCandidate = + extractLambdaCandidate(thenBranch, getCall.getMethodExpression().getQualifierExpression(), args[0], value); + if (lambdaCandidate == null) return; nameElement.replace(factory.createIdentifier("computeIfAbsent")); String varName = JavaCodeStyleManager.getInstance(project).suggestUniqueVariableName("k", lambdaCandidate, true); PsiExpression lambda = factory.createExpressionFromText(varName + " -> " + lambdaCandidate.getText(), lambdaCandidate); getCall.getArgumentList().add(lambda); - } else return; + } ifStatement.delete(); CodeStyleManager.getInstance(project).reformat(statement); comments.forEach(comment -> statement.getParent().addBefore(comment, statement)); diff --git a/java/java-tests/testData/inspection/java8MapGet/afterComputeIfAbsentSingleLine.java b/java/java-tests/testData/inspection/java8MapGet/afterComputeIfAbsentSingleLine.java new file mode 100644 index 000000000000..a6caf3d589d9 --- /dev/null +++ b/java/java-tests/testData/inspection/java8MapGet/afterComputeIfAbsentSingleLine.java @@ -0,0 +1,11 @@ +// "Replace with 'computeIfAbsent' method call" "true" +import java.util.ArrayList; +import java.util.List; +import java.util.Map; + +public class Main { + public void testMap(Map> map, String key, String value) { + List list = map.computeIfAbsent(key, k -> new ArrayList<>()); + list.add(value); + } +} \ No newline at end of file diff --git a/java/java-tests/testData/inspection/java8MapGet/beforeComputeIfAbsentSingleLine.java b/java/java-tests/testData/inspection/java8MapGet/beforeComputeIfAbsentSingleLine.java new file mode 100644 index 000000000000..96071e187676 --- /dev/null +++ b/java/java-tests/testData/inspection/java8MapGet/beforeComputeIfAbsentSingleLine.java @@ -0,0 +1,14 @@ +// "Replace with 'computeIfAbsent' method call" "true" +import java.util.ArrayList; +import java.util.List; +import java.util.Map; + +public class Main { + public void testMap(Map> map, String key, String value) { + List list = map.get(key); + if(list == null) { + map.put(key, list = new ArrayList<>()); + } + list.add(value); + } +} \ No newline at end of file diff --git a/java/java-tests/testData/inspection/java8MapGet/beforeComputeIfAbsentSingleLineWrongVar.java b/java/java-tests/testData/inspection/java8MapGet/beforeComputeIfAbsentSingleLineWrongVar.java new file mode 100644 index 000000000000..d37fc020c386 --- /dev/null +++ b/java/java-tests/testData/inspection/java8MapGet/beforeComputeIfAbsentSingleLineWrongVar.java @@ -0,0 +1,15 @@ +// "Replace with 'computeIfAbsent' method call" "false" +import java.util.ArrayList; +import java.util.List; +import java.util.Map; + +public class Main { + public void testMap(Map> map, String key, String value) { + List list2 = null; + List list = map.get(key); + if(list == null) { + map.put(key, list2 = new ArrayList<>()); + } + list.add(value); + } +} \ No newline at end of file