extract method: allow to extract @NotNull code with missed branches which would be replaced with 'return null' with check on call site (IDEA-120076)

This commit is contained in:
Anna Kozlova
2014-11-27 18:12:53 +01:00
parent 6ff126e422
commit 83d3df6a95
6 changed files with 112 additions and 41 deletions
@@ -143,8 +143,12 @@ public class ControlFlowWrapper {
public PsiVariable[] getOutputVariables() {
return getOutputVariables(myGenerateConditionalExit);
}
public PsiVariable[] getOutputVariables(boolean collectVariablesAtExitPoints) {
PsiVariable[] myOutputVariables = ControlFlowUtil.getOutputVariables(myControlFlow, myFlowStart, myFlowEnd, myExitPoints.toArray());
if (myGenerateConditionalExit) {
if (collectVariablesAtExitPoints) {
//variables declared in selected block used in return statements are to be considered output variables when extracting guard methods
final Set<PsiVariable> outputVariables = new HashSet<PsiVariable>(Arrays.asList(myOutputVariables));
for (PsiStatement statement : myExitStatements) {
@@ -17,11 +17,9 @@ package com.intellij.refactoring.extractMethod;
import com.intellij.codeInsight.ChangeContextUtil;
import com.intellij.codeInsight.ExceptionUtil;
import com.intellij.codeInsight.NullableNotNullManager;
import com.intellij.codeInsight.daemon.impl.analysis.JavaHighlightUtil;
import com.intellij.codeInsight.daemon.impl.quickfix.AnonymousTargetClassPreselectionUtil;
import com.intellij.codeInsight.highlighting.HighlightManager;
import com.intellij.codeInsight.intention.impl.AddNullableAnnotationFix;
import com.intellij.codeInsight.navigation.NavigationUtil;
import com.intellij.codeInspection.dataFlow.RunnerResult;
import com.intellij.codeInspection.dataFlow.StandardDataFlowRunner;
@@ -40,9 +38,7 @@ import com.intellij.openapi.editor.colors.EditorColorsManager;
import com.intellij.openapi.editor.markup.TextAttributes;
import com.intellij.openapi.progress.ProgressManager;
import com.intellij.openapi.project.Project;
import com.intellij.openapi.util.Comparing;
import com.intellij.openapi.util.Pass;
import com.intellij.openapi.util.TextRange;
import com.intellij.openapi.util.*;
import com.intellij.openapi.util.text.StringUtil;
import com.intellij.openapi.vfs.VirtualFile;
import com.intellij.openapi.wm.WindowManager;
@@ -132,6 +128,7 @@ public class ExtractMethodProcessor implements MatchProvider {
private PsiMethod myExtractedMethod;
private PsiMethodCallExpression myMethodCall;
private boolean myNullConditionalCheck = false;
private boolean myNotNullConditionalCheck = false;
public ExtractMethodProcessor(Project project,
Editor editor,
@@ -263,12 +260,12 @@ public class ExtractMethodProcessor implements MatchProvider {
}
myHasExpressionOutput = expressionType != PsiType.VOID;
PsiType returnStatementType = null;
if (myHasReturnStatement) {
returnStatementType = myCodeFragmentMember instanceof PsiMethod ? ((PsiMethod)myCodeFragmentMember).getReturnType()
: myCodeFragmentMember instanceof PsiLambdaExpression ? LambdaUtil.getFunctionalInterfaceReturnType((PsiLambdaExpression)myCodeFragmentMember) : null;
}
myHasReturnStatementOutput = returnStatementType != null && returnStatementType != PsiType.VOID;
final PsiType returnStatementType = myCodeFragmentMember instanceof PsiMethod
? ((PsiMethod)myCodeFragmentMember).getReturnType()
: myCodeFragmentMember instanceof PsiLambdaExpression
? LambdaUtil.getFunctionalInterfaceReturnType((PsiLambdaExpression)myCodeFragmentMember)
: null;
myHasReturnStatementOutput = myHasReturnStatement && returnStatementType != null && returnStatementType != PsiType.VOID;
if (myGenerateConditionalExit && myOutputVariables.length == 1) {
if (!(myOutputVariables[0].getType() instanceof PsiPrimitiveType)) {
@@ -276,21 +273,33 @@ public class ExtractMethodProcessor implements MatchProvider {
for (PsiStatement exitStatement : myExitStatements) {
if (exitStatement instanceof PsiReturnStatement) {
final PsiExpression returnValue = ((PsiReturnStatement)exitStatement).getReturnValue();
myNullConditionalCheck &= returnValue == null ||
returnValue instanceof PsiLiteralExpression && PsiType.NULL.equals(returnValue.getType());
myNullConditionalCheck &= returnValue == null || isNullInferred(returnValue.getText(), true);
}
}
myNullConditionalCheck &= isNullInferred(myOutputVariables[0].getName(), false);
}
if (insertNotNullCheckIfPossible() && myControlFlowWrapper.getOutputVariables(false).length == 0) {
myNotNullConditionalCheck = returnStatementType != null && returnStatementType != PsiType.VOID;
for (PsiStatement statement : myExitStatements) {
if (statement instanceof PsiReturnStatement) {
final PsiExpression returnValue = ((PsiReturnStatement)statement).getReturnValue();
myNotNullConditionalCheck &= returnValue != null && !isNullInferred(returnValue.getText(), true);
}
}
myNullConditionalCheck &= isNotNull(myOutputVariables[0]);
}
}
if (!myHasReturnStatementOutput && checkOutputVariablesCount() && !myNullConditionalCheck) {
if (!myHasReturnStatementOutput && checkOutputVariablesCount() && !myNullConditionalCheck && !myNotNullConditionalCheck) {
showMultipleOutputMessage(expressionType);
return false;
}
myOutputVariable = myOutputVariables.length > 0 ? myOutputVariables[0] : null;
if (myHasReturnStatementOutput) {
if (myNotNullConditionalCheck) {
myReturnType = returnStatementType instanceof PsiPrimitiveType ? ((PsiPrimitiveType)returnStatementType).getBoxedType(myCodeFragmentMember)
: returnStatementType;
} else if (myHasReturnStatementOutput) {
myReturnType = returnStatementType;
}
else if (myOutputVariable != null) {
@@ -328,20 +337,25 @@ public class ExtractMethodProcessor implements MatchProvider {
return true;
}
private boolean isNotNull(PsiVariable outputVariable) {
final PsiCodeBlock block = myElementFactory.createCodeBlock();
protected boolean insertNotNullCheckIfPossible() {
return true;
}
private boolean isNullInferred(String exprText, boolean trueSet) {
final PsiCodeBlock block = myElementFactory.createCodeBlockFromText("{}", myElements[0]);
for (PsiElement element : myElements) {
block.add(element);
}
final PsiIfStatement statementFromText = (PsiIfStatement)myElementFactory.createStatementFromText("if (" + outputVariable.getName() + " == null);", null);
final PsiIfStatement statementFromText = (PsiIfStatement)myElementFactory.createStatementFromText("if (" + exprText + " == null);", null);
block.add(statementFromText);
final StandardDataFlowRunner dfaRunner = new StandardDataFlowRunner();
final StandardInstructionVisitor visitor = new StandardInstructionVisitor();
final RunnerResult rc = dfaRunner.analyzeMethod(block, visitor);
if (rc == RunnerResult.OK) {
final Set<Instruction> falseSet = dfaRunner.getConstConditionalExpressions().getSecond();
for (Instruction instruction : falseSet) {
final Pair<Set<Instruction>, Set<Instruction>> expressions = dfaRunner.getConstConditionalExpressions();
final Set<Instruction> set = trueSet ? expressions.getFirst() : expressions.getSecond();
for (Instruction instruction : set) {
if (instruction instanceof BranchingInstruction) {
if (((BranchingInstruction)instruction).getPsiAnchor().getText().equals(statementFromText.getCondition().getText())) {
return true;
@@ -667,7 +681,7 @@ public class ExtractMethodProcessor implements MatchProvider {
hasNormalExit = true;
}
PsiStatement exitStatementCopy = myControlFlowWrapper.getExitStatementCopy(returnStatement, myElements);
PsiStatement exitStatementCopy = myNotNullConditionalCheck ? null : myControlFlowWrapper.getExitStatementCopy(returnStatement, myElements);
declareNecessaryVariablesInsideBody(body);
@@ -683,7 +697,11 @@ public class ExtractMethodProcessor implements MatchProvider {
body.addRange(myElements[0], myElements[myElements.length - 1]);
if (myNullConditionalCheck) {
body.add(myElementFactory.createStatementFromText("return " + myOutputVariable.getName() + ";", null));
} else if (myGenerateConditionalExit) {
}
else if (myNotNullConditionalCheck) {
body.add(myElementFactory.createStatementFromText("return null;", null));
}
else if (myGenerateConditionalExit) {
body.add(myElementFactory.createStatementFromText("return false;", null));
}
else if (!myHasReturnStatement && hasNormalExit && myOutputVariable != null) {
@@ -736,6 +754,11 @@ public class ExtractMethodProcessor implements MatchProvider {
ifStatement = (PsiIfStatement)addToMethodCallLocation(ifStatement);
CodeStyleManager.getInstance(myProject).reformat(ifStatement);
}
else if (myNotNullConditionalCheck) {
final String varName = myOutputVariable.getName();
declareVariableAtMethodCallLocation(varName, myReturnType instanceof PsiPrimitiveType ? ((PsiPrimitiveType)myReturnType).getBoxedType(myCodeFragmentMember) : myReturnType);
addToMethodCallLocation(myElementFactory.createStatementFromText("if (" + varName + " != null) return " + varName + ";", null));
}
else if (myGenerateConditionalExit) {
PsiIfStatement ifStatement = (PsiIfStatement)myElementFactory.createStatementFromText("if (a) b;", null);
ifStatement = (PsiIfStatement)addToMethodCallLocation(ifStatement);
@@ -775,7 +798,7 @@ public class ExtractMethodProcessor implements MatchProvider {
addToMethodCallLocation(exitStatementCopy);
}
if (!myNullConditionalCheck) {
if (!myNullConditionalCheck && !myNotNullConditionalCheck) {
declareNecessaryVariablesAfterCall(myOutputVariable);
}
@@ -855,8 +878,11 @@ public class ExtractMethodProcessor implements MatchProvider {
}
private void declareVariableAtMethodCallLocation(String name) {
PsiDeclarationStatement statement =
myElementFactory.createVariableDeclarationStatement(name, myOutputVariable.getType(), myMethodCall);
declareVariableAtMethodCallLocation(name, myOutputVariable.getType());
}
private void declareVariableAtMethodCallLocation(String name, PsiType type) {
PsiDeclarationStatement statement = myElementFactory.createVariableDeclarationStatement(name, type, myMethodCall);
statement = (PsiDeclarationStatement)addToMethodCallLocation(statement);
PsiVariable var = (PsiVariable)statement.getDeclaredElements()[0];
myMethodCall = (PsiMethodCallExpression)var.getInitializer();
@@ -1067,18 +1093,6 @@ public class ExtractMethodProcessor implements MatchProvider {
throwsList.add(JavaPsiFacade.getInstance(myManager.getProject()).getElementFactory().createReferenceElementByType(exception));
}
if (myNullConditionalCheck) {
final boolean isNullCheckReturnNull = (myHasExpressionOutput ? 1 : 0) + (myGenerateConditionalExit ? 1 : 0) + myOutputVariables.length <= 1;
if (isNullCheckReturnNull && PsiUtil.isLanguageLevel5OrHigher(myElements[0])) {
final NullableNotNullManager manager = NullableNotNullManager.getInstance(myProject);
final PsiClass nullableAnnotationClass =
JavaPsiFacade.getInstance(myProject).findClass(manager.getDefaultNullable(), GlobalSearchScope.allScope(myProject));
if (nullableAnnotationClass != null) {
new AddNullableAnnotationFix(newMethod).invoke(myProject, myTargetClass.getContainingFile(), newMethod, newMethod);
}
}
}
if (myTargetClass.isInterface() && PsiUtil.isLanguageLevel8OrHigher(myTargetClass)) {
PsiUtil.setModifierProperty(newMethod, PsiModifier.DEFAULT, true);
}
@@ -652,6 +652,11 @@ public class ExtractMethodObjectProcessor extends BaseRefactoringProcessor {
}
@Override
protected boolean insertNotNullCheckIfPossible() {
return false;
}
@Override
protected void apply(final AbstractExtractDialog dialog) {
super.apply(dialog);
@@ -0,0 +1,16 @@
class C {
public Object m() {
Object o = newMethod();
if (o != null) return o;
return null;
}
private Object newMethod() {
for (Object o : new ArrayList<Object>()) {
if (o != null) {
return o;
}
}
return null;
}
}
@@ -0,0 +1,32 @@
class Result {
private String _message;
public Result(String _message) {
this._message = _message;
}
}
class Main {
public static Result doIt(String name) {
Result result;
Result result = newMethod(name);
if (result != null) return result;
result = new Result("Name is " + name);
return result;
}
private static Result newMethod(String name) {
Result result;
if (name == null) {
result = new Result("Name is null");
return result;
}
if (name.length() == 0) {
result = new Result("Name is empty");
return result;
}
return null;
}
}
@@ -80,11 +80,11 @@ public class ExtractMethodTest extends LightCodeInsightTestCase {
}
public void testExitPoints8() throws Exception {
doExitPointsTest(false);
doTest();
}
public void testExitPoints9() throws Exception {
doExitPointsTest(false);
doTest();
}
public void testContinueInside() throws Exception {