From 9d2bd32bba2313167ee4849cd2158bd59ce34a32 Mon Sep 17 00:00:00 2001 From: Andrey Vlasovskikh Date: Sat, 12 Sep 2015 15:54:47 +0300 Subject: [PATCH] Extract async function (PY-16094) --- .../codeFragment/PyCodeFragment.java | 9 +++- .../codeFragment/PyCodeFragmentUtil.java | 4 +- .../python/psi/impl/PyFunctionBuilder.java | 9 ++++ .../extractmethod/PyExtractMethodUtil.java | 50 +++++++++++++++---- .../extractmethod/AsyncDef.after.py | 8 +++ .../extractmethod/AsyncDef.before.py | 3 ++ .../extractmethod/AwaitExpression.after.py | 7 +++ .../extractmethod/AwaitExpression.before.py | 3 ++ .../refactoring/PyExtractMethodTest.java | 8 +++ 9 files changed, 87 insertions(+), 14 deletions(-) create mode 100644 python/testData/refactoring/extractmethod/AsyncDef.after.py create mode 100644 python/testData/refactoring/extractmethod/AsyncDef.before.py create mode 100644 python/testData/refactoring/extractmethod/AwaitExpression.after.py create mode 100644 python/testData/refactoring/extractmethod/AwaitExpression.before.py diff --git a/python/src/com/jetbrains/python/codeInsight/codeFragment/PyCodeFragment.java b/python/src/com/jetbrains/python/codeInsight/codeFragment/PyCodeFragment.java index c39a1f5e42c3..bb5d9eefdb74 100644 --- a/python/src/com/jetbrains/python/codeInsight/codeFragment/PyCodeFragment.java +++ b/python/src/com/jetbrains/python/codeInsight/codeFragment/PyCodeFragment.java @@ -26,17 +26,20 @@ public class PyCodeFragment extends CodeFragment { private final Set myGlobalWrites; private final Set myNonlocalWrites; private final boolean myYieldInside; + private final boolean myAsync; public PyCodeFragment(final Set input, final Set output, final Set globalWrites, final Set 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 getGlobalWrites() { @@ -50,4 +53,8 @@ public class PyCodeFragment extends CodeFragment { public boolean isYieldInside() { return myYieldInside; } + + public boolean isAsync() { + return myAsync; + } } diff --git a/python/src/com/jetbrains/python/codeInsight/codeFragment/PyCodeFragmentUtil.java b/python/src/com/jetbrains/python/codeInsight/codeFragment/PyCodeFragmentUtil.java index ca97e7b4e213..eb1f6a0f05e8 100644 --- a/python/src/com/jetbrains/python/codeInsight/codeFragment/PyCodeFragmentUtil.java +++ b/python/src/com/jetbrains/python/codeInsight/codeFragment/PyCodeFragmentUtil.java @@ -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) { diff --git a/python/src/com/jetbrains/python/psi/impl/PyFunctionBuilder.java b/python/src/com/jetbrains/python/psi/impl/PyFunctionBuilder.java index 4a7abae08a89..8be934854256 100644 --- a/python/src/com/jetbrains/python/psi/impl/PyFunctionBuilder.java +++ b/python/src/com/jetbrains/python/psi/impl/PyFunctionBuilder.java @@ -44,6 +44,7 @@ public class PyFunctionBuilder { private String[] myDocStringLines = null; @NotNull private final Map myDecoratorValues = new HashMap(); + 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, ", ")); diff --git a/python/src/com/jetbrains/python/refactoring/extractmethod/PyExtractMethodUtil.java b/python/src/com/jetbrains/python/refactoring/extractmethod/PyExtractMethodUtil.java index a91efa182050..23b4243aee7d 100644 --- a/python/src/com/jetbrains/python/refactoring/extractmethod/PyExtractMethodUtil.java +++ b/python/src/com/jetbrains/python/refactoring/extractmethod/PyExtractMethodUtil.java @@ -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 newMethodElements = new ArrayList(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 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))); diff --git a/python/testData/refactoring/extractmethod/AsyncDef.after.py b/python/testData/refactoring/extractmethod/AsyncDef.after.py new file mode 100644 index 000000000000..87a5afc9b77f --- /dev/null +++ b/python/testData/refactoring/extractmethod/AsyncDef.after.py @@ -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 diff --git a/python/testData/refactoring/extractmethod/AsyncDef.before.py b/python/testData/refactoring/extractmethod/AsyncDef.before.py new file mode 100644 index 000000000000..e655232e4112 --- /dev/null +++ b/python/testData/refactoring/extractmethod/AsyncDef.before.py @@ -0,0 +1,3 @@ +async def foo(x): + y = await x + return await y diff --git a/python/testData/refactoring/extractmethod/AwaitExpression.after.py b/python/testData/refactoring/extractmethod/AwaitExpression.after.py new file mode 100644 index 000000000000..1fdadf73f1b9 --- /dev/null +++ b/python/testData/refactoring/extractmethod/AwaitExpression.after.py @@ -0,0 +1,7 @@ +async def foo(x): + y = await bar(x) + return y + + +async def bar(x_new): + return await x_new + 1 diff --git a/python/testData/refactoring/extractmethod/AwaitExpression.before.py b/python/testData/refactoring/extractmethod/AwaitExpression.before.py new file mode 100644 index 000000000000..7b8c39b8c44a --- /dev/null +++ b/python/testData/refactoring/extractmethod/AwaitExpression.before.py @@ -0,0 +1,3 @@ +async def foo(x): + y = await x + 1 + return y diff --git a/python/testSrc/com/jetbrains/python/refactoring/PyExtractMethodTest.java b/python/testSrc/com/jetbrains/python/refactoring/PyExtractMethodTest.java index 1ebebf000efe..23e9f78bd958 100644 --- a/python/testSrc/com/jetbrains/python/refactoring/PyExtractMethodTest.java +++ b/python/testSrc/com/jetbrains/python/refactoring/PyExtractMethodTest.java @@ -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); + } }