IDEA-200789: Suggest to use Map.compute() instead of Map.get() + Map.put()

(cherry picked from commit 484148216a2c29613043695ef3f69238118946a9)

IJ-CR-117655

GitOrigin-RevId: 8a95550370e3ee03ec02420254fe2e2287eedf51
This commit is contained in:
Daria Nikolskaia
2023-10-30 20:55:19 +00:00
committed by intellij-monorepo-bot
parent 0e2bcac8fc
commit dd7f92d48a
8 changed files with 176 additions and 2 deletions
@@ -4,7 +4,10 @@ package com.intellij.codeInspection.java18api;
import com.intellij.codeInsight.Nullability;
import com.intellij.codeInsight.PsiEquivalenceUtil;
import com.intellij.codeInsight.daemon.QuickFixBundle;
import com.intellij.codeInspection.*;
import com.intellij.codeInspection.AbstractBaseJavaLocalInspectionTool;
import com.intellij.codeInspection.LambdaCanBeMethodReferenceInspection;
import com.intellij.codeInspection.ProblemHighlightType;
import com.intellij.codeInspection.ProblemsHolder;
import com.intellij.codeInspection.dataFlow.NullabilityUtil;
import com.intellij.codeInspection.options.OptPane;
import com.intellij.codeInspection.util.LambdaGenerationUtil;
@@ -35,6 +38,7 @@ import java.util.List;
import static com.intellij.codeInspection.options.OptPane.checkbox;
import static com.intellij.codeInspection.options.OptPane.pane;
import static com.siyeh.ig.psiutils.EquivalenceChecker.getCanonicalPsiEquivalence;
import static com.siyeh.ig.psiutils.Java8MigrationUtils.*;
import static com.siyeh.ig.psiutils.Java8MigrationUtils.MapCheckCondition.fromConditional;
@@ -87,6 +91,49 @@ public class Java8MapApiInspection extends AbstractBaseJavaLocalInspectionTool {
PsiExpression noneBranch = condition.getNoneBranch(expression.getThenExpression(), expression.getElseExpression());
processGetPut(condition, existsBranch, existsBranch, noneBranch);
}
@Override
public void visitLocalVariable(@NotNull PsiLocalVariable variable) {
PsiExpression expression = variable.getInitializer();
PsiMethodCallExpression getCall = extractMapMethodCall(expression, "get");
if (getCall == null) return;
List<PsiReferenceExpression> references = VariableAccessUtils
.getVariableReferences(variable, PsiTreeUtil.getParentOfType(variable, PsiCodeBlock.class));
if (references.isEmpty()) return;
PsiMethodCallExpression commonPutCall = findPutMethodParent(references.get(0).getElement());
if (commonPutCall == null || !isCommonPutCallForAllReferences(references, commonPutCall)) return;
PsiExpression getCallQualifierExpression = getCall.getMethodExpression().getQualifierExpression();
PsiExpression putCallQualifierExpression = commonPutCall.getMethodExpression().getQualifierExpression();
EquivalenceChecker equivalenceChecker = getCanonicalPsiEquivalence();
if (! equivalenceChecker.expressionsAreEquivalent(getCallQualifierExpression, putCallQualifierExpression)) return;
PsiStatement variableDeclarationStatement = PsiTreeUtil.getParentOfType(variable, PsiDeclarationStatement.class);
PsiElement nextSibling = PsiTreeUtil.skipWhitespacesAndCommentsForward(variableDeclarationStatement);
if (! (nextSibling instanceof PsiStatement)) return;
PsiExpressionStatement putCallStatement = ObjectUtils.tryCast(commonPutCall.getParent(), PsiExpressionStatement.class);
if (nextSibling != putCallStatement) return;
PsiExpression[] getArgs = getCall.getArgumentList().getExpressions();
PsiExpression[] putArgs = commonPutCall.getArgumentList().getExpressions();
if (getArgs.length != 1 || putArgs.length != 2 ||
! equivalenceChecker.expressionsAreEquivalent(getArgs[0], putArgs[0])) return;
PsiExpression value = putArgs[1];
if (LambdaGenerationUtil.canBeUncheckedLambda(value)) {
GetPutToComputeFix fix = new GetPutToComputeFix(variable);
holder.registerProblem(commonPutCall,
QuickFixBundle.message("java.8.map.api.inspection.description", "compute"), fix);
}
}
@Override
public void visitIfStatement(@NotNull PsiIfStatement statement) {
@@ -101,7 +148,7 @@ public class Java8MapApiInspection extends AbstractBaseJavaLocalInspectionTool {
processMerge(condition, existsBranch, noneBranch);
}
if(condition.hasVariable()) return;
EquivalenceChecker.Match match = EquivalenceChecker.getCanonicalPsiEquivalence().statementsMatch(noneBranch, existsBranch);
EquivalenceChecker.Match match = getCanonicalPsiEquivalence().statementsMatch(noneBranch, existsBranch);
processGetPut(condition, existsBranch, match.getRightDiff(), match.getLeftDiff());
}
@@ -128,6 +175,27 @@ public class Java8MapApiInspection extends AbstractBaseJavaLocalInspectionTool {
QuickFixBundle.message("java.8.map.api.inspection.description", fix.myMethodName), fix);
}
private static PsiMethodCallExpression findPutMethodParent(PsiElement element) {
while (element != null && !(element instanceof PsiMethod)) {
if (element instanceof PsiExpression expression) {
PsiMethodCallExpression putCall = extractMapMethodCall(expression, "put");
if (putCall != null) return putCall;
}
element = element.getParent();
}
return null;
}
private static boolean isCommonPutCallForAllReferences(List<PsiReferenceExpression> references, PsiMethodCallExpression commonPutCall) {
for (PsiReferenceExpression reference : references) {
PsiMethodCallExpression putCall = findPutMethodParent(reference);
if (putCall != commonPutCall) {
return false;
}
}
return true;
}
private static boolean hasMapUsages(@NotNull MapLoopCondition condition, @Nullable PsiExpression value) {
return !VariableAccessUtils.getVariableReferences(condition.getMap(), value).stream()
.map(ExpressionUtils::getCallForQualifier)
@@ -452,6 +520,45 @@ public class Java8MapApiInspection extends AbstractBaseJavaLocalInspectionTool {
return new ReplaceWithSingleMapOperation(methodName, call, value, result);
}
}
private static class GetPutToComputeFix extends PsiUpdateModCommandQuickFix {
private final SmartPsiElementPointer<PsiLocalVariable> variablePointer;
private GetPutToComputeFix(PsiLocalVariable variable) {
variablePointer = SmartPointerManager.createPointer(variable);
}
@Override
public @NotNull String getName() {
return QuickFixBundle.message("java.8.map.api.inspection.fix.text", "compute");
}
@Override
public @NotNull String getFamilyName() {
return QuickFixBundle.message("java.8.map.api.inspection.fix.family.name");
}
@Override
protected void applyFix(@NotNull Project project, @NotNull PsiElement element, @NotNull ModPsiUpdater updater) {
CommentTracker commentTracker = new CommentTracker();
PsiMethodCallExpression call = (PsiMethodCallExpression) element;
PsiLocalVariable variable = updater.getWritable(variablePointer.getElement());
if (variable == null) return;
ExpressionUtils.bindCallTo(call, "compute");
String variableName = variable.getName();
PsiExpressionList argsList = call.getArgumentList();
PsiExpression[] args = argsList.getExpressions();
if(args.length != 2) return;
PsiExpression exp = args[1];
VariableNameGenerator generator = new VariableNameGenerator(call, VariableKind.PARAMETER);
String keyName = generator.byName("k", "key").generate(true);
String lambdaParameters = "(" + keyName + ", " + variableName + ")";
String lambdaExpressionText = lambdaParameters + " -> " + commentTracker.text(exp);
commentTracker.delete(variable);
commentTracker.replaceExpressionAndRestoreComments(exp, lambdaExpressionText);
}
}
private static void register(MapCheckCondition condition, ProblemsHolder holder, boolean informationLevel, ReplaceWithSingleMapOperation fix) {
if (informationLevel && !holder.isOnTheFly()) return;
@@ -0,0 +1,8 @@
// "Replace with 'compute' method call" "true"
import java.util.Map;
public class Main {
public void testCompute(Map<String, Integer> map, String key) {
map.compute(key, (k, value) -> value == null ? 0 : value + 1);
}
}
@@ -0,0 +1,8 @@
// "Replace with 'compute' method call" "true"
import java.util.Map;
public class Main {
public void testCompute(Map<String, Integer> map, String key) {
map/*7*/./*8*/compute/*9*/(/*10*/key/*11*/, /*1*/ /*2*/ /*3*/ /*4*/ /*5*/ /*6*/ (k, value) -> value /*12*/ == /*13*/null ? 0 : /*14*/value + 1)/*15*/;
}
}
@@ -0,0 +1,9 @@
// "Replace with 'compute' method call" "true"
import java.util.Map;
public class Main {
public void testCompute(Map<String, Integer> map, String key) {
Integer value = map.get(key);
map.pu<caret>t(key, value == null ? 0 : value + 1);
}
}
@@ -0,0 +1,9 @@
// "Replace with 'compute' method call" "false"
import java.util.Map;
public class Main {
public void testCompute(Map<String, Integer> map, String key) {
Integer value = map.get();
map.pu<caret>t(key, value == null ? 0 : value + 1);
}
}
@@ -0,0 +1,10 @@
// "Replace with 'compute' method call" "false"
import java.util.Map;
public class Main {
public void testCompute(Map<String, Integer> map, String key) {
Integer value = map.get(key);
%
map.pu<caret>t(key, value == null ? 0 : value + 1);
}
}
@@ -0,0 +1,14 @@
// "Replace with 'compute' method call" "false"
import java.util.Map;
public class Main {
public Integer sum(Integer a, Integer b) {
return a + b;
}
public void testCompute(Map<String, Integer> map, String key) {
Integer value = map.get(key);
sum(6, map.pu<caret>t(key, value == null ? 0 : value + 1));
}
}
@@ -0,0 +1,9 @@
// "Replace with 'compute' method call" "true"
import java.util.Map;
public class Main {
public void testCompute(Map<String, Integer> map, String key) {
Integer value = map/*1*/./*2*/get/*3*/(/*4*/key/*5*/)/*6*/;
map/*7*/./*8*/pu<caret>t/*9*/(/*10*/key/*11*/, value /*12*/== /*13*/null ? 0 : /*14*/value + 1)/*15*/;
}
}