DFA: MethodReference handling refactoring; DfaCallArguments introduced (IDEA-CR-20606)

This commit is contained in:
Tagir Valeev
2017-04-26 17:41:53 +07:00
parent 95207fd561
commit 23b93b37ec
6 changed files with 121 additions and 69 deletions
@@ -1376,8 +1376,8 @@ public class ControlFlowAnalyzer extends JavaElementVisitor {
finishElement(expression);
}
private static List<? extends MethodContract> getMethodCallContracts(@NotNull final PsiMethod method,
@NotNull PsiMethodCallExpression call) {
static List<? extends MethodContract> getMethodCallContracts(@NotNull final PsiMethod method,
@Nullable PsiMethodCallExpression call) {
List<MethodContract> contracts = HardcodedContracts.getHardcodedContracts(method, call);
return !contracts.isEmpty() ? contracts : getMethodContracts(method);
}
@@ -37,73 +37,71 @@ import static com.siyeh.ig.callMatcher.CallMatcher.*;
*/
public class CustomMethodHandlers {
interface CustomMethodHandler {
List<DfaMemoryState> handle(DfaValue qualifier, DfaValue[] args, DfaMemoryState memState, DfaValueFactory factory);
List<DfaMemoryState> handle(DfaCallArguments callArguments, DfaMemoryState memState, DfaValueFactory factory);
}
private static final CallMapper<CustomMethodHandler> CUSTOM_METHOD_HANDLERS = new CallMapper<CustomMethodHandler>()
.register(instanceCall(JAVA_LANG_STRING, "isEmpty").parameterCount(0),
(qualifier, args, memState, factory) -> stringIsEmpty(qualifier, memState, factory))
(args, memState, factory) -> stringIsEmpty(args.myQualifier, memState, factory))
.register(instanceCall(JAVA_LANG_STRING, "indexOf", "lastIndexOf"),
(qualifier, args, memState, factory) -> stringIndexOf(qualifier, memState, factory))
(args, memState, factory) -> stringIndexOf(args.myQualifier, memState, factory))
.register(instanceCall(JAVA_LANG_STRING, "equals").parameterCount(1),
(qualifier, args, memState, factory) -> stringEquals(qualifier, args, memState, factory, false))
(args, memState, factory) -> stringEquals(args, memState, factory, false))
.register(instanceCall(JAVA_LANG_STRING, "equalsIgnoreCase").parameterCount(1),
(qualifier, args, memState, factory) -> stringEquals(qualifier, args, memState, factory, true))
(args, memState, factory) -> stringEquals(args, memState, factory, true))
.register(instanceCall(JAVA_LANG_STRING, "startsWith").parameterCount(1),
(qualifier, args, memState, factory) -> stringStartsEnds(qualifier, args, memState, factory, false))
(args, memState, factory) -> stringStartsEnds(args, memState, factory, false))
.register(instanceCall(JAVA_LANG_STRING, "endsWith").parameterCount(1),
(qualifier, args, memState, factory) -> stringStartsEnds(qualifier, args, memState, factory, true))
(args, memState, factory) -> stringStartsEnds(args, memState, factory, true))
.register(anyOf(staticCall(JAVA_LANG_MATH, "max").parameterTypes("int", "int"),
staticCall(JAVA_LANG_MATH, "max").parameterTypes("long", "long"),
staticCall(JAVA_LANG_INTEGER, "max").parameterTypes("int", "int"),
staticCall(JAVA_LANG_LONG, "max").parameterTypes("long", "long")),
(qualifier, args, memState, factory) -> mathMinMax(args, memState, factory, true))
(args, memState, factory) -> mathMinMax(args.myArguments, memState, factory, true))
.register(anyOf(staticCall(JAVA_LANG_MATH, "min").parameterTypes("int", "int"),
staticCall(JAVA_LANG_MATH, "min").parameterTypes("long", "long"),
staticCall(JAVA_LANG_INTEGER, "min").parameterTypes("int", "int"),
staticCall(JAVA_LANG_LONG, "min").parameterTypes("long", "long")),
(qualifier, args, memState, factory) -> mathMinMax(args, memState, factory, false))
(args, memState, factory) -> mathMinMax(args.myArguments, memState, factory, false))
.register(staticCall(JAVA_LANG_MATH, "abs").parameterTypes("int"),
(qualifier, args, memState, factory) -> mathAbs(args, memState, factory, false))
(args, memState, factory) -> mathAbs(args.myArguments, memState, factory, false))
.register(staticCall(JAVA_LANG_MATH, "abs").parameterTypes("long"),
(qualifier, args, memState, factory) -> mathAbs(args, memState, factory, true));
(args, memState, factory) -> mathAbs(args.myArguments, memState, factory, true));
public static CustomMethodHandler find(PsiMethodCallExpression call) {
return CUSTOM_METHOD_HANDLERS.mapFirst(call);
}
private static List<DfaMemoryState> stringStartsEnds(DfaValue qualifier,
DfaValue[] args,
private static List<DfaMemoryState> stringStartsEnds(DfaCallArguments args,
DfaMemoryState memState,
DfaValueFactory factory,
boolean ends) {
DfaValue arg = ArrayUtil.getFirstElement(args);
DfaValue arg = ArrayUtil.getFirstElement(args.myArguments);
if (arg == null) return Collections.emptyList();
String leftConst = ObjectUtils.tryCast(getConstantValue(memState, qualifier), String.class);
String leftConst = ObjectUtils.tryCast(getConstantValue(memState, args.myQualifier), String.class);
String rightConst = ObjectUtils.tryCast(getConstantValue(memState, arg), String.class);
if (leftConst != null && rightConst != null) {
return singleResult(memState, factory.getBoolean(ends ? leftConst.endsWith(rightConst) : leftConst.startsWith(rightConst)));
}
DfaValue leftLength = factory.createStringLength(qualifier);
DfaValue leftLength = factory.createStringLength(args.myQualifier);
DfaValue rightLength = factory.createStringLength(arg);
DfaValue trueRelation = factory.createCondition(leftLength, RelationType.GE, rightLength);
DfaValue falseRelation = factory.createCondition(leftLength, RelationType.LT, rightLength);
return applyCondition(memState, trueRelation, DfaUnknownValue.getInstance(), falseRelation, factory.getBoolean(false));
}
private static List<DfaMemoryState> stringEquals(DfaValue qualifier,
DfaValue[] args,
private static List<DfaMemoryState> stringEquals(DfaCallArguments args,
DfaMemoryState memState,
DfaValueFactory factory,
boolean ignoreCase) {
DfaValue arg = ArrayUtil.getFirstElement(args);
DfaValue arg = ArrayUtil.getFirstElement(args.myArguments);
if (arg == null) return Collections.emptyList();
String leftConst = ObjectUtils.tryCast(getConstantValue(memState, qualifier), String.class);
String leftConst = ObjectUtils.tryCast(getConstantValue(memState, args.myQualifier), String.class);
String rightConst = ObjectUtils.tryCast(getConstantValue(memState, arg), String.class);
if (leftConst != null && rightConst != null) {
return singleResult(memState, factory.getBoolean(ignoreCase ? leftConst.equalsIgnoreCase(rightConst) : leftConst.equals(rightConst)));
}
DfaValue leftLength = factory.createStringLength(qualifier);
DfaValue leftLength = factory.createStringLength(args.myQualifier);
DfaValue rightLength = factory.createStringLength(arg);
DfaValue trueRelation = factory.createCondition(leftLength, RelationType.EQ, rightLength);
DfaValue falseRelation = factory.createCondition(leftLength, RelationType.NE, rightLength);
@@ -0,0 +1,31 @@
/*
* Copyright 2000-2017 JetBrains s.r.o.
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* http://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package com.intellij.codeInspection.dataFlow;
import com.intellij.codeInspection.dataFlow.value.DfaValue;
/**
* @author Tagir Valeev
*/
/* package */ class DfaCallArguments {
final DfaValue myQualifier;
final DfaValue[] myArguments;
public DfaCallArguments(DfaValue qualifier, DfaValue[] arguments) {
myQualifier = qualifier;
myArguments = arguments;
}
}
@@ -15,11 +15,9 @@
*/
package com.intellij.codeInspection.dataFlow;
import com.intellij.codeInspection.dataFlow.value.DfaConstValue;
import com.intellij.codeInspection.dataFlow.value.*;
import com.intellij.codeInspection.dataFlow.value.DfaRelationValue.RelationType;
import com.intellij.codeInspection.dataFlow.value.DfaValue;
import com.intellij.codeInspection.dataFlow.value.DfaValueFactory;
import com.intellij.psi.PsiType;
import org.jetbrains.annotations.NotNull;
import org.jetbrains.annotations.Nullable;
import java.util.List;
@@ -43,21 +41,24 @@ public abstract class MethodContract {
public abstract ValueConstraint getReturnValue();
/**
* Returns DfaValue describing the return value of this contract or null if this contract assumes that any value could be returned.
* Returns DfaValue describing the return value of this contract
*
* @param factory factory to create values
* @param resultType method result type
* @param defaultResult default result value for the called method
* @return a DfaValue describing the return value of this contract
*/
@Nullable
DfaValue getDfaReturnValue(DfaValueFactory factory, PsiType resultType) {
@NotNull
DfaValue getDfaReturnValue(DfaValueFactory factory, DfaValue defaultResult) {
switch (getReturnValue()) {
case NULL_VALUE: return factory.getConstFactory().getNull();
case NOT_NULL_VALUE: return factory.createTypeValue(resultType, Nullness.NOT_NULL);
case NOT_NULL_VALUE:
return defaultResult instanceof DfaTypeValue
? ((DfaTypeValue)defaultResult).withNullness(Nullness.NOT_NULL)
: DfaUnknownValue.getInstance();
case TRUE_VALUE: return factory.getConstFactory().getTrue();
case FALSE_VALUE: return factory.getConstFactory().getFalse();
case THROW_EXCEPTION: return factory.getConstFactory().getContractFail();
default: return null;
default: return defaultResult;
}
}
@@ -40,6 +40,7 @@ import org.jetbrains.annotations.NotNull;
import org.jetbrains.annotations.Nullable;
import java.util.*;
import java.util.stream.Stream;
/**
* @author peter
@@ -147,12 +148,23 @@ public class StandardInstructionVisitor extends InstructionVisitor {
JavaResolveResult resolveResult = methodRef.advancedResolve(false);
PsiMethod method = ObjectUtils.tryCast(resolveResult.getElement(), PsiMethod.class);
if (method == null || !ControlFlowAnalyzer.isPure(method)) return;
List<? extends MethodContract> contracts = HardcodedContracts.getHardcodedContracts(method, null);
if (contracts.isEmpty()) {
contracts = ControlFlowAnalyzer.getMethodContracts(method);
}
List<? extends MethodContract> contracts = ControlFlowAnalyzer.getMethodCallContracts(method, null);
if (contracts.isEmpty()) return;
PsiSubstitutor substitutor = resolveResult.getSubstitutor();
DfaCallArguments callArguments = getMethodReferenceCallArguments(methodRef, qualifier, runner, sam, method, substitutor);
PsiType returnType = substitutor.substitute(method.getReturnType());
DfaValue defaultResult = runner.getFactory().createTypeValue(returnType, DfaPsiUtil.getElementNullability(returnType, method));
Stream<DfaValue> returnValues = possibleReturnValues(callArguments, state, contracts, runner.getFactory(), defaultResult);
returnValues.forEach(res -> processMethodReferenceResult(methodRef, res));
}
@NotNull
private static DfaCallArguments getMethodReferenceCallArguments(PsiMethodReferenceExpression methodRef,
DfaValue qualifier,
DataFlowRunner runner,
PsiMethod sam,
PsiMethod method,
PsiSubstitutor substitutor) {
PsiParameter[] samParameters = sam.getParameterList().getParameters();
boolean isStatic = method.hasModifierProperty(PsiModifier.STATIC);
boolean instanceBound = !isStatic && !PsiMethodReferenceUtil.isStaticallyReferenced(methodRef);
@@ -173,20 +185,21 @@ public class StandardInstructionVisitor extends InstructionVisitor {
}
}
}
return new DfaCallArguments(qualifier, arguments);
}
private static Stream<DfaValue> possibleReturnValues(DfaCallArguments callArguments,
DfaMemoryState state,
List<? extends MethodContract> contracts,
DfaValueFactory factory, DfaValue defaultResult) {
LinkedHashSet<DfaMemoryState> currentStates = ContainerUtil.newLinkedHashSet(state.createClosureState());
Set<DfaMemoryState> finalStates = ContainerUtil.newLinkedHashSet();
PsiType returnType = substitutor.substitute(method.getReturnType());
DfaValue defaultResult = runner.getFactory().createTypeValue(returnType, DfaPsiUtil.getElementNullability(returnType, method));
for (MethodContract contract : contracts) {
DfaValue result = contract.getDfaReturnValue(runner.getFactory(), returnType);
if (result == null) {
result = defaultResult;
}
currentStates = addContractResults(qualifier, arguments, contract, currentStates, runner.getFactory(), finalStates, result);
DfaValue result = contract.getDfaReturnValue(factory, defaultResult);
currentStates = addContractResults(callArguments, contract, currentStates, factory, finalStates, result);
}
StreamEx.of(finalStates).map(DfaMemoryState::peek)
.append(currentStates.isEmpty() ? StreamEx.empty() : StreamEx.of(defaultResult))
.distinct().forEach(res -> processMethodReferenceResult(methodRef, res));
return StreamEx.of(finalStates).map(DfaMemoryState::peek)
.append(currentStates.isEmpty() ? StreamEx.empty() : StreamEx.of(defaultResult)).distinct();
}
protected void processMethodReferenceResult(PsiMethodReferenceExpression methodRef, DfaValue res) {
@@ -266,17 +279,14 @@ public class StandardInstructionVisitor extends InstructionVisitor {
finalStates.addAll(handleKnownMethods(instruction, runner, memState));
if (finalStates.isEmpty()) {
DfaValue[] argValues = popCallArguments(instruction, runner, memState, true);
final DfaValue qualifier = popQualifier(instruction, runner, memState);
DfaCallArguments callArguments = popCall(instruction, runner, memState, true);
LinkedHashSet<DfaMemoryState> currentStates = ContainerUtil.newLinkedHashSet(memState);
if (argValues != null) {
if (callArguments.myArguments != null) {
for (MethodContract contract : instruction.getContracts()) {
DfaValue returnValue = contract.getDfaReturnValue(runner.getFactory(), instruction.getResultType());
if (returnValue == null) {
returnValue = getMethodResultValue(instruction, qualifier, runner.getFactory());
}
currentStates = addContractResults(qualifier, argValues, contract, currentStates, runner.getFactory(), finalStates, returnValue);
DfaValue returnValue = getMethodResultValue(instruction, callArguments.myQualifier, runner.getFactory());
returnValue = contract.getDfaReturnValue(runner.getFactory(), returnValue);
currentStates = addContractResults(callArguments, contract, currentStates, runner.getFactory(), finalStates, returnValue);
if (currentStates.size() + finalStates.size() > DataFlowRunner.MAX_STATES_PER_BRANCH) {
if (LOG.isDebugEnabled()) {
LOG.debug("Too complex contract on " + instruction.getContext() + ", skipping contract processing");
@@ -288,7 +298,7 @@ public class StandardInstructionVisitor extends InstructionVisitor {
}
}
for (DfaMemoryState state : currentStates) {
state.push(getMethodResultValue(instruction, qualifier, runner.getFactory()));
state.push(getMethodResultValue(instruction, callArguments.myQualifier, runner.getFactory()));
finalStates.add(state);
}
}
@@ -309,13 +319,12 @@ public class StandardInstructionVisitor extends InstructionVisitor {
PsiMethodCallExpression call = ObjectUtils.tryCast(instruction.getCallExpression(), PsiMethodCallExpression.class);
CustomMethodHandlers.CustomMethodHandler handler = CustomMethodHandlers.find(call);
if (handler == null) return Collections.emptyList();
DfaValue[] arguments = popCallArguments(instruction, runner, memState, false);
DfaValue qualifier = popQualifier(instruction, runner, memState);
DfaCallArguments callArguments = popCall(instruction, runner, memState, false);
List<DfaMemoryState> states =
arguments == null ? Collections.emptyList() :
handler.handle(qualifier, arguments, memState, runner.getFactory());
callArguments.myArguments == null ? Collections.emptyList() :
handler.handle(callArguments, memState, runner.getFactory());
if (states.isEmpty()) {
memState.push(getMethodResultValue(instruction, qualifier, runner.getFactory()));
memState.push(getMethodResultValue(instruction, callArguments.myQualifier, runner.getFactory()));
return Collections.singletonList(memState);
}
return states;
@@ -330,8 +339,8 @@ public class StandardInstructionVisitor extends InstructionVisitor {
PsiMethod method = call.resolveMethod();
if (method == null || !TypeUtils.isOptional(method.getContainingClass())) return Collections.emptyList();
List<DfaMemoryState> closures = runner.getStackTopClosures();
DfaValue[] argValues = popCallArguments(instruction, runner, memState, false);
DfaValue qualifier = popQualifier(instruction, runner, memState);
DfaCallArguments arguments = popCall(instruction, runner, memState, false);
DfaValue[] argValues = arguments.myArguments;
DfaValue result = null;
DfaValueFactory factory = runner.getFactory();
switch (methodName) {
@@ -350,7 +359,7 @@ public class StandardInstructionVisitor extends InstructionVisitor {
if (argValues != null && argValues.length == 1) {
DfaMemoryState falseState = memState.createCopy();
DfaOptionalValue optional = factory.getOptionalFactory().getOptional(true);
DfaValue relation = factory.createCondition(qualifier, RelationType.IS, optional);
DfaValue relation = factory.createCondition(arguments.myQualifier, RelationType.IS, optional);
List<DfaMemoryState> states = new ArrayList<>(2);
if (memState.applyCondition(relation)) {
memState.push(factory.createTypeValue(instruction.getResultType(), Nullness.NOT_NULL));
@@ -371,7 +380,7 @@ public class StandardInstructionVisitor extends InstructionVisitor {
case "orElseGet":
case "transform": {
DfaOptionalValue optional = factory.getOptionalFactory().getOptional(!methodName.startsWith("or"));
DfaValue relation = factory.createCondition(qualifier, RelationType.IS, optional);
DfaValue relation = factory.createCondition(arguments.myQualifier, RelationType.IS, optional);
for (DfaMemoryState closure : closures) {
closure.applyCondition(relation);
}
@@ -379,10 +388,20 @@ public class StandardInstructionVisitor extends InstructionVisitor {
}
default:
}
memState.push(result == null ? getMethodResultValue(instruction, qualifier, factory) : result);
memState.push(result == null ? getMethodResultValue(instruction, arguments.myQualifier, factory) : result);
return Collections.singletonList(memState);
}
@NotNull
private DfaCallArguments popCall(MethodCallInstruction instruction,
DataFlowRunner runner,
DfaMemoryState memState,
boolean contractOnly) {
DfaValue[] argValues = popCallArguments(instruction, runner, memState, contractOnly);
final DfaValue qualifier = popQualifier(instruction, runner, memState);
return new DfaCallArguments(qualifier, argValues);
}
@Nullable
private DfaValue[] popCallArguments(MethodCallInstruction instruction,
DataFlowRunner runner,
@@ -440,14 +459,13 @@ public class StandardInstructionVisitor extends InstructionVisitor {
return qualifier;
}
private static LinkedHashSet<DfaMemoryState> addContractResults(DfaValue qualifier,
DfaValue[] argValues,
private static LinkedHashSet<DfaMemoryState> addContractResults(DfaCallArguments callParameters,
MethodContract contract,
LinkedHashSet<DfaMemoryState> states,
DfaValueFactory factory,
Set<DfaMemoryState> finalStates,
DfaValue returnValue) {
List<DfaValue> conditions = contract.getConditions(factory, qualifier, argValues);
List<DfaValue> conditions = contract.getConditions(factory, callParameters.myQualifier, callParameters.myArguments);
if (StreamEx.of(conditions).allMatch(factory.getConstFactory().getTrue()::equals)) {
for (DfaMemoryState state : states) {
state.push(returnValue);
@@ -91,6 +91,10 @@ public class DfaTypeValue extends DfaValue {
return myNullness;
}
public DfaTypeValue withNullness(Nullness nullness) {
return nullness == myNullness ? this : myFactory.getTypeFactory().createTypeValue(myType, nullness);
}
@NonNls
public String toString() {
return myType + ", nullable=" + myNullness;