diff --git a/python/src/com/jetbrains/python/refactoring/makeFunctionTopLevel/PyBaseMakeFunctionTopLevelProcessor.java b/python/src/com/jetbrains/python/refactoring/makeFunctionTopLevel/PyBaseMakeFunctionTopLevelProcessor.java index c34a475619e5..0fb9a680b572 100644 --- a/python/src/com/jetbrains/python/refactoring/makeFunctionTopLevel/PyBaseMakeFunctionTopLevelProcessor.java +++ b/python/src/com/jetbrains/python/refactoring/makeFunctionTopLevel/PyBaseMakeFunctionTopLevelProcessor.java @@ -21,27 +21,32 @@ import com.intellij.openapi.application.ApplicationManager; import com.intellij.openapi.util.text.StringUtil; import com.intellij.psi.PsiElement; import com.intellij.psi.PsiFile; +import com.intellij.psi.PsiNamedElement; import com.intellij.psi.util.PsiTreeUtil; import com.intellij.refactoring.BaseRefactoringProcessor; import com.intellij.refactoring.ui.UsageViewDescriptorAdapter; import com.intellij.usageView.UsageInfo; import com.intellij.usageView.UsageViewDescriptor; import com.intellij.util.ArrayUtil; +import com.intellij.util.containers.HashSet; import com.jetbrains.python.PyNames; import com.jetbrains.python.codeInsight.controlflow.ControlFlowCache; import com.jetbrains.python.codeInsight.controlflow.ReadWriteInstruction; import com.jetbrains.python.codeInsight.controlflow.ScopeOwner; import com.jetbrains.python.codeInsight.dataflow.scope.ScopeUtil; +import com.jetbrains.python.codeInsight.imports.AddImportHelper; import com.jetbrains.python.psi.*; import com.jetbrains.python.psi.impl.PyPsiUtils; import com.jetbrains.python.psi.resolve.PyResolveContext; import com.jetbrains.python.psi.types.TypeEvalContext; import com.jetbrains.python.refactoring.PyRefactoringUtil; +import com.jetbrains.python.refactoring.classes.PyClassRefactoringUtil; import org.jetbrains.annotations.NotNull; import java.util.ArrayList; import java.util.Collection; import java.util.List; +import java.util.Set; import static com.jetbrains.python.psi.PyUtil.as; @@ -50,9 +55,11 @@ import static com.jetbrains.python.psi.PyUtil.as; */ public abstract class PyBaseMakeFunctionTopLevelProcessor extends BaseRefactoringProcessor { protected final PyFunction myFunction; + protected final PsiFile mySourceFile; protected final PyResolveContext myResolveContext; protected final PyElementGenerator myGenerator; protected final String myDestinationPath; + protected final List myExternalReads = new ArrayList(); public PyBaseMakeFunctionTopLevelProcessor(@NotNull PyFunction targetFunction, @NotNull String destinationPath) { super(targetFunction.getProject()); @@ -61,6 +68,7 @@ public abstract class PyBaseMakeFunctionTopLevelProcessor extends BaseRefactorin final TypeEvalContext typeEvalContext = TypeEvalContext.userInitiated(myProject, targetFunction.getContainingFile()); myResolveContext = PyResolveContext.defaultContext().withTypeEvalContext(typeEvalContext); myGenerator = PyElementGenerator.getInstance(myProject); + mySourceFile = myFunction.getContainingFile(); } @NotNull @@ -97,8 +105,36 @@ public abstract class PyBaseMakeFunctionTopLevelProcessor extends BaseRefactorin assert ApplicationManager.getApplication().isWriteAccessAllowed(); + // We should update usages before we generate and insert new function, because we have to update its usages inside + // (e.g. recursive calls) it first updateUsages(newParameters, usages); - replaceFunction(createNewFunction(newParameters)); + final PyFile newFile = PyUtil.getOrCreateFile(myDestinationPath, myProject); + final PyFunction newFunction = insertFunction(createNewFunction(newParameters), newFile); + + myFunction.delete(); + + updateImports(newFunction, usages); + } + + private void updateImports(@NotNull PyFunction newFunction, @NotNull UsageInfo[] usages) { + final Set usageFiles = new HashSet(); + for (UsageInfo usage : usages) { + usageFiles.add(usage.getFile()); + } + for (PsiFile file : usageFiles) { + if (file != newFunction.getContainingFile()) { + PyClassRefactoringUtil.insertImport(file, newFunction, null, true); + } + } + // References inside the body of function + if (newFunction.getContainingFile() != mySourceFile) { + for (PsiElement read : myExternalReads) { + if (read instanceof PsiNamedElement && read.isValid()) { + PyClassRefactoringUtil.insertImport(newFunction, (PsiNamedElement)read, null, true); + } + } + PyClassRefactoringUtil.optimizeImports(mySourceFile); + } } @NotNull @@ -137,12 +173,17 @@ public abstract class PyBaseMakeFunctionTopLevelProcessor extends BaseRefactorin } @NotNull - protected PyFunction replaceFunction(@NotNull PyFunction newFunction) { - final PsiFile file = myFunction.getContainingFile(); - final PsiElement anchor = PyPsiUtils.getParentRightBefore(myFunction, file); - - myFunction.delete(); - return (PyFunction)file.addAfter(newFunction, anchor); + protected PyFunction insertFunction(@NotNull PyFunction newFunction, PyFile newFile) { + final PyFunction replacement; + if (mySourceFile == newFile) { + final PsiElement anchor; + anchor = PyPsiUtils.getParentRightBefore(myFunction, mySourceFile); + replacement = (PyFunction)mySourceFile.addAfter(newFunction, anchor); + } + else { + replacement = (PyFunction)newFile.addAfter(newFunction, AddImportHelper.getFileInsertPosition(newFile)); + } + return replacement; } @NotNull @@ -165,6 +206,9 @@ public abstract class PyBaseMakeFunctionTopLevelProcessor extends BaseRefactorin if (isFromEnclosingScope(resolved)) { result.readsFromEnclosingScope.add(element); } + else if (!PsiTreeUtil.isAncestor(myFunction, resolved, false)) { + myExternalReads.add(resolved); + } if (resolved instanceof PyParameter && ((PyParameter)resolved).isSelf()) { if (PsiTreeUtil.getParentOfType(resolved, PyFunction.class) == myFunction) { result.readsOfSelfParameter.add(element); @@ -200,7 +244,7 @@ public abstract class PyBaseMakeFunctionTopLevelProcessor extends BaseRefactorin } private boolean isFromEnclosingScope(@NotNull PsiElement element) { - return element.getContainingFile() == myFunction.getContainingFile() && + return element.getContainingFile() == mySourceFile && !PsiTreeUtil.isAncestor(myFunction, element, false) && !(ScopeUtil.getScopeOwner(element) instanceof PsiFile); } diff --git a/python/src/com/jetbrains/python/refactoring/makeFunctionTopLevel/PyMakeMethodTopLevelProcessor.java b/python/src/com/jetbrains/python/refactoring/makeFunctionTopLevel/PyMakeMethodTopLevelProcessor.java index f7d86c1b4f46..6a2d65e70dca 100644 --- a/python/src/com/jetbrains/python/refactoring/makeFunctionTopLevel/PyMakeMethodTopLevelProcessor.java +++ b/python/src/com/jetbrains/python/refactoring/makeFunctionTopLevel/PyMakeMethodTopLevelProcessor.java @@ -19,7 +19,6 @@ import com.google.common.collect.Iterables; import com.google.common.collect.Lists; import com.intellij.openapi.util.Comparing; import com.intellij.psi.PsiElement; -import com.intellij.psi.PsiFile; import com.intellij.psi.util.PsiTreeUtil; import com.intellij.usageView.UsageInfo; import com.intellij.util.Function; @@ -30,10 +29,7 @@ import com.intellij.util.containers.MultiMap; import com.jetbrains.python.PyBundle; import com.jetbrains.python.PyNames; import com.jetbrains.python.codeInsight.controlflow.ScopeOwner; -import com.jetbrains.python.codeInsight.imports.AddImportHelper; -import com.jetbrains.python.codeInsight.imports.AddImportHelper.ImportPriority; import com.jetbrains.python.psi.*; -import com.jetbrains.python.psi.resolve.QualifiedNameFinder; import com.jetbrains.python.refactoring.NameSuggesterUtil; import com.jetbrains.python.refactoring.introduce.IntroduceValidator; import org.jetbrains.annotations.NotNull; @@ -138,16 +134,6 @@ public class PyMakeMethodTopLevelProcessor extends PyBaseMakeFunctionTopLevelPro } } - final PsiFile usageFile = usage.getFile(); - final PsiFile origFile = myFunction.getContainingFile(); - if (usageFile != origFile) { - final String funcName = myFunction.getName(); - final String origModuleName = QualifiedNameFinder.findShortestImportableName(origFile, origFile.getVirtualFile()); - if (usageFile != null && origModuleName != null && funcName != null) { - AddImportHelper.addOrUpdateFromImportStatement(usageFile, origModuleName, funcName, null, ImportPriority.PROJECT, null); - } - } - // Will replace/invalidate entire expression removeQualifier((PyReferenceExpression)usageElem); } diff --git a/python/testData/refactoring/makeFunctionTopLevel/methodImportUpdates/after/main.py b/python/testData/refactoring/makeFunctionTopLevel/methodImportUpdates/after/main.py index 203ab74aa8b2..4495d9681e68 100644 --- a/python/testData/refactoring/makeFunctionTopLevel/methodImportUpdates/after/main.py +++ b/python/testData/refactoring/makeFunctionTopLevel/methodImportUpdates/after/main.py @@ -1,3 +1,5 @@ +import sys + class C: def __init__(self): self.foo = 42 @@ -5,3 +7,4 @@ class C: def method(foo, x): print(foo) + print(sys.path) diff --git a/python/testData/refactoring/makeFunctionTopLevel/methodImportUpdates/before/main.py b/python/testData/refactoring/makeFunctionTopLevel/methodImportUpdates/before/main.py index 5335d08fac7f..981927ee2347 100644 --- a/python/testData/refactoring/makeFunctionTopLevel/methodImportUpdates/before/main.py +++ b/python/testData/refactoring/makeFunctionTopLevel/methodImportUpdates/before/main.py @@ -1,6 +1,9 @@ +import sys + class C: def __init__(self): self.foo = 42 def method(self, x): print(self.foo) + print(sys.path) diff --git a/python/testData/refactoring/makeFunctionTopLevel/methodMoveToOtherFile/after/main.py b/python/testData/refactoring/makeFunctionTopLevel/methodMoveToOtherFile/after/main.py new file mode 100644 index 000000000000..b0a94ca32b2a --- /dev/null +++ b/python/testData/refactoring/makeFunctionTopLevel/methodMoveToOtherFile/after/main.py @@ -0,0 +1,6 @@ +class C: + def __init__(self): + self.foo = 42 + + +C().main('spam') \ No newline at end of file diff --git a/python/testData/refactoring/makeFunctionTopLevel/methodMoveToOtherFile/after/util.py b/python/testData/refactoring/makeFunctionTopLevel/methodMoveToOtherFile/after/util.py new file mode 100644 index 000000000000..e902aaa3a9ba --- /dev/null +++ b/python/testData/refactoring/makeFunctionTopLevel/methodMoveToOtherFile/after/util.py @@ -0,0 +1,6 @@ +import sys + + +def method(foo, x): + print(foo) + print(sys.path) \ No newline at end of file diff --git a/python/testData/refactoring/makeFunctionTopLevel/methodMoveToOtherFile/before/main.py b/python/testData/refactoring/makeFunctionTopLevel/methodMoveToOtherFile/before/main.py new file mode 100644 index 000000000000..6aae84c2e3de --- /dev/null +++ b/python/testData/refactoring/makeFunctionTopLevel/methodMoveToOtherFile/before/main.py @@ -0,0 +1,12 @@ +import sys + +class C: + def __init__(self): + self.foo = 42 + + def method(self, x): + print(self.foo) + print(sys.path) + + +C().main('spam') \ No newline at end of file diff --git a/python/testSrc/com/jetbrains/python/refactoring/PyMakeFunctionTopLevelTest.java b/python/testSrc/com/jetbrains/python/refactoring/PyMakeFunctionTopLevelTest.java index d1dec70a4a4b..711cf04a6751 100644 --- a/python/testSrc/com/jetbrains/python/refactoring/PyMakeFunctionTopLevelTest.java +++ b/python/testSrc/com/jetbrains/python/refactoring/PyMakeFunctionTopLevelTest.java @@ -16,6 +16,10 @@ package com.jetbrains.python.refactoring; import com.intellij.openapi.command.WriteCommandAction; +import com.intellij.openapi.roots.ModuleRootManager; +import com.intellij.openapi.util.io.FileUtil; +import com.intellij.openapi.vfs.VirtualFile; +import com.intellij.testFramework.PlatformTestUtil; import com.intellij.util.IncorrectOperationException; import com.jetbrains.python.PyBundle; import com.jetbrains.python.fixtures.PyTestCase; @@ -27,47 +31,21 @@ import com.jetbrains.python.refactoring.makeFunctionTopLevel.PyMakeMethodTopLeve import org.jetbrains.annotations.NotNull; import org.jetbrains.annotations.Nullable; +import java.io.IOException; + /** * @author Mikhail Golubev */ public class PyMakeFunctionTopLevelTest extends PyTestCase { - public void doTest(@Nullable String message) { + public void doTest(@Nullable String errorMessage) { myFixture.configureByFile(getTestName(true) + ".py"); - final PyFunction function = assertInstanceOf(myFixture.getElementAtCaret(), PyFunction.class); - final String destination = PyPsiUtils.getContainingFilePath(function); - assertNotNull(destination); - try { - WriteCommandAction.runWriteCommandAction(myFixture.getProject(), new Runnable() { - @Override - public void run() { - if (function.getContainingClass() != null) { - new PyMakeMethodTopLevelProcessor(function, destination).run(); - } - else { - new PyMakeLocalFunctionTopLevelProcessor(function, destination).run(); - } - } - }); + runRefactoring(null, errorMessage); + if (errorMessage == null) { myFixture.checkResultByFile(getTestName(true) + ".after.py"); } - catch (IncorrectOperationException e) { - if (message == null) { - fail("Refactoring failed unexpectedly with message: " + e.getMessage()); - } - assertEquals(message, e.getMessage()); - } } - //private void doMultiFileTest() throws IOException { - // final String rootBeforePath = getTestName(true) + "/before"; - // final String rootAfterPath = getTestName(true) + "/after"; - // final VirtualFile copiedDirectory = myFixture.copyDirectoryToProject(rootBeforePath, ""); - // myFixture.configureByFile("main.py"); - // myFixture.testAction(new PyMakeFunctionTopLevelRefactoring()); - // PlatformTestUtil.assertDirectoriesEqual(getVirtualFileByName(getTestDataPath() + rootAfterPath), copiedDirectory); - //} - private void doTestSuccess() { doTest(null); } @@ -76,6 +54,49 @@ public class PyMakeFunctionTopLevelTest extends PyTestCase { doTest(message); } + private void runRefactoring(@Nullable String destination, @Nullable String errorMessage) { + final PyFunction function = assertInstanceOf(myFixture.getElementAtCaret(), PyFunction.class); + if (destination == null) { + destination = PyPsiUtils.getContainingFilePath(function); + } + else { + final VirtualFile srcRoot = ModuleRootManager.getInstance(myFixture.getModule()).getSourceRoots()[0]; + destination = FileUtil.join(srcRoot.getPath(), destination); + } + assertNotNull(destination); + final String finalDestination = destination; + try { + WriteCommandAction.runWriteCommandAction(myFixture.getProject(), new Runnable() { + @Override + public void run() { + if (function.getContainingClass() != null) { + new PyMakeMethodTopLevelProcessor(function, finalDestination).run(); + } + else { + new PyMakeLocalFunctionTopLevelProcessor(function, finalDestination).run(); + } + } + }); + } + catch (IncorrectOperationException e) { + if (errorMessage == null) { + fail("Refactoring failed unexpectedly with message: " + e.getMessage()); + } + assertEquals(errorMessage, e.getMessage()); + } + } + + private void doMultiFileTest(@Nullable String destination, @Nullable String errorMessage) throws IOException { + final String rootBeforePath = getTestName(true) + "/before"; + final String rootAfterPath = getTestName(true) + "/after"; + final VirtualFile copiedDirectory = myFixture.copyDirectoryToProject(rootBeforePath, ""); + myFixture.configureByFile("main.py"); + runRefactoring(destination, errorMessage); + if (errorMessage == null) { + PlatformTestUtil.assertDirectoriesEqual(getVirtualFileByName(getTestDataPath() + rootAfterPath), copiedDirectory); + } + } + //private static boolean isActionEnabled() { // final PyMakeFunctionTopLevelRefactoring action = new PyMakeFunctionTopLevelRefactoring(); // final TestActionEvent event = new TestActionEvent(action); @@ -184,9 +205,13 @@ public class PyMakeFunctionTopLevelTest extends PyTestCase { doTestSuccess(); } - //public void testMethodImportUpdates() throws IOException { - // doMultiFileTest(); - //} + public void testMethodImportUpdates() throws IOException { + doMultiFileTest(null, null); + } + + public void testMethodMoveToOtherFile() throws IOException { + doMultiFileTest("util.py", null); + } public void testMethodCalledViaClass() { doTestSuccess();