diff --git a/java/java-analysis-impl/src/com/intellij/codeInspection/dataFlow/ContractReturnValue.java b/java/java-analysis-impl/src/com/intellij/codeInspection/dataFlow/ContractReturnValue.java index 8c62de450c42..7ddfeb434b47 100644 --- a/java/java-analysis-impl/src/com/intellij/codeInspection/dataFlow/ContractReturnValue.java +++ b/java/java-analysis-impl/src/com/intellij/codeInspection/dataFlow/ContractReturnValue.java @@ -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> 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; } diff --git a/java/java-analysis-impl/src/com/intellij/codeInspection/dataFlow/DfaCallState.java b/java/java-analysis-impl/src/com/intellij/codeInspection/dataFlow/DfaCallState.java new file mode 100644 index 000000000000..5f18482d60f2 --- /dev/null +++ b/java/java-analysis-impl/src/com/intellij/codeInspection/dataFlow/DfaCallState.java @@ -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(); + } +} 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 c55280faaa9f..c29e2a5102ed 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 @@ -271,7 +271,7 @@ public class StandardInstructionVisitor extends InstructionVisitor { DfaMemoryState state, List contracts, DfaValueFactory factory, DfaValue defaultResult) { - Set currentStates = Collections.singleton(new CallState(state.createClosureState(), callArguments)); + Set currentStates = Collections.singleton(new DfaCallState(state.createClosureState(), callArguments)); Set 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 currentStates = Collections.singleton(new CallState(memState, callArguments)); + Set 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 addContractResults(MethodContract contract, - Set states, - DfaValueFactory factory, - Set finalStates, - DfaValue defaultResult) { + private static Set addContractResults(MethodContract contract, + Set states, + DfaValueFactory factory, + Set 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 falseStates = new LinkedHashSet<>(); + Set 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); } } diff --git a/java/java-tests/testData/inspection/dataFlow/fixture/ContractReturnValues.java b/java/java-tests/testData/inspection/dataFlow/fixture/ContractReturnValues.java index b035c52ce48d..0141029af57b 100644 --- a/java/java-tests/testData/inspection/dataFlow/fixture/ContractReturnValues.java +++ b/java/java-tests/testData/inspection/dataFlow/fixture/ContractReturnValues.java @@ -72,4 +72,27 @@ class ContractReturnValues { } private static final String SOMETHING = "foo"; + + void testCheckString() { + String s2 = nullable(); + checkString(s2); + if (s2 == null) { + System.out.println("impossible"); + } + String s = nullable(); + if (checkString(s) == null) { + 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); }