diff --git a/python/python-psi-impl/src/com/jetbrains/python/codeInsight/codeFragment/PyCodeFragmentUtil.java b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/codeFragment/PyCodeFragmentUtil.java index 5e5fcde4b0e5..28d1cb1dbd18 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/codeInsight/codeFragment/PyCodeFragmentUtil.java +++ b/python/python-psi-impl/src/com/jetbrains/python/codeInsight/codeFragment/PyCodeFragmentUtil.java @@ -42,7 +42,6 @@ public final class PyCodeFragmentUtil { final List subGraph = getFragmentSubGraph(graph, start, end); final AnalysisResult subGraphAnalysis = analyseSubGraph(subGraph, start, end); if ((subGraphAnalysis.regularExits > 0 && subGraphAnalysis.returns > 0) || - subGraphAnalysis.targetInstructions > 1 || subGraphAnalysis.outerLoopBreaks > 0) { throw new CannotCreateCodeFragmentException(PyPsiBundle.message("refactoring.extract.method.error.interrupted.execution.flow")); } @@ -162,15 +161,13 @@ public final class PyCodeFragmentUtil { private static class AnalysisResult { private final int starImports; - private final int targetInstructions; private final int regularExits; private final int returns; private final int outerLoopBreaks; private final int yieldExpressions; - AnalysisResult(int starImports, int targetInstructions, int returns, int regularExits, int outerLoopBreaks, int yieldExpressions) { + AnalysisResult(int starImports, int returns, int regularExits, int outerLoopBreaks, int yieldExpressions) { this.starImports = starImports; - this.targetInstructions = targetInstructions; this.regularExits = regularExits; this.returns = returns; this.outerLoopBreaks = outerLoopBreaks; @@ -182,7 +179,6 @@ public final class PyCodeFragmentUtil { private static AnalysisResult analyseSubGraph(@NotNull List subGraph, int start, int end) { int returnSources = 0; int regularSources = 0; - final Set targetInstructions = new HashSet<>(); int starImports = 0; int outerLoopBreaks = 0; int yieldExpressions = 0; @@ -195,18 +191,6 @@ public final class PyCodeFragmentUtil { final PyReturnStatement returnStatement = PsiTreeUtil.getParentOfType(source, PyReturnStatement.class, false); final boolean isExceptTarget = target instanceof PyExceptPart || target instanceof PyFinallyPart; - final boolean isLoopTarget = target instanceof PyWhileStatement || PyForStatementNavigator.getPyForStatementByIterable(target) != null; - - final PyBinaryExpression binaryExpression = PsiTreeUtil.getParentOfType(source, PyBinaryExpression.class); - final boolean isOppositeBinaryTarget = - binaryExpression != null && - ArrayUtil.contains(binaryExpression.getOperator(), PyTokenTypes.AND_KEYWORD, PyTokenTypes.OR_KEYWORD) && - binaryExpression.getLeftExpression() == source && - binaryExpression.getRightExpression() == target; - - if (target != null && !isExceptTarget && !isLoopTarget && !isOppositeBinaryTarget) { - targetInstructions.add(targetInstruction); - } if (returnStatement != null && CodeFragmentUtil.getPosition(returnStatement, start, end) == Position.INSIDE) { returnSources++; @@ -234,7 +218,7 @@ public final class PyCodeFragmentUtil { } } - return new AnalysisResult(starImports, targetInstructions.size(), returnSources, regularSources, outerLoopBreaks, yieldExpressions); + return new AnalysisResult(starImports, returnSources, regularSources, outerLoopBreaks, yieldExpressions); } @NotNull diff --git a/python/python-psi-impl/src/com/jetbrains/python/refactoring/extractmethod/PyExtractMethodUtil.java b/python/python-psi-impl/src/com/jetbrains/python/refactoring/extractmethod/PyExtractMethodUtil.java index e6a28f3675da..20bfe4244ef7 100644 --- a/python/python-psi-impl/src/com/jetbrains/python/refactoring/extractmethod/PyExtractMethodUtil.java +++ b/python/python-psi-impl/src/com/jetbrains/python/refactoring/extractmethod/PyExtractMethodUtil.java @@ -44,6 +44,7 @@ import com.jetbrains.python.psi.impl.PyFunctionBuilder; import com.jetbrains.python.psi.impl.PyPsiUtils; import com.jetbrains.python.refactoring.PyRefactoringUiService; import com.jetbrains.python.refactoring.PyReplaceExpressionUtil; +import one.util.streamex.StreamEx; import org.jetbrains.annotations.NotNull; import org.jetbrains.annotations.Nullable; @@ -189,13 +190,13 @@ public final class PyExtractMethodUtil { final ScopeOwner owner = ScopeUtil.getScopeOwner(anchor); if (owner instanceof PsiFile) return Collections.emptyList(); final List scope = new ArrayList<>(); - if (owner instanceof PyFunction) { - scope.add(owner); + if (owner instanceof PyFunction pyFunction) { + scope.add(pyFunction.getStatementList()); final PyClass containingClass = ((PyFunction)owner).getContainingClass(); if (containingClass != null) { for (PyFunction function : containingClass.getMethods()) { if (!function.equals(owner) && !function.equals(generatedMethod)) { - scope.add(function); + scope.add(function.getStatementList()); } } } @@ -339,7 +340,12 @@ public final class PyExtractMethodUtil { } PyPsiUtils.assertValid(expression); - final List duplicates = collectDuplicates(finder, expression, insertedMethod); + List duplicates = collectDuplicates(finder, expression, insertedMethod); + // When a single reference is extracted into an identity function, prevent unrelated expressions being replaced with calls to it + if (expression instanceof PyReferenceExpression) { + duplicates = ContainerUtil.filter(duplicates, it -> it.getStartElement() == it.getEndElement() && + expression.getText().equals(it.getStartElement().getText())); + } // replace statements with call PsiElement insertedCallElement = null; diff --git a/python/testData/refactoring/extractmethod/ConditionOfConditionalExpression.after.py b/python/testData/refactoring/extractmethod/ConditionOfConditionalExpression.after.py new file mode 100644 index 000000000000..184d8a845d89 --- /dev/null +++ b/python/testData/refactoring/extractmethod/ConditionOfConditionalExpression.after.py @@ -0,0 +1,6 @@ +def f(n): + return n * 2 if bar(n) else n + 1 + + +def bar(n_new): + return n_new \ No newline at end of file diff --git a/python/testData/refactoring/extractmethod/ConditionOfConditionalExpression.before.py b/python/testData/refactoring/extractmethod/ConditionOfConditionalExpression.before.py new file mode 100644 index 000000000000..5df14f415ca9 --- /dev/null +++ b/python/testData/refactoring/extractmethod/ConditionOfConditionalExpression.before.py @@ -0,0 +1,2 @@ +def f(n): + return n * 2 if n else n + 1 \ No newline at end of file diff --git a/python/testData/refactoring/extractmethod/IfConditionExpression.after.py b/python/testData/refactoring/extractmethod/IfConditionExpression.after.py new file mode 100644 index 000000000000..95fc301aa08b --- /dev/null +++ b/python/testData/refactoring/extractmethod/IfConditionExpression.after.py @@ -0,0 +1,7 @@ +def foo(some_var): + if bar(some_var): + print('w00t') + + +def bar(some_var_new): + return some_var_new \ No newline at end of file diff --git a/python/testData/refactoring/extractmethod/IfConditionExpression.before.py b/python/testData/refactoring/extractmethod/IfConditionExpression.before.py new file mode 100644 index 000000000000..5c7c4384a2da --- /dev/null +++ b/python/testData/refactoring/extractmethod/IfConditionExpression.before.py @@ -0,0 +1,3 @@ +def foo(some_var): + if some_var: + print('w00t') \ No newline at end of file diff --git a/python/testData/refactoring/extractmethod/IfElseConditionExpression.after.py b/python/testData/refactoring/extractmethod/IfElseConditionExpression.after.py new file mode 100644 index 000000000000..c6fbdbb39121 --- /dev/null +++ b/python/testData/refactoring/extractmethod/IfElseConditionExpression.after.py @@ -0,0 +1,9 @@ +def foo(some_var): + if bar(some_var): + print('w00t') + else: + print(42) + + +def bar(some_var_new): + return some_var_new \ No newline at end of file diff --git a/python/testData/refactoring/extractmethod/IfElseConditionExpression.before.py b/python/testData/refactoring/extractmethod/IfElseConditionExpression.before.py new file mode 100644 index 000000000000..7ebe14872953 --- /dev/null +++ b/python/testData/refactoring/extractmethod/IfElseConditionExpression.before.py @@ -0,0 +1,5 @@ +def foo(some_var): + if some_var: + print('w00t') + else: + print(42) \ No newline at end of file diff --git a/python/testData/refactoring/extractmethod/SimilarBinaryExpressions.after.py b/python/testData/refactoring/extractmethod/SimilarBinaryExpressions.after.py new file mode 100644 index 000000000000..ff172cbffc0e --- /dev/null +++ b/python/testData/refactoring/extractmethod/SimilarBinaryExpressions.after.py @@ -0,0 +1,7 @@ +def compound_duplicate(p1, p2): + print(bar(p1)) + print(bar(p2)) + + +def bar(p1_new): + return p1_new + 1 \ No newline at end of file diff --git a/python/testData/refactoring/extractmethod/SimilarBinaryExpressions.before.py b/python/testData/refactoring/extractmethod/SimilarBinaryExpressions.before.py new file mode 100644 index 000000000000..fc88d894d3d5 --- /dev/null +++ b/python/testData/refactoring/extractmethod/SimilarBinaryExpressions.before.py @@ -0,0 +1,3 @@ +def compound_duplicate(p1, p2): + print(p1 + 1) + print(p2 + 1) \ No newline at end of file diff --git a/python/testSrc/com/jetbrains/python/refactoring/PyExtractMethodTest.java b/python/testSrc/com/jetbrains/python/refactoring/PyExtractMethodTest.java index a6f2bf775147..c33c38908678 100644 --- a/python/testSrc/com/jetbrains/python/refactoring/PyExtractMethodTest.java +++ b/python/testSrc/com/jetbrains/python/refactoring/PyExtractMethodTest.java @@ -255,6 +255,25 @@ public class PyExtractMethodTest extends LightMarkedTestCase { doFail("foo", "Cannot perform refactoring at a class level"); } + // PY-9045 + public void testIfConditionExpression() { + doTest("bar"); + } + + // PY-9045 + public void testIfElseConditionExpression() { + doTest("bar"); + } + + // PY-9045 + public void testConditionOfConditionalExpression() { + doTest("bar"); + } + + public void testSimilarBinaryExpressions() { + doTest("bar"); + } + public void testAsyncDef() { doTest("bar"); }