added Replace duplicates in Extract Method

This commit is contained in:
Ekaterina Tuzova
2013-04-18 15:47:29 +04:00
parent 188f394f72
commit d0de628451
8 changed files with 221 additions and 14 deletions
@@ -536,6 +536,7 @@ refactoring.extract.method.error.cannot.perform.refactoring.when.execution.flow.
refactoring.extract.method.error.cannot.perform.refactoring.when.from.import.inside=Cannot perform refactoring with from import statement inside code block
refactoring.extract.method.error.cannot.perform.refactoring.using.selected.elements=Cannot perform extract method using selected element(s)
refactoring.extract.method.error.name.clash=Method name clashes with already existing name
refactoring.extract.method.error.cannot.perform.refactoring.with.local=Cannot perform refactoring from expression with local variables modifications and return instructions inside code fragment
# extract superclass
refactoring.extract.super.target.path.outside.roots=Target directory is outside the project.<br>Must be within content roots
@@ -0,0 +1,97 @@
package com.jetbrains.python.refactoring.extractmethod;
import com.intellij.codeInsight.PsiEquivalenceUtil;
import com.intellij.openapi.util.Pair;
import com.intellij.psi.PsiComment;
import com.intellij.psi.PsiElement;
import com.intellij.psi.PsiWhiteSpace;
import com.intellij.psi.util.PsiTreeUtil;
import com.jetbrains.python.psi.PyFunction;
import org.jetbrains.annotations.NotNull;
import org.jetbrains.annotations.Nullable;
import java.util.ArrayList;
import java.util.List;
/**
* User : ktisha
*/
public class PyDuplicatesFinder {
private final ArrayList<PsiElement> myPattern;
public PyDuplicatesFinder(@NotNull final PsiElement statement1, @NotNull final PsiElement statement2) {
myPattern = new ArrayList<PsiElement>();
PsiElement sibling = statement1;
do {
myPattern.add(sibling);
if (sibling == statement2) break;
sibling = PsiTreeUtil.skipSiblingsForward(sibling, PsiWhiteSpace.class, PsiComment.class);
} while (sibling != null);
}
public List<Pair<PsiElement, PsiElement>> findDuplicates(@Nullable final PsiElement scope, @NotNull final PyFunction generatedMethod) {
final ArrayList<Pair<PsiElement, PsiElement>> result = new ArrayList<Pair<PsiElement, PsiElement>>();
if (scope != null) {
findPatternOccurrences(result, scope, generatedMethod);
}
return result;
}
private void findPatternOccurrences(@NotNull final List<Pair<PsiElement, PsiElement>> array, @NotNull final PsiElement scope,
@NotNull final PyFunction generatedMethod) {
if (scope == generatedMethod) return;
final PsiElement[] children = scope.getChildren();
for (PsiElement child : children) {
final Pair<PsiElement, PsiElement> match = isDuplicateFragment(child);
if (match != null) {
array.add(match);
continue;
}
findPatternOccurrences(array, child, generatedMethod);
}
}
@Nullable
private Pair<PsiElement, PsiElement> isDuplicateFragment(@NotNull final PsiElement candidate) {
for (PsiElement pattern : myPattern) {
if (PsiTreeUtil.isAncestor(pattern, candidate, false)) return null;
}
PsiElement sibling = candidate;
final ArrayList<PsiElement> candidates = new ArrayList<PsiElement>();
for (int i = 0; i != myPattern.size(); ++i) {
if (sibling == null) return null;
candidates.add(sibling);
sibling = PsiTreeUtil.skipSiblingsForward(sibling, PsiWhiteSpace.class, PsiComment.class);
}
if (myPattern.size() != candidates.size()) return null;
if (candidates.size() <= 0) return null;
final Pair<PsiElement, PsiElement> match = new Pair<PsiElement, PsiElement>(candidates.get(0), candidates.get(candidates.size() - 1));
for (int i = 0; i < myPattern.size(); i++) {
if (!matchPattern(myPattern.get(i), candidates.get(i))) return null;
}
return match;
}
private static boolean matchPattern(@Nullable final PsiElement pattern,
@Nullable final PsiElement candidate) {
if (pattern == null || candidate == null) return pattern == candidate;
final PsiElement[] children1 = PsiEquivalenceUtil.getFilteredChildren(pattern, null, true);
final PsiElement[] children2 = PsiEquivalenceUtil.getFilteredChildren(candidate, null, true);
if (children1.length != children2.length) return false;
for (int i = 0; i < children1.length; i++) {
PsiElement child1 = children1[i];
PsiElement child2 = children2[i];
if (!matchPattern(child1, child2)) return false;
}
if (children1.length == 0) {
if (!pattern.textMatches(candidate)) return false;
}
return true;
}
}
@@ -1,20 +1,27 @@
package com.jetbrains.python.refactoring.extractmethod;
import com.intellij.codeInsight.CodeInsightUtilBase;
import com.intellij.codeInsight.codeFragment.CodeFragment;
import com.intellij.codeInsight.highlighting.HighlightManager;
import com.intellij.find.FindManager;
import com.intellij.lang.LanguageNamesValidation;
import com.intellij.lang.refactoring.NamesValidator;
import com.intellij.openapi.application.ApplicationManager;
import com.intellij.openapi.application.ApplicationNamesInfo;
import com.intellij.openapi.command.CommandProcessor;
import com.intellij.openapi.editor.Editor;
import com.intellij.openapi.editor.LogicalPosition;
import com.intellij.openapi.editor.ScrollType;
import com.intellij.openapi.editor.colors.EditorColors;
import com.intellij.openapi.editor.colors.EditorColorsManager;
import com.intellij.openapi.editor.markup.TextAttributes;
import com.intellij.openapi.project.Project;
import com.intellij.openapi.ui.DialogWrapper;
import com.intellij.openapi.ui.Messages;
import com.intellij.openapi.util.Pair;
import com.intellij.openapi.util.TextRange;
import com.intellij.openapi.util.text.StringUtil;
import com.intellij.psi.PsiElement;
import com.intellij.psi.PsiNamedElement;
import com.intellij.psi.PsiRecursiveElementVisitor;
import com.intellij.psi.PsiWhiteSpace;
import com.intellij.psi.*;
import com.intellij.psi.codeStyle.CodeStyleManager;
import com.intellij.psi.impl.source.codeStyle.CodeEditUtil;
import com.intellij.psi.util.PsiTreeUtil;
@@ -26,6 +33,7 @@ import com.intellij.refactoring.extractMethod.ExtractMethodValidator;
import com.intellij.refactoring.listeners.RefactoringElementListenerComposite;
import com.intellij.refactoring.rename.RenameUtil;
import com.intellij.refactoring.util.CommonRefactoringUtil;
import com.intellij.ui.ReplacePromptDialog;
import com.intellij.usageView.UsageInfo;
import com.intellij.util.Function;
import com.intellij.util.IncorrectOperationException;
@@ -56,19 +64,19 @@ public class PyExtractMethodUtil {
private PyExtractMethodUtil() {
}
public static void extractFromStatements(final Project project,
final Editor editor,
final PyCodeFragment fragment,
final PsiElement statement1,
final PsiElement statement2) {
public static void extractFromStatements(@NotNull final Project project,
@NotNull final Editor editor,
@NotNull final PyCodeFragment fragment,
@NotNull final PsiElement statement1,
@NotNull final PsiElement statement2) {
if (!fragment.getOutputVariables().isEmpty() && fragment.isReturnInstructionInside()) {
CommonRefactoringUtil.showErrorHint(project, editor,
"Cannot perform refactoring from expression with local variables modifications and return instructions inside code fragment",
PyBundle.message("refactoring.extract.method.error.cannot.perform.refactoring.with.local"),
RefactoringBundle.message("error.title"), "refactoring.extractMethod");
return;
}
PyFunction function = PsiTreeUtil.getParentOfType(statement1, PyFunction.class);
final PyFunction function = PsiTreeUtil.getParentOfType(statement1, PyFunction.class);
final PyUtil.MethodFlags flags = function == null ? null : PyUtil.MethodFlags.of(function);
final boolean isClassMethod = flags != null && flags.isClassMethod();
final boolean isStaticMethod = flags != null && flags.isStaticMethod();
@@ -90,6 +98,8 @@ public class PyExtractMethodUtil {
final String methodName = data.first;
final AbstractVariableData[] variableData = data.second;
final PyDuplicatesFinder finder = new PyDuplicatesFinder(statement1, statement2);
if (fragment.getOutputVariables().isEmpty()) {
CommandProcessor.getInstance().executeCommand(project, new Runnable() {
public void run() {
@@ -123,6 +133,8 @@ public class PyExtractMethodUtil {
// Replace statements with call
callElement = replaceElements(elementsRange, callElement);
callElement = CodeInsightUtilBase.forcePsiPostprocessAndRestoreElement(callElement);
processDuplicates(callElement, generatedMethod, finder, editor);
// Set editor
setSelectionAndCaret(editor, callElement);
@@ -174,6 +186,8 @@ public class PyExtractMethodUtil {
// replace statements with call
callElement = replaceElements(elementsRange, callElement);
callElement = CodeInsightUtilBase.forcePsiPostprocessAndRestoreElement(callElement);
processDuplicates(callElement, generatedMethod, finder, editor);
// Set editor
setSelectionAndCaret(editor, callElement);
@@ -184,7 +198,66 @@ public class PyExtractMethodUtil {
}
}
private static void processGlobalWrites(@NotNull PyFunction function, @NotNull PyCodeFragment fragment) {
private static void processDuplicates(@NotNull final PsiElement callElement,
@NotNull final PyFunction generatedMethod,
@NotNull final PyDuplicatesFinder finder,
@NotNull final Editor editor) {
final ScopeOwner owner = ScopeUtil.getScopeOwner(callElement);
if (owner instanceof PsiFile) return;
final List<Pair<PsiElement, PsiElement>> duplicates = finder.findDuplicates(owner, generatedMethod);
if (duplicates.size() > 0) {
final String message = RefactoringBundle.message("0.has.detected.1.code.fragments.in.this.file.that.can.be.replaced.with.a.call.to.extracted.method",
ApplicationNamesInfo.getInstance().getProductName(), duplicates.size());
final boolean isUnittest = ApplicationManager.getApplication().isUnitTestMode();
final int exitCode = !isUnittest ? Messages.showYesNoDialog(callElement.getProject(), message,
RefactoringBundle.message("refactoring.extract.method.dialog.title"), Messages.getInformationIcon()) :
DialogWrapper.OK_EXIT_CODE;
if (exitCode == DialogWrapper.OK_EXIT_CODE) {
boolean replaceAll = false;
for (Pair<PsiElement, PsiElement> match : duplicates) {
final List<PsiElement> elementsRange = PyPsiUtils.collectElements(match.getFirst(), match.getSecond());
if (!replaceAll) {
highlightInEditor(callElement.getProject(), match, editor);
int promptResult = FindManager.PromptResult.ALL;
if (!isUnittest) {
ReplacePromptDialog promptDialog = new ReplacePromptDialog(false, RefactoringBundle.message("replace.fragment"), callElement.getProject());
promptDialog.show();
promptResult = promptDialog.getExitCode();
}
if (promptResult == FindManager.PromptResult.SKIP) continue;
if (promptResult == FindManager.PromptResult.CANCEL) break;
if (promptResult == FindManager.PromptResult.OK) {
replaceElements(elementsRange, callElement);
}
else if (promptResult == FindManager.PromptResult.ALL) {
replaceElements(elementsRange, callElement);
replaceAll = true;
}
}
else {
replaceElements(elementsRange, callElement);
}
}
}
}
}
private static void highlightInEditor(@NotNull final Project project, @NotNull final Pair<PsiElement, PsiElement> pair,
@NotNull final Editor editor) {
final HighlightManager highlightManager = HighlightManager.getInstance(project);
final EditorColorsManager colorsManager = EditorColorsManager.getInstance();
final TextAttributes attributes = colorsManager.getGlobalScheme().getAttributes(EditorColors.SEARCH_RESULT_ATTRIBUTES);
final int startOffset = pair.getFirst().getTextRange().getStartOffset();
final int endOffset = pair.getSecond().getTextRange().getEndOffset();
highlightManager.addRangeHighlight(editor, startOffset, endOffset, attributes, true, null);
final LogicalPosition logicalPosition = editor.offsetToLogicalPosition(startOffset);
editor.getScrollingModel().scrollTo(logicalPosition, ScrollType.MAKE_VISIBLE);
}
private static void processGlobalWrites(@NotNull final PyFunction function, @NotNull final PyCodeFragment fragment) {
final Set<String> globalWrites = fragment.getGlobalWrites();
final Set<String> newGlobalNames = new LinkedHashSet<String>();
final Scope scope = ControlFlowCache.getScope(function);
@@ -239,7 +312,7 @@ public class PyExtractMethodUtil {
builder.append(".");
}
public static void extractFromExpression(final Project project,
public static void extractFromExpression(@NotNull final Project project,
final Editor editor,
final PyCodeFragment fragment,
final PsiElement expression) {
@@ -301,7 +374,7 @@ public class PyExtractMethodUtil {
final PyElement generated = generator.createFromText(LanguageLevel.getDefault(), PyElement.class, builder.toString());
PsiElement callElement = null;
if (generated instanceof PyReturnStatement) {
callElement = fragment.isReturnInstructionInside() ? generated : ((PyReturnStatement)generated).getExpression();
callElement = ((PyReturnStatement)generated).getExpression();
}
else if (generated instanceof PyExpressionStatement) {
callElement = ((PyExpressionStatement)generated).getExpression();
@@ -0,0 +1,8 @@
def foo():
a = 1
print a
def bar():
foo()
foo()
@@ -0,0 +1,5 @@
def bar():
<selection>a = 1
print a</selection>
a = 1
print a
@@ -0,0 +1,10 @@
def foo():
a = 1
return a
def bar():
a = foo()
print a
a = foo()
print a
@@ -0,0 +1,5 @@
def bar():
<selection>a = 1</selection>
print a
a = 1
print a
@@ -250,4 +250,12 @@ public class PyExtractMethodTest extends LightMarkedTestCase {
public void testYieldFrom33() {
doTest("bar", LanguageLevel.PYTHON33);
}
public void testDuplicateSingleLine() {
doTest("foo");
}
public void testDuplicateMultiLine() {
doTest("foo");
}
}