PY-17265 Correctly update usages inside function escalated to other file

This commit is contained in:
Mikhail Golubev
2016-10-25 00:03:49 +03:00
parent f3653451e0
commit 9f39149814
8 changed files with 141 additions and 56 deletions
@@ -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<PsiElement> myExternalReads = new ArrayList<PsiElement>();
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<PsiFile> usageFiles = new HashSet<PsiFile>();
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);
}
@@ -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);
}
@@ -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)
@@ -1,6 +1,9 @@
import sys
class C:
def __init__(self):
self.foo = 42
def me<caret>thod(self, x):
print(self.foo)
print(sys.path)
@@ -0,0 +1,6 @@
class C:
def __init__(self):
self.foo = 42
C().main('spam')
@@ -0,0 +1,6 @@
import sys
def method(foo, x):
print(foo)
print(sys.path)
@@ -0,0 +1,12 @@
import sys
class C:
def __init__(self):
self.foo = 42
def me<caret>thod(self, x):
print(self.foo)
print(sys.path)
C().main('spam')
@@ -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();