diff --git a/python/src/com/jetbrains/python/refactoring/extractmethod/PyExtractMethodUtil.java b/python/src/com/jetbrains/python/refactoring/extractmethod/PyExtractMethodUtil.java index 323d8808e10d..f79bf2ceb543 100644 --- a/python/src/com/jetbrains/python/refactoring/extractmethod/PyExtractMethodUtil.java +++ b/python/src/com/jetbrains/python/refactoring/extractmethod/PyExtractMethodUtil.java @@ -83,7 +83,7 @@ public class PyExtractMethodUtil { final String methodName = data.first; final AbstractVariableData[] variableData = data.second; - final SimpleDuplicatesFinder finder = new SimpleDuplicatesFinder(statement1, statement2, variableData); + final SimpleDuplicatesFinder finder = new SimpleDuplicatesFinder(statement1, statement2, variableData, fragment.getOutputVariables()); CommandProcessor.getInstance().executeCommand(project, new Runnable() { public void run() { @@ -236,7 +236,7 @@ public class PyExtractMethodUtil { public static void extractFromExpression(@NotNull final Project project, final Editor editor, final PyCodeFragment fragment, - final PsiElement expression) { + @NotNull final PsiElement expression) { if (!fragment.getOutputVariables().isEmpty()){ CommonRefactoringUtil.showErrorHint(project, editor, "Cannot perform refactoring from expression with local variables modifications inside code fragment", @@ -264,6 +264,7 @@ public class PyExtractMethodUtil { final String methodName = data.first; final AbstractVariableData[] variableData = data.second; + final SimpleDuplicatesFinder finder = new SimpleDuplicatesFinder(expression, expression, variableData, fragment.getOutputVariables()); if (fragment.getOutputVariables().isEmpty()) { CommandProcessor.getInstance().executeCommand(project, new Runnable() { @Override @@ -305,7 +306,8 @@ public class PyExtractMethodUtil { if (callElement != null) { callElement = PyReplaceExpressionUtil.replaceExpression(expression, callElement); } - + if (callElement != null) + processDuplicates(callElement, generatedMethod, finder, editor); // Set editor setSelectionAndCaret(editor, callElement); } @@ -328,27 +330,40 @@ public class PyExtractMethodUtil { return callElement; } - private static PsiElement replaceElements(@NotNull final SimpleMatch match, @NotNull final PsiElement callElement) { + private static PsiElement replaceElements(@NotNull final SimpleMatch match, @NotNull final PsiElement element) { final List elementsRange = PyPsiUtils.collectElements(match.getStartElement(), match.getEndElement()); final Map changedParameters = match.getChangedParameters(); - if (callElement instanceof PyExpressionStatement) { - final PyExpression expression = ((PyExpressionStatement)callElement).getExpression(); - if (expression instanceof PyCallExpression) { - final Set keys = changedParameters.keySet(); - final PyArgumentList argumentList = ((PyCallExpression)expression).getArgumentList(); - if (argumentList != null) { - for (PyExpression arg : argumentList.getArguments()) { - final String argText = arg.getText(); - if (argText != null && keys.contains(argText)) { - arg.replace(PyElementGenerator.getInstance(callElement.getProject()).createExpressionFromText( - LanguageLevel.forElement(callElement), - changedParameters.get(argText))); - } + PsiElement callElement = element; + final PyElementGenerator generator = PyElementGenerator.getInstance(callElement.getProject()); + if (element instanceof PyAssignmentStatement) { + final PyExpression value = ((PyAssignmentStatement)element).getAssignedValue(); + if (value != null) callElement = value; + final PyExpression[] targets = ((PyAssignmentStatement)element).getTargets(); + if (targets.length == 1) { + final String output = match.getChangedOutput(); + final PyExpression text = generator.createFromText(LanguageLevel.forElement(callElement), PyAssignmentStatement.class, + output + " = 1").getTargets()[0]; + targets[0].replace(text); + } + } + if (element instanceof PyExpressionStatement) { + callElement = ((PyExpressionStatement)element).getExpression(); + } + if (callElement instanceof PyCallExpression) { + final Set keys = changedParameters.keySet(); + final PyArgumentList argumentList = ((PyCallExpression)callElement).getArgumentList(); + if (argumentList != null) { + for (PyExpression arg : argumentList.getArguments()) { + final String argText = arg.getText(); + if (argText != null && keys.contains(argText)) { + arg.replace(generator.createExpressionFromText( + LanguageLevel.forElement(callElement), + changedParameters.get(argText))); } } } } - return replaceElements(elementsRange, callElement); + return replaceElements(elementsRange, element); } // Creates string for call diff --git a/python/testData/refactoring/extractmethod/DuplicateCheckParam.after.py b/python/testData/refactoring/extractmethod/DuplicateCheckParam.after.py new file mode 100644 index 000000000000..c4f329969215 --- /dev/null +++ b/python/testData/refactoring/extractmethod/DuplicateCheckParam.after.py @@ -0,0 +1,11 @@ +def foo(a_new): + b1 = do_smth_with(a_new) + return b1 + + +def f(): + a = do_smth() + b1 = foo(a) + a = do_smth() + b = foo(a + 1) + do_smth_with(b1, b) diff --git a/python/testData/refactoring/extractmethod/DuplicateCheckParam.before.py b/python/testData/refactoring/extractmethod/DuplicateCheckParam.before.py new file mode 100644 index 000000000000..7db0066886ff --- /dev/null +++ b/python/testData/refactoring/extractmethod/DuplicateCheckParam.before.py @@ -0,0 +1,6 @@ +def f(): + a = do_smth() + b1 = do_smth_with(a) + a = do_smth() + b = do_smth_with(a + 1) + do_smth_with(b1, b) diff --git a/python/testSrc/com/jetbrains/python/refactoring/PyExtractMethodTest.java b/python/testSrc/com/jetbrains/python/refactoring/PyExtractMethodTest.java index e76c417e0247..8ebe4ed1a215 100644 --- a/python/testSrc/com/jetbrains/python/refactoring/PyExtractMethodTest.java +++ b/python/testSrc/com/jetbrains/python/refactoring/PyExtractMethodTest.java @@ -266,4 +266,8 @@ public class PyExtractMethodTest extends LightMarkedTestCase { public void testDuplicateWithRename() { doTest("foo"); } + + public void testDuplicateCheckParam() { + doTest("foo"); + } }