diff --git a/java/java-analysis-impl/src/com/intellij/codeInspection/dataFlow/DfaMemoryState.java b/java/java-analysis-impl/src/com/intellij/codeInspection/dataFlow/DfaMemoryState.java index f8d8ef30ed02..e59cf69e61d6 100644 --- a/java/java-analysis-impl/src/com/intellij/codeInspection/dataFlow/DfaMemoryState.java +++ b/java/java-analysis-impl/src/com/intellij/codeInspection/dataFlow/DfaMemoryState.java @@ -19,6 +19,7 @@ import com.intellij.codeInspection.dataFlow.value.DfaConstValue; import com.intellij.codeInspection.dataFlow.value.DfaRelationValue; import com.intellij.codeInspection.dataFlow.value.DfaValue; import com.intellij.codeInspection.dataFlow.value.DfaVariableValue; +import com.intellij.psi.PsiType; import org.jetbrains.annotations.NotNull; import org.jetbrains.annotations.Nullable; @@ -55,6 +56,13 @@ public interface DfaMemoryState { @Nullable T getValueFact(@NotNull DfaFactType factType, @NotNull DfaValue value); + /** + * @param value to determine its type + * @return value type at this state if known (possibly erased) + */ + @Nullable + PsiType getValueType(DfaValue value); + void flushFields(); void flushVariable(DfaVariableValue variable); diff --git a/java/java-analysis-impl/src/com/intellij/codeInspection/dataFlow/DfaMemoryStateImpl.java b/java/java-analysis-impl/src/com/intellij/codeInspection/dataFlow/DfaMemoryStateImpl.java index ad60aab3cfff..3ea7c590c4df 100644 --- a/java/java-analysis-impl/src/com/intellij/codeInspection/dataFlow/DfaMemoryStateImpl.java +++ b/java/java-analysis-impl/src/com/intellij/codeInspection/dataFlow/DfaMemoryStateImpl.java @@ -1068,6 +1068,23 @@ public class DfaMemoryStateImpl implements DfaMemoryState { return factType.fromDfaValue(value); } + @Nullable + @Override + public PsiType getValueType(DfaValue value) { + if (value instanceof DfaTypeValue) { + return ((DfaTypeValue)value).getDfaType().getPsiType(); + } + if (value instanceof DfaVariableValue) { + DfaVariableState state = getVariableState((DfaVariableValue)value); + Set values = state.getInstanceofValues(); + if (!values.isEmpty()) { + return values.iterator().next().getPsiType(); + } + return ((DfaVariableValue)value).getVariableType(); + } + return null; + } + @NotNull private DfaValue resolveVariableValue(DfaVariableValue var) { DfaConstValue constValue = getConstantValue(var); diff --git a/java/java-analysis-impl/src/com/intellij/codeInspection/dataFlow/StandardInstructionVisitor.java b/java/java-analysis-impl/src/com/intellij/codeInspection/dataFlow/StandardInstructionVisitor.java index c25fd349dcd4..957a3dc01807 100644 --- a/java/java-analysis-impl/src/com/intellij/codeInspection/dataFlow/StandardInstructionVisitor.java +++ b/java/java-analysis-impl/src/com/intellij/codeInspection/dataFlow/StandardInstructionVisitor.java @@ -24,6 +24,7 @@ import com.intellij.openapi.diagnostic.Logger; import com.intellij.openapi.util.Pair; import com.intellij.psi.*; import com.intellij.psi.tree.IElementType; +import com.intellij.psi.util.InheritanceUtil; import com.intellij.psi.util.PsiTreeUtil; import com.intellij.psi.util.PsiUtil; import com.intellij.psi.util.TypeConversionUtil; @@ -302,10 +303,10 @@ public class StandardInstructionVisitor extends InstructionVisitor { DfaCallArguments callArguments = popCall(instruction, runner, memState, true); LinkedHashSet currentStates = ContainerUtil.newLinkedHashSet(memState); + DfaValue resultValue = getMethodResultValue(instruction, callArguments.myQualifier, memState, runner.getFactory()); if (callArguments.myArguments != null) { for (MethodContract contract : instruction.getContracts()) { - DfaValue returnValue = getMethodResultValue(instruction, callArguments.myQualifier, runner.getFactory()); - returnValue = contract.getDfaReturnValue(runner.getFactory(), returnValue); + DfaValue returnValue = contract.getDfaReturnValue(runner.getFactory(), resultValue); currentStates = addContractResults(callArguments, contract, currentStates, runner.getFactory(), finalStates, returnValue); if (currentStates.size() + finalStates.size() > DataFlowRunner.MAX_STATES_PER_BRANCH) { if (LOG.isDebugEnabled()) { @@ -318,7 +319,7 @@ public class StandardInstructionVisitor extends InstructionVisitor { } } for (DfaMemoryState state : currentStates) { - state.push(getMethodResultValue(instruction, callArguments.myQualifier, runner.getFactory())); + state.push(resultValue); finalStates.add(state); } } @@ -348,7 +349,7 @@ public class StandardInstructionVisitor extends InstructionVisitor { callArguments.myArguments == null ? Collections.emptyList() : handler.handle(callArguments, memState, runner.getFactory()); if (states.isEmpty()) { - memState.push(getMethodResultValue(instruction, callArguments.myQualifier, runner.getFactory())); + memState.push(getMethodResultValue(instruction, callArguments.myQualifier, memState, runner.getFactory())); return Collections.singletonList(memState); } return states; @@ -479,13 +480,13 @@ public class StandardInstructionVisitor extends InstructionVisitor { @NotNull private static DfaValue getMethodResultValue(MethodCallInstruction instruction, @Nullable DfaValue qualifierValue, - DfaValueFactory factory) { + DfaMemoryState state, DfaValueFactory factory) { DfaValue precalculated = instruction.getPrecalculatedReturnValue(); if (precalculated != null) { return precalculated; } - final PsiType type = instruction.getResultType(); + PsiType type = instruction.getResultType(); final MethodCallInstruction.MethodType methodType = instruction.getMethodType(); if (methodType == MethodCallInstruction.MethodType.METHOD_REFERENCE_CALL && qualifierValue instanceof DfaVariableValue) { @@ -518,8 +519,25 @@ public class StandardInstructionVisitor extends InstructionVisitor { if (type != null && !(type instanceof PsiPrimitiveType)) { Nullness nullability = instruction.getReturnNullability(); PsiMethod targetMethod = instruction.getTargetMethod(); - if (nullability == Nullness.UNKNOWN && targetMethod != null) { - nullability = factory.suggestNullabilityForNonAnnotatedMember(targetMethod); + if (targetMethod != null) { + PsiClass specificQualifierClass = PsiUtil.resolveClassInClassTypeOnly(state.getValueType(qualifierValue)); + PsiClass qualifierClass = targetMethod.getContainingClass(); + if (specificQualifierClass != null && qualifierClass != null && + !specificQualifierClass.equals(qualifierClass) && + InheritanceUtil.isInheritorOrSelf(specificQualifierClass, qualifierClass, true)) { + PsiMethod realMethod = specificQualifierClass.findMethodBySignature(targetMethod, true); + if (realMethod != null && realMethod != targetMethod) { + nullability = DfaPsiUtil.getElementNullability(type, realMethod); + PsiType returnType = realMethod.getReturnType(); + if(returnType != null && TypeConversionUtil.erasure(type).isAssignableFrom(returnType)) { + // possibly covariant return type + type = returnType; + } + } + } + if (nullability == Nullness.UNKNOWN) { + nullability = factory.suggestNullabilityForNonAnnotatedMember(targetMethod); + } } return factory.createTypeValue(type, nullability); } diff --git a/java/java-tests/testData/inspection/dataFlow/fixture/CovariantReturn.java b/java/java-tests/testData/inspection/dataFlow/fixture/CovariantReturn.java new file mode 100644 index 000000000000..76cff4cb724d --- /dev/null +++ b/java/java-tests/testData/inspection/dataFlow/fixture/CovariantReturn.java @@ -0,0 +1,36 @@ +import org.jetbrains.annotations.Nullable; +import org.jetbrains.annotations.NotNull; + +class CovariantReturn { + public interface Base { + @Nullable + String foo(); + + Object baz(); + } + + public interface Derived extends Base { + @Override + @NotNull + String foo(); + + String baz(); + } + + public static boolean bar(@NotNull String s) { + return true; + } + + // IDEA-178773 Can't get to green code without suppressing an inspection in case of restricted nullability in overridden method + public static boolean test1(Base base) { + return base instanceof Derived && bar(base.foo()); + } + + public static boolean test2(Base base) { + return base instanceof Derived || bar(base.foo()); + } + + boolean testType(Base base) { + return base instanceof Derived && base.baz() instanceof String; + } +} \ No newline at end of file diff --git a/java/java-tests/testSrc/com/intellij/java/codeInspection/DataFlowInspectionTest.java b/java/java-tests/testSrc/com/intellij/java/codeInspection/DataFlowInspectionTest.java index a545f1848822..092dce8af4ed 100644 --- a/java/java-tests/testSrc/com/intellij/java/codeInspection/DataFlowInspectionTest.java +++ b/java/java-tests/testSrc/com/intellij/java/codeInspection/DataFlowInspectionTest.java @@ -540,4 +540,5 @@ public class DataFlowInspectionTest extends DataFlowInspectionTestCase { public void testFieldAssignedNegative() { doTest(); } public void testIteratePositiveCheck() { doTest(); } public void testInnerClass() { doTest(); } + public void testCovariantReturn() { doTest(); } }