Fixed extract method for fragments with nonlocal variables (PY-6625)

This commit is contained in:
Andrey Vlasovskikh
2012-05-23 20:22:38 +04:00
parent f4f66b0f32
commit bd42933bd2
10 changed files with 126 additions and 32 deletions
@@ -8,14 +8,24 @@ import java.util.Set;
* @author vlan
*/
public class PyCodeFragment extends CodeFragment {
private final Set<String> myGlobals;
private final Set<String> myGlobalWrites;
private final Set<String> myNonlocalWrites;
public PyCodeFragment(final Set<String> input, final Set<String> output, final Set<String> globals, final boolean returnInside) {
public PyCodeFragment(final Set<String> input,
final Set<String> output,
final Set<String> globalWrites,
final Set<String> nonlocalWrites,
final boolean returnInside) {
super(input, output, returnInside);
myGlobals = globals;
myGlobalWrites = globalWrites;
myNonlocalWrites = nonlocalWrites;
}
public Set<String> getGlobals() {
return myGlobals;
public Set<String> getGlobalWrites() {
return myGlobalWrites;
}
public Set<String> getNonlocalWrites() {
return myNonlocalWrites;
}
}
@@ -54,6 +54,7 @@ public class PyCodeFragmentUtil {
}
final Set<String> globalWrites = getGlobalWrites(subGraph, owner);
final Set<String> nonlocalWrites = getNonlocalWrites(subGraph, owner);
final Set<String> inputNames = new HashSet<String>();
for (PsiElement element : filterElementsInScope(getInputElements(subGraph, graph), owner)) {
@@ -63,7 +64,7 @@ public class PyCodeFragmentUtil {
if (resolvesToBoundMethodParameter(element)) {
continue;
}
if (globalWrites.contains(name)) {
if (globalWrites.contains(name) || nonlocalWrites.contains(name)) {
continue;
}
inputNames.add(name);
@@ -74,12 +75,14 @@ public class PyCodeFragmentUtil {
for (PsiElement element : getOutputElements(subGraph, graph)) {
final String name = getName(element);
if (name != null) {
if (globalWrites.contains(name) || nonlocalWrites.contains(name)) {
continue;
}
outputNames.add(name);
globalWrites.remove(name);
}
}
return new PyCodeFragment(inputNames, outputNames, globalWrites, subGraphAnalysis.returns > 0);
return new PyCodeFragment(inputNames, outputNames, globalWrites, nonlocalWrites, subGraphAnalysis.returns > 0);
}
private static boolean resolvesToBoundMethodParameter(@NotNull PsiElement element) {
@@ -333,6 +336,21 @@ public class PyCodeFragmentUtil {
return globalWrites;
}
@NotNull
private static Set<String> getNonlocalWrites(@NotNull List<Instruction> instructions, @NotNull ScopeOwner owner) {
final Scope scope = ControlFlowCache.getScope(owner);
final Set<String> nonlocalWrites = new LinkedHashSet<String>();
for (Instruction instruction : getWriteInstructions(instructions)) {
if (instruction instanceof ReadWriteInstruction) {
final String name = ((ReadWriteInstruction)instruction).getName();
if (scope.isNonlocal(name)) {
nonlocalWrites.add(name);
}
}
}
return nonlocalWrites;
}
@NotNull
private static List<PsiElement> multiResolve(@NotNull PsiReference reference) {
if (reference instanceof PsiPolyVariantReference) {
@@ -44,10 +44,7 @@ import com.jetbrains.python.psi.impl.PyPsiUtils;
import org.jetbrains.annotations.NotNull;
import org.jetbrains.annotations.Nullable;
import java.util.ArrayList;
import java.util.List;
import java.util.Map;
import java.util.Set;
import java.util.*;
/**
* @author oleg
@@ -105,8 +102,8 @@ public class PyExtractMethodUtil {
final PsiElement firstElement = elementsRange.get(0);
final boolean isMethod = PyPsiUtils.isMethodContext(firstElement);
processParameters(project, generatedMethod, variableData, isMethod, isClassMethod, isStaticMethod);
processGlobalWrites(generatedMethod, fragment);
processNonlocalWrites(generatedMethod, fragment);
// Generating call element
final StringBuilder builder = new StringBuilder();
@@ -156,8 +153,8 @@ public class PyExtractMethodUtil {
// Process parameters
final boolean isMethod = PyPsiUtils.isMethodContext(elementsRange.get(0));
processParameters(project, generatedMethod, variableData, isMethod, isClassMethod, isStaticMethod);
processGlobalWrites(generatedMethod, fragment);
processNonlocalWrites(generatedMethod, fragment);
// Generate call element
builder.append(" = ");
@@ -181,12 +178,19 @@ public class PyExtractMethodUtil {
}
private static void processGlobalWrites(@NotNull PyFunction function, @NotNull PyCodeFragment fragment) {
final Set<String> globals = fragment.getGlobals();
if (!globals.isEmpty()) {
final Set<String> globalWrites = fragment.getGlobalWrites();
final Set<String> newGlobalNames = new LinkedHashSet<String>();
final Scope scope = ControlFlowCache.getScope(function);
for (String name : globalWrites) {
if (!scope.isGlobal(name)) {
newGlobalNames.add(name);
}
}
if (!newGlobalNames.isEmpty()) {
final PyElementGenerator generator = PyElementGenerator.getInstance(function.getProject());
final PyGlobalStatement globalStatement = generator.createFromText(LanguageLevel.forElement(function),
PyGlobalStatement.class,
"global " + StringUtil.join(globals, ", "));
"global " + StringUtil.join(newGlobalNames, ", "));
final PyStatementList statementList = function.getStatementList();
if (statementList != null) {
statementList.addBefore(globalStatement, statementList.getFirstChild());
@@ -194,6 +198,28 @@ public class PyExtractMethodUtil {
}
}
private static void processNonlocalWrites(@NotNull PyFunction function, @NotNull PyCodeFragment fragment) {
final Set<String> nonlocalWrites = fragment.getNonlocalWrites();
final Set<String> newNonlocalNames = new LinkedHashSet<String>();
final Scope scope = ControlFlowCache.getScope(function);
for (String name : nonlocalWrites) {
if (!scope.isNonlocal(name)) {
newNonlocalNames.add(name);
}
}
if (!newNonlocalNames.isEmpty()) {
final PyElementGenerator generator = PyElementGenerator.getInstance(function.getProject());
final PyNonlocalStatement nonlocalStatement = generator.createFromText(LanguageLevel.forElement(function),
PyNonlocalStatement.class,
"nonlocal " + StringUtil.join(newNonlocalNames, ", "));
final PyStatementList statementList = function.getStatementList();
if (statementList != null) {
statementList.addBefore(nonlocalStatement, statementList.getFirstChild());
}
}
}
private static void appendSelf(PsiElement firstElement, StringBuilder builder, boolean staticMethod) {
if (staticMethod) {
final PyClass containingClass = PsiTreeUtil.getParentOfType(firstElement, PyClass.class);
@@ -1,8 +1,9 @@
print("start")
<begin>
import foo
<end>
foo.bar
def foo():
print("start")
<begin>
import foo
<end>
foo.bar
<result>
In:
Out:
@@ -1,8 +1,9 @@
print("start")
<begin>
aaa = 123
<end>
print(aaa)
def foo():
print("start")
<begin>
aaa = 123
<end>
print(aaa)
<result>
In:
Out:
@@ -1,10 +1,10 @@
a = 2
b = 3
def bar():
global a
a = a + 1
global a, c
a = a + b
c = a * 2
return c
c = bar()
bar()
print(c)
@@ -1,4 +1,5 @@
a = 2
<selection>a = a + 1
b = 3
<selection>a = a + b
c = a * 2</selection>
print(c)
@@ -0,0 +1,13 @@
def foo():
x = 1
def baz():
nonlocal x
x = 2
def bar():
nonlocal x
baz()
print(x)
bar()
foo()
@@ -0,0 +1,8 @@
def foo():
x = 1
def bar():
nonlocal x
<selection>x = 2</selection>
print(x)
bar()
foo()
@@ -7,12 +7,23 @@ import com.intellij.openapi.editor.ex.EditorEx;
import com.intellij.refactoring.RefactoringActionHandler;
import com.jetbrains.python.PythonLanguage;
import com.jetbrains.python.fixtures.LightMarkedTestCase;
import com.jetbrains.python.psi.LanguageLevel;
import com.jetbrains.python.refactoring.extractmethod.PyExtractMethodUtil;
/**
* @author oleg
*/
public class PyExtractMethodTest extends LightMarkedTestCase {
private void doTest(String newName, LanguageLevel level) {
setLanguageLevel(level);
try {
doTest(newName);
}
finally {
setLanguageLevel(null);
}
}
private void doTest(String newName) {
final String testName = getTestName(false);
final String beforeName = testName + ".before.py";
@@ -219,4 +230,9 @@ public class PyExtractMethodTest extends LightMarkedTestCase {
public void testClassWithoutInit() {
doTest("bar");
}
// PY-6625
public void testNonlocal() {
doTest("baz", LanguageLevel.PYTHON30);
}
}