diff --git a/java/java-analysis-impl/src/com/intellij/codeInspection/dataFlow/ContractInference.java b/java/java-analysis-impl/src/com/intellij/codeInspection/dataFlow/ContractInference.java index b76177e58880..64b11f231eeb 100644 --- a/java/java-analysis-impl/src/com/intellij/codeInspection/dataFlow/ContractInference.java +++ b/java/java-analysis-impl/src/com/intellij/codeInspection/dataFlow/ContractInference.java @@ -16,10 +16,13 @@ package com.intellij.codeInspection.dataFlow; import com.intellij.codeInspection.dataFlow.MethodContract.ValueConstraint; +import com.intellij.openapi.util.Computable; import com.intellij.openapi.util.Condition; +import com.intellij.openapi.util.RecursionManager; import com.intellij.psi.*; import com.intellij.psi.tree.IElementType; import com.intellij.util.Function; +import com.intellij.util.NullableFunction; import com.intellij.util.containers.ContainerUtil; import org.jetbrains.annotations.NotNull; import org.jetbrains.annotations.Nullable; @@ -49,15 +52,78 @@ class ContractInferenceInterpreter { List inferContracts() { PsiCodeBlock body = myMethod.getBody(); - if (body == null) return Collections.emptyList(); - - PsiStatement[] statements = body.getStatements(); + PsiStatement[] statements = body == null ? PsiStatement.EMPTY_ARRAY : body.getStatements(); if (statements.length == 0) return Collections.emptyList(); + if (statements.length == 1) { + if (statements[0] instanceof PsiReturnStatement) { + List result = handleDelegation(((PsiReturnStatement)statements[0]).getReturnValue(), false); + if (result != null) return result; + } + else if (statements[0] instanceof PsiExpressionStatement && ((PsiExpressionStatement)statements[0]).getExpression() instanceof PsiMethodCallExpression) { + List result = handleDelegation(((PsiExpressionStatement)statements[0]).getExpression(), false); + if (result != null) return result; + } + } + ValueConstraint[] emptyState = MethodContract.createConstraintArray(myMethod.getParameterList().getParametersCount()); return visitStatements(Collections.singletonList(emptyState), statements); } + @Nullable + private List handleDelegation(final PsiExpression expression, final boolean negated) { + if (expression instanceof PsiParenthesizedExpression) { + return handleDelegation(((PsiParenthesizedExpression)expression).getExpression(), negated); + } + + if (expression instanceof PsiPrefixExpression && ((PsiPrefixExpression)expression).getOperationTokenType() == JavaTokenType.EXCL) { + return handleDelegation(((PsiPrefixExpression)expression).getOperand(), !negated); + } + + if (expression instanceof PsiMethodCallExpression) { + return handleCallDelegation((PsiMethodCallExpression)expression, negated); + } + + return null; + } + + private List handleCallDelegation(PsiMethodCallExpression expression, final boolean negated) { + final PsiMethod targetMethod = expression.resolveMethod(); + if (targetMethod == null) return Collections.emptyList(); + + final PsiExpression[] arguments = expression.getArgumentList().getExpressions(); + return RecursionManager.doPreventingRecursion(myMethod, true, new Computable>() { + @Override + public List compute() { + List delegateContracts = ContractInference.inferContracts(targetMethod); //todo use explicit contracts, too + return ContainerUtil.mapNotNull(delegateContracts, new NullableFunction() { + @Nullable + @Override + public MethodContract fun(MethodContract delegateContract) { + ValueConstraint[] answer = MethodContract.createConstraintArray(myMethod.getParameterList().getParametersCount()); + for (int i = 0; i < delegateContract.arguments.length; i++) { + if (i >= arguments.length) return null; + + ValueConstraint argConstraint = delegateContract.arguments[i]; + if (argConstraint != ANY_VALUE) { + int paramIndex = resolveParameter(arguments[i]); + if (paramIndex < 0) { + if (argConstraint != getLiteralConstraint(arguments[i])) { + return null; + } + } + else { + answer = withConstraint(answer, paramIndex, argConstraint); + } + } + } + return new MethodContract(answer, negated ? negateConstraint(delegateContract.returnValue) : delegateContract.returnValue); + } + }); + } + }); + } + @NotNull private List visitExpression(final List states, @Nullable PsiExpression expr) { if (states.isEmpty()) return Collections.emptyList(); diff --git a/java/java-tests/testSrc/com/intellij/codeInspection/ContractInferenceFromSourceTest.groovy b/java/java-tests/testSrc/com/intellij/codeInspection/ContractInferenceFromSourceTest.groovy index ab9a5660b8ad..32982f3687df 100644 --- a/java/java-tests/testSrc/com/intellij/codeInspection/ContractInferenceFromSourceTest.groovy +++ b/java/java-tests/testSrc/com/intellij/codeInspection/ContractInferenceFromSourceTest.groovy @@ -171,6 +171,55 @@ class ContractInferenceFromSourceTest extends LightCodeInsightFixtureTestCase { assert c == [] } + public void "test plain delegation"() { + def c = inferContracts(""" + boolean delegating(Object o) { + return smth(o); + } + boolean smth(Object o) { + assert o instanceof String; + return true; + } +""") + assert c == ['null -> fail'] + } + + public void "test arg swapping delegation"() { + def c = inferContracts(""" + boolean delegating(Object o, Object o1) { + return smth(o1, o); + } + boolean smth(Object o, Object o1) { + return o == null && o1 != null; + } +""") + assert c == ['_, !null -> false', 'null, null -> false', '!null, null -> true'] + } + + public void "test negating delegation"() { + def c = inferContracts(""" + boolean delegating(Object o) { + return !smth(o); + } + boolean smth(Object o) { + return o == null; + } +""") + assert c == ['null -> false', '!null -> true'] + } + + public void "test delegation with constant"() { + def c = inferContracts(""" + boolean delegating(Object o) { + return smth(null); + } + boolean smth(Object o) { + return o == null; + } +""") + assert c == ['_ -> true'] + } + private String inferContract(String method) { return assertOneElement(inferContracts(method)) }