Extract async function (PY-16094)

This commit is contained in:
Andrey Vlasovskikh
2015-09-12 15:54:47 +03:00
parent 17db5621d6
commit 9d2bd32bba
9 changed files with 87 additions and 14 deletions
@@ -26,17 +26,20 @@ public class PyCodeFragment extends CodeFragment {
private final Set<String> myGlobalWrites;
private final Set<String> myNonlocalWrites;
private final boolean myYieldInside;
private final boolean myAsync;
public PyCodeFragment(final Set<String> input,
final Set<String> output,
final Set<String> globalWrites,
final Set<String> nonlocalWrites,
final boolean returnInside,
final boolean yieldInside) {
final boolean yieldInside,
final boolean isAsync) {
super(input, output, returnInside);
myGlobalWrites = globalWrites;
myNonlocalWrites = nonlocalWrites;
myYieldInside = yieldInside;
myAsync = isAsync;
}
public Set<String> getGlobalWrites() {
@@ -50,4 +53,8 @@ public class PyCodeFragment extends CodeFragment {
public boolean isYieldInside() {
return myYieldInside;
}
public boolean isAsync() {
return myAsync;
}
}
@@ -99,13 +99,13 @@ public class PyCodeFragmentUtil {
}
}
final boolean yieldsFound = subGraphAnalysis.yieldExpressions > 0;
if (yieldsFound && LanguageLevel.forElement(owner).isOlderThan(LanguageLevel.PYTHON33)) {
throw new CannotCreateCodeFragmentException(PyBundle.message("refactoring.extract.method.error.yield"));
}
final boolean isAsync = owner instanceof PyFunction && ((PyFunction)owner).isAsync();
return new PyCodeFragment(inputNames, outputNames, globalWrites, nonlocalWrites, subGraphAnalysis.returns > 0, yieldsFound);
return new PyCodeFragment(inputNames, outputNames, globalWrites, nonlocalWrites, subGraphAnalysis.returns > 0, yieldsFound, isAsync);
}
private static boolean resolvesToBoundMethodParameter(@NotNull PsiElement element) {
@@ -44,6 +44,7 @@ public class PyFunctionBuilder {
private String[] myDocStringLines = null;
@NotNull
private final Map<String, String> myDecoratorValues = new HashMap<String, String>();
private boolean myAsync = false;
/**
* Creates builder copying signature and doc from another one.
@@ -136,6 +137,11 @@ public class PyFunctionBuilder {
return this;
}
public PyFunctionBuilder makeAsync() {
myAsync = true;
return this;
}
public PyFunctionBuilder statement(String text) {
myStatements.add(text);
return this;
@@ -166,6 +172,9 @@ public class PyFunctionBuilder {
}
decoratorAppender.append("\n");
}
if (myAsync) {
builder.append("async ");
}
builder.append("def ");
builder.append(myName).append("(");
builder.append(StringUtil.join(myParameters, ", "));
@@ -117,6 +117,11 @@ public class PyExtractMethodUtil {
.refactoringStarted(getRefactoringId(), beforeData);
final StringBuilder builder = new StringBuilder();
final boolean isAsync = fragment.isAsync();
if (isAsync) {
builder.append("async ");
}
builder.append("def f():\n ");
final List<PsiElement> newMethodElements = new ArrayList<PsiElement>(elementsRange);
final boolean hasOutputVariables = !fragment.getOutputVariables().isEmpty();
@@ -124,14 +129,17 @@ public class PyExtractMethodUtil {
final LanguageLevel languageLevel = LanguageLevel.forElement(statement1);
if (hasOutputVariables) {
// Generate return modified variables statements
StringUtil.join(fragment.getOutputVariables(), ", ", builder);
final String outputVariables = StringUtil.join(fragment.getOutputVariables(), ", ");
String newMethodText = builder + "return " + outputVariables;
builder.append(outputVariables);
final PsiElement returnStatement = generator.createFromText(languageLevel, PyElement.class, "return " + builder.toString());
final PyFunction function = generator.createFromText(languageLevel, PyFunction.class, newMethodText);
final PsiElement returnStatement = function.getStatementList().getStatements()[0];
newMethodElements.add(returnStatement);
}
// Generate method
PyFunction generatedMethod = generateMethodFromElements(project, methodName, variableData, newMethodElements, flags);
PyFunction generatedMethod = generateMethodFromElements(project, methodName, variableData, newMethodElements, flags, isAsync);
generatedMethod = insertGeneratedMethod(statement1, generatedMethod);
// Process parameters
@@ -148,7 +156,10 @@ public class PyExtractMethodUtil {
else if (fragment.isReturnInstructionInside()) {
builder.append("return ");
}
if (fragment.isYieldInside()) {
if (isAsync) {
builder.append("await ");
}
else if (fragment.isYieldInside()) {
builder.append("yield from ");
}
if (isMethod) {
@@ -156,7 +167,8 @@ public class PyExtractMethodUtil {
}
builder.append(methodName).append("(");
builder.append(createCallArgsString(variableData)).append(")");
PsiElement callElement = generator.createFromText(languageLevel, PyElement.class, builder.toString());
final PyFunction function = generator.createFromText(languageLevel, PyFunction.class, builder.toString());
PsiElement callElement = function.getStatementList().getStatements()[0];
// replace statements with call
callElement = replaceElements(elementsRange, callElement);
@@ -297,7 +309,8 @@ public class PyExtractMethodUtil {
@Override
public void run() {
// Generate method
PyFunction generatedMethod = generateMethodFromExpression(project, methodName, variableData, expression, flags);
final boolean isAsync = fragment.isAsync();
PyFunction generatedMethod = generateMethodFromExpression(project, methodName, variableData, expression, flags, isAsync);
generatedMethod = insertGeneratedMethod(expression, generatedMethod);
// Process parameters
@@ -306,7 +319,14 @@ public class PyExtractMethodUtil {
// Generating call element
final StringBuilder builder = new StringBuilder();
if (fragment.isYieldInside()) {
if (isAsync) {
builder.append("async ");
}
builder.append("def f():\n ");
if (isAsync) {
builder.append("await ");
}
else if (fragment.isYieldInside()) {
builder.append("yield from ");
}
else {
@@ -318,8 +338,9 @@ public class PyExtractMethodUtil {
builder.append(methodName);
builder.append("(").append(createCallArgsString(variableData)).append(")");
final PyElementGenerator generator = PyElementGenerator.getInstance(project);
final PyElement generated =
generator.createFromText(LanguageLevel.forElement(expression), PyElement.class, builder.toString());
final PyFunction function = generator.createFromText(LanguageLevel.forElement(expression), PyFunction.class,
builder.toString());
final PyElement generated = function.getStatementList().getStatements()[0];
PsiElement callElement = null;
if (generated instanceof PyReturnStatement) {
callElement = ((PyReturnStatement)generated).getExpression();
@@ -495,10 +516,13 @@ public class PyExtractMethodUtil {
@NotNull final String methodName,
@NotNull final AbstractVariableData[] variableData,
@NotNull final PsiElement expression,
@Nullable final PyUtil.MethodFlags flags) {
@Nullable final PyUtil.MethodFlags flags, boolean isAsync) {
final PyFunctionBuilder builder = new PyFunctionBuilder(methodName);
addDecorators(builder, flags);
addFakeParameters(builder, variableData);
if (isAsync) {
builder.makeAsync();
}
final String text;
if (expression instanceof PyYieldExpression) {
text = String.format("(%s)", expression.getText());
@@ -515,10 +539,14 @@ public class PyExtractMethodUtil {
@NotNull final String methodName,
@NotNull final AbstractVariableData[] variableData,
@NotNull final List<PsiElement> elementsRange,
@Nullable PyUtil.MethodFlags flags) {
@Nullable PyUtil.MethodFlags flags,
boolean isAsync) {
assert !elementsRange.isEmpty() : "Empty statements list was selected!";
final PyFunctionBuilder builder = new PyFunctionBuilder(methodName);
if (isAsync) {
builder.makeAsync();
}
addDecorators(builder, flags);
addFakeParameters(builder, variableData);
final PyFunction method = builder.buildFunction(project, LanguageLevel.forElement(elementsRange.get(0)));
@@ -0,0 +1,8 @@
async def foo(x):
y = await bar(x)
return await y
async def bar(x_new):
y = await x_new
return y
@@ -0,0 +1,3 @@
async def foo(x):
<selection>y = await x</selection>
return await y
@@ -0,0 +1,7 @@
async def foo(x):
y = await bar(x)
return y
async def bar(x_new):
return await x_new + 1
@@ -0,0 +1,3 @@
async def foo(x):
y = <selection>await x + 1</selection>
return y
@@ -278,4 +278,12 @@ public class PyExtractMethodTest extends LightMarkedTestCase {
public void testProhibitedAtClassLevel() {
doFail("foo", "Cannot perform refactoring at class level");
}
public void testAsyncDef() {
doTest("bar", LanguageLevel.PYTHON35);
}
public void testAwaitExpression() {
doTest("bar", LanguageLevel.PYTHON35);
}
}