diff --git a/java/java-analysis-impl/src/com/intellij/codeInspection/dataFlow/DfaUtil.java b/java/java-analysis-impl/src/com/intellij/codeInspection/dataFlow/DfaUtil.java index 007e0b5638b0..bf6b96ad5d48 100644 --- a/java/java-analysis-impl/src/com/intellij/codeInspection/dataFlow/DfaUtil.java +++ b/java/java-analysis-impl/src/com/intellij/codeInspection/dataFlow/DfaUtil.java @@ -24,6 +24,7 @@ import com.intellij.psi.*; import com.intellij.psi.tree.IElementType; import com.intellij.psi.util.CachedValueProvider; import com.intellij.psi.util.CachedValuesManager; +import com.intellij.psi.util.PsiTreeUtil; import com.intellij.psi.util.PsiUtil; import com.intellij.util.IncorrectOperationException; import com.intellij.util.containers.ContainerUtil; @@ -130,6 +131,21 @@ public class DfaUtil { return Nullness.UNKNOWN; } + return inferBlockNullity(body, InferenceFromSourceUtil.suppressNullable(method)); + } + + @NotNull + public static Nullness inferLambdaNullity(PsiLambdaExpression lambda) { + final PsiElement body = lambda.getBody(); + if (body == null || LambdaUtil.getFunctionalInterfaceReturnType(lambda) == null) { + return Nullness.UNKNOWN; + } + + return inferBlockNullity(body, false); + } + + @NotNull + private static Nullness inferBlockNullity(PsiElement body, boolean suppressNullable) { final AtomicBoolean hasNulls = new AtomicBoolean(); final AtomicBoolean hasNotNulls = new AtomicBoolean(); final AtomicBoolean hasUnknowns = new AtomicBoolean(); @@ -140,15 +156,17 @@ public class DfaUtil { public DfaInstructionState[] visitCheckReturnValue(CheckReturnValueInstruction instruction, DataFlowRunner runner, DfaMemoryState memState) { - DfaValue returned = memState.peek(); - if (memState.isNull(returned)) { - hasNulls.set(true); - } - else if (memState.isNotNull(returned)) { - hasNotNulls.set(true); - } - else { - hasUnknowns.set(true); + if(PsiTreeUtil.isAncestor(body, instruction.getReturn(), false)) { + DfaValue returned = memState.peek(); + if (memState.isNull(returned)) { + hasNulls.set(true); + } + else if (memState.isNotNull(returned)) { + hasNotNulls.set(true); + } + else { + hasUnknowns.set(true); + } } return super.visitCheckReturnValue(instruction, runner, memState); } @@ -156,7 +174,7 @@ public class DfaUtil { if (rc == RunnerResult.OK) { if (hasNulls.get()) { - return InferenceFromSourceUtil.suppressNullable(method) ? Nullness.UNKNOWN : Nullness.NULLABLE; + return suppressNullable ? Nullness.UNKNOWN : Nullness.NULLABLE; } if (hasNotNulls.get() && !hasUnknowns.get()) { return Nullness.NOT_NULL; diff --git a/java/java-impl/src/com/intellij/codeInspection/SimplifyStreamApiCallChainsInspection.java b/java/java-impl/src/com/intellij/codeInspection/SimplifyStreamApiCallChainsInspection.java index 40aa506b1d00..3adb2a6d2201 100644 --- a/java/java-impl/src/com/intellij/codeInspection/SimplifyStreamApiCallChainsInspection.java +++ b/java/java-impl/src/com/intellij/codeInspection/SimplifyStreamApiCallChainsInspection.java @@ -15,6 +15,9 @@ */ package com.intellij.codeInspection; +import com.intellij.codeInsight.NullableNotNullManager; +import com.intellij.codeInspection.dataFlow.DfaUtil; +import com.intellij.codeInspection.dataFlow.Nullness; import com.intellij.openapi.diagnostic.Logger; import com.intellij.openapi.project.Project; import com.intellij.openapi.util.TextRange; @@ -1158,21 +1161,28 @@ public class SimplifyStreamApiCallChainsInspection extends BaseJavaBatchLocalIns static CallHandler handler() { return CallHandler.of(STREAM_MATCH, call -> { + PsiMethodCallExpression qualifierCall = getQualifierMethodCall(call); + if (!STREAM_MAP.test(qualifierCall)) return null; + PsiExpression qualifierArg = PsiUtil.skipParenthesizedExprDown(qualifierCall.getArgumentList().getExpressions()[0]); PsiExpression predicate = call.getArgumentList().getExpressions()[0]; boolean invert = false; if (!isBooleanIdentity(predicate)) { - if (!"anyMatch".equals(call.getMethodExpression().getReferenceName()) && - Boolean.FALSE.equals(getBooleanEqualsTarget(predicate))) { - invert = true; + Boolean target = getBooleanEqualsTarget(predicate); + if (target == null || (!target && "anyMatch".equals(call.getMethodExpression().getReferenceName()))) return null; + invert = !target; + if (qualifierArg instanceof PsiMethodReferenceExpression) { + PsiMethod method = tryCast(((PsiMethodReferenceExpression)qualifierArg).resolve(), PsiMethod.class); + if (method == null) return null; + if (!PsiType.BOOLEAN.equals(method.getReturnType()) && !NullableNotNullManager.isNotNull(method)) return null; } - else { + else if (!(qualifierArg instanceof PsiLambdaExpression) || + DfaUtil.inferLambdaNullity((PsiLambdaExpression)qualifierArg) != Nullness.NOT_NULL) { return null; } } - PsiMethodCallExpression qualifierCall = getQualifierMethodCall(call); - if (!STREAM_MAP.test(qualifierCall)) return null; - PsiExpression qualifierArg = qualifierCall.getArgumentList().getExpressions()[0]; - if (adaptToPredicate(qualifierArg) == null) return null; + else { + if (adaptToPredicate(qualifierArg) == null) return null; + } return new RemoveBooleanIdentityFix(invert); }); } @@ -1182,8 +1192,7 @@ public class SimplifyStreamApiCallChainsInspection extends BaseJavaBatchLocalIns if (FunctionalExpressionUtils.isFunctionalReferenceTo(arg, CommonClassNames.JAVA_LANG_BOOLEAN, PsiType.BOOLEAN, "booleanValue", PsiType.EMPTY_ARRAY) || FunctionalExpressionUtils.isFunctionalReferenceTo(arg, CommonClassNames.JAVA_LANG_BOOLEAN, null, - "valueOf", PsiType.BOOLEAN) || - Boolean.TRUE.equals(getBooleanEqualsTarget(arg))) { + "valueOf", PsiType.BOOLEAN)) { return true; } return arg instanceof PsiLambdaExpression && LambdaUtil.isIdentityLambda((PsiLambdaExpression)arg); diff --git a/java/java-tests/testData/inspection/streamApiCallChains/afterBooleanIdentity.java b/java/java-tests/testData/inspection/streamApiCallChains/afterBooleanIdentity.java index 04b263402fde..11fd6636cb6f 100644 --- a/java/java-tests/testData/inspection/streamApiCallChains/afterBooleanIdentity.java +++ b/java/java-tests/testData/inspection/streamApiCallChains/afterBooleanIdentity.java @@ -1,5 +1,5 @@ // "Fix all 'Simplify stream API call chains' problems in file" "true" -import java.util.List; +import java.util.*; import java.util.function.Function; import java.util.stream.Stream; @@ -76,4 +76,16 @@ public class Main { public boolean noneMatchBooleanEqualsLambda(List list) { return list.stream().map(String::isEmpty).noneMatch(b -> Boolean.FALSE.equals((!b))); } + + interface Source { + Optional containsDescendantOf(String name); + } + + public static boolean testNullityInference(Stream sources, String name) { + return sources.noneMatch((s) -> s.containsDescendantOf(name).orElse(true)); + } + + public static boolean testNullityInference2(Stream sources, String name) { + return sources.map((s) -> s.containsDescendantOf(name).orElse(null)).allMatch(Boolean.FALSE::equals); + } } \ No newline at end of file diff --git a/java/java-tests/testData/inspection/streamApiCallChains/beforeBooleanIdentity.java b/java/java-tests/testData/inspection/streamApiCallChains/beforeBooleanIdentity.java index 958f54c97577..331f4aa16e41 100644 --- a/java/java-tests/testData/inspection/streamApiCallChains/beforeBooleanIdentity.java +++ b/java/java-tests/testData/inspection/streamApiCallChains/beforeBooleanIdentity.java @@ -1,5 +1,5 @@ // "Fix all 'Simplify stream API call chains' problems in file" "true" -import java.util.List; +import java.util.*; import java.util.function.Function; import java.util.stream.Stream; @@ -78,4 +78,16 @@ public class Main { public boolean noneMatchBooleanEqualsLambda(List list) { return list.stream().map(String::isEmpty).noneMatch(b -> Boolean.FALSE.equals((!b))); } + + interface Source { + Optional containsDescendantOf(String name); + } + + public static boolean testNullityInference(Stream sources, String name) { + return sources.map((s) -> s.containsDescendantOf(name).orElse(true)).allMatch(Boolean.FALSE::equals); + } + + public static boolean testNullityInference2(Stream sources, String name) { + return sources.map((s) -> s.containsDescendantOf(name).orElse(null)).allMatch(Boolean.FALSE::equals); + } } \ No newline at end of file