ContractReturnValue: respect NotNull annotation for ->paramN contracts

This commit is contained in:
Tagir Valeev
2018-05-07 13:01:45 +07:00
parent 6a3aecbce2
commit 28b1de4b70
4 changed files with 112 additions and 65 deletions
@@ -2,6 +2,7 @@
package com.intellij.codeInspection.dataFlow;
import com.intellij.codeInspection.dataFlow.value.*;
import com.intellij.openapi.util.Pair;
import com.intellij.psi.*;
import org.jetbrains.annotations.NotNull;
import org.jetbrains.annotations.Nullable;
@@ -85,15 +86,33 @@ public abstract class ContractReturnValue {
abstract Stream<Function<PsiMethod, String>> validators();
static DfaValue merge(DfaValue defaultValue, DfaValue newValue, DfaMemoryState memState) {
if (defaultValue == null || defaultValue == DfaUnknownValue.getInstance()) return newValue;
if (newValue == null || newValue == DfaUnknownValue.getInstance()) return defaultValue;
if (defaultValue instanceof DfaFactMapValue) {
DfaFactMap defaultFacts = ((DfaFactMapValue)defaultValue).getFacts();
if (newValue instanceof DfaFactMapValue) {
DfaFactMap intersection = defaultFacts.intersect(((DfaFactMapValue)newValue).getFacts());
if (intersection != null) {
return defaultValue.getFactory().getFactFactory().createValue(intersection);
}
}
if (newValue instanceof DfaVariableValue) {
defaultFacts.facts(Pair::create).forEach(fact -> memState.applyFact(newValue, fact.getFirst(), fact.getSecond()));
}
}
return newValue;
}
/**
* Converts this return value to the most suitable {@link DfaValue} which represents the same constraints.
*
* @param factory a {@link DfaValueFactory} which can be used to create new values if necessary
* @param defaultValue a default method return type value in the absence of the contracts (may contain method type information)
* @param arguments arguments passed
* @param callState call state
* @return a value which represents the constraints of this contract return value.
*/
public abstract DfaValue toDfaValue(DfaValueFactory factory, DfaValue defaultValue, DfaCallArguments arguments);
public abstract DfaValue getDfaValue(DfaValueFactory factory, DfaValue defaultValue, DfaCallState callState);
/**
* Returns true if the supplied {@link DfaValue} could be compatible with this return value. If false is returned, then
@@ -304,7 +323,7 @@ public abstract class ContractReturnValue {
}
@Override
public DfaValue toDfaValue(DfaValueFactory factory, DfaValue defaultValue, DfaCallArguments arguments) {
public DfaValue getDfaValue(DfaValueFactory factory, DfaValue defaultValue, DfaCallState callState) {
return defaultValue;
}
@@ -321,7 +340,7 @@ public abstract class ContractReturnValue {
}
@Override
public DfaValue toDfaValue(DfaValueFactory factory, DfaValue defaultValue, DfaCallArguments arguments) {
public DfaValue getDfaValue(DfaValueFactory factory, DfaValue defaultValue, DfaCallState callState) {
return factory.getConstFactory().getContractFail();
}
@@ -338,7 +357,7 @@ public abstract class ContractReturnValue {
}
@Override
public DfaValue toDfaValue(DfaValueFactory factory, DfaValue defaultValue, DfaCallArguments arguments) {
public DfaValue getDfaValue(DfaValueFactory factory, DfaValue defaultValue, DfaCallState callState) {
return factory.getConstFactory().getNull();
}
@@ -360,7 +379,11 @@ public abstract class ContractReturnValue {
}
@Override
public DfaValue toDfaValue(DfaValueFactory factory, DfaValue defaultValue, DfaCallArguments arguments) {
public DfaValue getDfaValue(DfaValueFactory factory, DfaValue defaultValue, DfaCallState callState) {
if (defaultValue instanceof DfaVariableValue) {
callState.myMemoryState.forceVariableFact((DfaVariableValue)defaultValue, DfaFactType.CAN_BE_NULL, false);
return defaultValue;
}
return factory.withFact(defaultValue, DfaFactType.CAN_BE_NULL, false);
}
@@ -377,10 +400,13 @@ public abstract class ContractReturnValue {
}
@Override
public DfaValue toDfaValue(DfaValueFactory factory, DfaValue defaultValue, DfaCallArguments arguments) {
if (defaultValue instanceof DfaVariableValue) return defaultValue;
public DfaValue getDfaValue(DfaValueFactory factory, DfaValue defaultValue, DfaCallState callState) {
if (defaultValue instanceof DfaVariableValue) {
callState.myMemoryState.forceVariableFact((DfaVariableValue)defaultValue, DfaFactType.CAN_BE_NULL, false);
return defaultValue;
}
DfaValue value = factory.withFact(defaultValue, DfaFactType.CAN_BE_NULL, false);
if (arguments.myPure) {
if (callState.myCallArguments.myPure) {
boolean unmodifiableView =
value instanceof DfaFactMapValue && ((DfaFactMapValue)value).get(DfaFactType.MUTABILITY) == Mutability.UNMODIFIABLE_VIEW;
// Unmodifiable view methods like Collections.unmodifiableList create new object, but their special field "size" is
@@ -414,9 +440,11 @@ public abstract class ContractReturnValue {
}
@Override
public DfaValue toDfaValue(DfaValueFactory factory, DfaValue defaultValue, DfaCallArguments arguments) {
DfaValue qualifier = arguments.myQualifier;
if (qualifier != null && qualifier != DfaUnknownValue.getInstance()) return qualifier;
public DfaValue getDfaValue(DfaValueFactory factory, DfaValue defaultValue, DfaCallState callState) {
DfaValue qualifier = callState.myCallArguments.myQualifier;
if (qualifier != null && qualifier != DfaUnknownValue.getInstance()) {
return merge(defaultValue, qualifier, callState.myMemoryState);
}
return factory.withFact(defaultValue, DfaFactType.CAN_BE_NULL, false);
}
@@ -456,7 +484,7 @@ public abstract class ContractReturnValue {
}
@Override
public DfaValue toDfaValue(DfaValueFactory factory, DfaValue defaultValue, DfaCallArguments arguments) {
public DfaValue getDfaValue(DfaValueFactory factory, DfaValue defaultValue, DfaCallState callState) {
return factory.getBoolean(myValue);
}
@@ -512,18 +540,10 @@ public abstract class ContractReturnValue {
}
@Override
public DfaValue toDfaValue(DfaValueFactory factory, DfaValue defaultValue, DfaCallArguments arguments) {
if (arguments.myArguments != null && arguments.myArguments.length > myParamNumber) {
DfaValue argument = arguments.myArguments[myParamNumber];
if (argument != null && argument != DfaUnknownValue.getInstance()) {
if (defaultValue instanceof DfaFactMapValue && argument instanceof DfaFactMapValue) {
DfaFactMap intersection = ((DfaFactMapValue)defaultValue).getFacts().intersect(((DfaFactMapValue)argument).getFacts());
if (intersection != null) {
return factory.getFactFactory().createValue(intersection);
}
}
return argument;
}
public DfaValue getDfaValue(DfaValueFactory factory, DfaValue defaultValue, DfaCallState callState) {
if (callState.myCallArguments.myArguments != null && callState.myCallArguments.myArguments.length > myParamNumber) {
DfaValue argument = callState.myCallArguments.myArguments[myParamNumber];
return merge(defaultValue, argument, callState.myMemoryState);
}
return defaultValue;
}
@@ -0,0 +1,27 @@
// Copyright 2000-2018 JetBrains s.r.o. Use of this source code is governed by the Apache 2.0 license that can be found in the LICENSE file.
package com.intellij.codeInspection.dataFlow;
import org.jetbrains.annotations.NotNull;
final class DfaCallState {
@NotNull final DfaMemoryState myMemoryState;
@NotNull final DfaCallArguments myCallArguments;
DfaCallState(@NotNull DfaMemoryState state, @NotNull DfaCallArguments arguments) {
myMemoryState = state;
myCallArguments = arguments;
}
@Override
public boolean equals(Object o) {
if (this == o) return true;
if (!(o instanceof DfaCallState)) return false;
DfaCallState that = (DfaCallState)o;
return myMemoryState.equals(that.myMemoryState) && myCallArguments.equals(that.myCallArguments);
}
@Override
public int hashCode() {
return 31 * myMemoryState.hashCode() + myCallArguments.hashCode();
}
}
@@ -271,7 +271,7 @@ public class StandardInstructionVisitor extends InstructionVisitor {
DfaMemoryState state,
List<? extends MethodContract> contracts,
DfaValueFactory factory, DfaValue defaultResult) {
Set<CallState> currentStates = Collections.singleton(new CallState(state.createClosureState(), callArguments));
Set<DfaCallState> currentStates = Collections.singleton(new DfaCallState(state.createClosureState(), callArguments));
Set<DfaMemoryState> finalStates = ContainerUtil.newLinkedHashSet();
for (MethodContract contract : contracts) {
currentStates = addContractResults(contract, currentStates, factory, finalStates, defaultResult);
@@ -311,7 +311,7 @@ public class StandardInstructionVisitor extends InstructionVisitor {
if (finalStates.isEmpty()) {
DfaCallArguments callArguments = popCall(instruction, runner, memState, true);
Set<CallState> currentStates = Collections.singleton(new CallState(memState, callArguments));
Set<DfaCallState> currentStates = Collections.singleton(new DfaCallState(memState, callArguments));
DfaValue defaultResult = getMethodResultValue(instruction, callArguments.myQualifier, memState, runner.getFactory());
if (callArguments.myArguments != null) {
for (MethodContract contract : instruction.getContracts()) {
@@ -321,12 +321,12 @@ public class StandardInstructionVisitor extends InstructionVisitor {
LOG.debug("Too complex contract on " + instruction.getContext() + ", skipping contract processing");
}
finalStates.clear();
currentStates = Collections.singleton(new CallState(memState, callArguments));
currentStates = Collections.singleton(new DfaCallState(memState, callArguments));
break;
}
}
}
for (CallState callState : currentStates) {
for (DfaCallState callState : currentStates) {
callState.myMemoryState.push(defaultResult);
finalStates.add(callState.myMemoryState);
}
@@ -444,46 +444,23 @@ public class StandardInstructionVisitor extends InstructionVisitor {
return value;
}
private static final class CallState {
@NotNull final DfaMemoryState myMemoryState;
@NotNull final DfaCallArguments myCallArguments;
private CallState(@NotNull DfaMemoryState state, @NotNull DfaCallArguments arguments) {
myMemoryState = state;
myCallArguments = arguments;
}
@Override
public boolean equals(Object o) {
if (this == o) return true;
if (!(o instanceof CallState)) return false;
CallState that = (CallState)o;
return myMemoryState.equals(that.myMemoryState) && myCallArguments.equals(that.myCallArguments);
}
@Override
public int hashCode() {
return 31 * myMemoryState.hashCode() + myCallArguments.hashCode();
}
}
private static Set<CallState> addContractResults(MethodContract contract,
Set<CallState> states,
DfaValueFactory factory,
Set<DfaMemoryState> finalStates,
DfaValue defaultResult) {
private static Set<DfaCallState> addContractResults(MethodContract contract,
Set<DfaCallState> states,
DfaValueFactory factory,
Set<DfaMemoryState> finalStates,
DfaValue defaultResult) {
if(contract.isTrivial()) {
for (CallState callState : states) {
DfaValue returnValue = contract.getReturnValue().toDfaValue(factory, defaultResult, callState.myCallArguments);
callState.myMemoryState.push(returnValue);
for (DfaCallState callState : states) {
DfaValue result = contract.getReturnValue().getDfaValue(factory, defaultResult, callState);;
callState.myMemoryState.push(result);
finalStates.add(callState.myMemoryState);
}
return Collections.emptySet();
}
Set<CallState> falseStates = new LinkedHashSet<>();
Set<DfaCallState> falseStates = new LinkedHashSet<>();
for (CallState callState : states) {
for (DfaCallState callState : states) {
DfaMemoryState state = callState.myMemoryState;
DfaCallArguments arguments = callState.myCallArguments;
for (ContractValue contractValue : contract.getConditions()) {
@@ -494,7 +471,7 @@ public class StandardInstructionVisitor extends InstructionVisitor {
DfaMemoryState falseState = state.createCopy();
if (falseState.applyContractCondition(condition.createNegated())) {
DfaCallArguments falseArguments = contractValue.updateArguments(arguments, true);
falseStates.add(new CallState(falseState, falseArguments));
falseStates.add(new DfaCallState(falseState, falseArguments));
}
if (!state.applyContractCondition(condition)) {
state = null;
@@ -503,8 +480,8 @@ public class StandardInstructionVisitor extends InstructionVisitor {
arguments = contractValue.updateArguments(arguments, false);
}
if(state != null) {
DfaValue returnValue = contract.getReturnValue().toDfaValue(factory, defaultResult, arguments);
state.push(returnValue);
DfaValue result = contract.getReturnValue().getDfaValue(factory, defaultResult, new DfaCallState(state, arguments));
state.push(result);
finalStates.add(state);
}
}
@@ -72,4 +72,27 @@ class ContractReturnValues {
}
private static final String SOMETHING = "foo";
void testCheckString() {
String s2 = nullable();
checkString(s2);
if (<warning descr="Condition 's2 == null' is always 'false'">s2 == null</warning>) {
System.out.println("impossible");
}
String s = nullable();
if (<warning descr="Condition 'checkString(s) == null' is always 'false'">checkString(s) == null</warning>) {
System.out.println("impossible");
}
}
// inferred contract: "_ -> param1"
// we should respect "notnull" even if we cannot infer failing contract
@NotNull
static String checkString(@Nullable String s) {
check(s); // probably does a null-check
return s;
}
@Contract("null -> fail")
native static void check(@Nullable String s);
}