diff --git a/java/java-impl-inspections/src/com/intellij/codeInspection/java18api/Java8MapApiInspection.java b/java/java-impl-inspections/src/com/intellij/codeInspection/java18api/Java8MapApiInspection.java index 29b7982e8e7f..d74de38a49b0 100644 --- a/java/java-impl-inspections/src/com/intellij/codeInspection/java18api/Java8MapApiInspection.java +++ b/java/java-impl-inspections/src/com/intellij/codeInspection/java18api/Java8MapApiInspection.java @@ -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 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 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 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; diff --git a/java/java-tests/testData/inspection/java8MapApi/afterCompute.java b/java/java-tests/testData/inspection/java8MapApi/afterCompute.java new file mode 100644 index 000000000000..42e78ef0c73c --- /dev/null +++ b/java/java-tests/testData/inspection/java8MapApi/afterCompute.java @@ -0,0 +1,8 @@ +// "Replace with 'compute' method call" "true" +import java.util.Map; + +public class Main { + public void testCompute(Map map, String key) { + map.compute(key, (k, value) -> value == null ? 0 : value + 1); + } +} \ No newline at end of file diff --git a/java/java-tests/testData/inspection/java8MapApi/afterComputeWithComments.java b/java/java-tests/testData/inspection/java8MapApi/afterComputeWithComments.java new file mode 100644 index 000000000000..7f3a64bc64af --- /dev/null +++ b/java/java-tests/testData/inspection/java8MapApi/afterComputeWithComments.java @@ -0,0 +1,8 @@ +// "Replace with 'compute' method call" "true" +import java.util.Map; + +public class Main { + public void testCompute(Map 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*/; + } +} \ No newline at end of file diff --git a/java/java-tests/testData/inspection/java8MapApi/beforeCompute.java b/java/java-tests/testData/inspection/java8MapApi/beforeCompute.java new file mode 100644 index 000000000000..8fb88097adfc --- /dev/null +++ b/java/java-tests/testData/inspection/java8MapApi/beforeCompute.java @@ -0,0 +1,9 @@ +// "Replace with 'compute' method call" "true" +import java.util.Map; + +public class Main { + public void testCompute(Map map, String key) { + Integer value = map.get(key); + map.put(key, value == null ? 0 : value + 1); + } +} \ No newline at end of file diff --git a/java/java-tests/testData/inspection/java8MapApi/beforeComputeIncompleteCode.java b/java/java-tests/testData/inspection/java8MapApi/beforeComputeIncompleteCode.java new file mode 100644 index 000000000000..f187ae341a51 --- /dev/null +++ b/java/java-tests/testData/inspection/java8MapApi/beforeComputeIncompleteCode.java @@ -0,0 +1,9 @@ +// "Replace with 'compute' method call" "false" +import java.util.Map; + +public class Main { + public void testCompute(Map map, String key) { + Integer value = map.get(); + map.put(key, value == null ? 0 : value + 1); + } +} \ No newline at end of file diff --git a/java/java-tests/testData/inspection/java8MapApi/beforeComputeNotSiblingStatements.java b/java/java-tests/testData/inspection/java8MapApi/beforeComputeNotSiblingStatements.java new file mode 100644 index 000000000000..e613f21737d9 --- /dev/null +++ b/java/java-tests/testData/inspection/java8MapApi/beforeComputeNotSiblingStatements.java @@ -0,0 +1,10 @@ +// "Replace with 'compute' method call" "false" +import java.util.Map; + +public class Main { + public void testCompute(Map map, String key) { + Integer value = map.get(key); + % + map.put(key, value == null ? 0 : value + 1); + } +} \ No newline at end of file diff --git a/java/java-tests/testData/inspection/java8MapApi/beforeComputePutInsideExpression.java b/java/java-tests/testData/inspection/java8MapApi/beforeComputePutInsideExpression.java new file mode 100644 index 000000000000..887d1c3e7676 --- /dev/null +++ b/java/java-tests/testData/inspection/java8MapApi/beforeComputePutInsideExpression.java @@ -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 map, String key) { + Integer value = map.get(key); + sum(6, map.put(key, value == null ? 0 : value + 1)); + } +} \ No newline at end of file diff --git a/java/java-tests/testData/inspection/java8MapApi/beforeComputeWithComments.java b/java/java-tests/testData/inspection/java8MapApi/beforeComputeWithComments.java new file mode 100644 index 000000000000..59259a742c5f --- /dev/null +++ b/java/java-tests/testData/inspection/java8MapApi/beforeComputeWithComments.java @@ -0,0 +1,9 @@ +// "Replace with 'compute' method call" "true" +import java.util.Map; + +public class Main { + public void testCompute(Map map, String key) { + Integer value = map/*1*/./*2*/get/*3*/(/*4*/key/*5*/)/*6*/; + map/*7*/./*8*/put/*9*/(/*10*/key/*11*/, value /*12*/== /*13*/null ? 0 : /*14*/value + 1)/*15*/; + } +} \ No newline at end of file