Simplify coroutine generation in JS backend

Stop making aliasing suspend function descriptor with reference to
instance of state machine. This may cause problems in some cases,
for example, when compiling recursive suspend function. See KT-17281.
Instead, make alias for synthetic continuation parameter. This
additionally required some refactoring, e.g. *always* generating
continuation parameter during codegen.
This commit is contained in:
Alexey Andreev
2017-04-11 16:01:27 +03:00
parent f4a4a41525
commit 43c084fde3
13 changed files with 95 additions and 80 deletions
@@ -0,0 +1,42 @@
// IGNORE_BACKEND: NATIVE
// WITH_RUNTIME
// WITH_COROUTINES
import kotlin.coroutines.experimental.*
import kotlin.coroutines.experimental.intrinsics.*
fun box(): String {
var result = 0
builder {
result = factorial(4)
}
while (postponed != null) {
postponed!!()
}
if (result != 24) return "fail1: $result"
if (log != "1;1;2;6;24;") return "fail2: $log"
return "OK"
}
suspend fun factorial(a: Int): Int = if (a > 0) suspendHere(factorial(a - 1) * a) else suspendHere(1)
suspend fun suspendHere(value: Int): Int = suspendCoroutineOrReturn { x ->
postponed = {
log += "$value;"
x.resume(value)
}
COROUTINE_SUSPENDED
}
fun builder(c: suspend () -> Unit) {
c.startCoroutine(handleResultContinuation {
postponed = null
})
}
var postponed: (() -> Unit)? = { }
var log = ""
@@ -5114,6 +5114,12 @@ public class IrBlackBoxCodegenTestGenerated extends AbstractIrBlackBoxCodegenTes
doTest(fileName); doTest(fileName);
} }
@TestMetadata("recursiveSuspend.kt")
public void testRecursiveSuspend() throws Exception {
String fileName = KotlinTestUtils.navigationMetadata("compiler/testData/codegen/box/coroutines/recursiveSuspend.kt");
doTest(fileName);
}
@TestMetadata("returnByLabel.kt") @TestMetadata("returnByLabel.kt")
public void testReturnByLabel() throws Exception { public void testReturnByLabel() throws Exception {
String fileName = KotlinTestUtils.navigationMetadata("compiler/testData/codegen/box/coroutines/returnByLabel.kt"); String fileName = KotlinTestUtils.navigationMetadata("compiler/testData/codegen/box/coroutines/returnByLabel.kt");
@@ -5114,6 +5114,12 @@ public class BlackBoxCodegenTestGenerated extends AbstractBlackBoxCodegenTest {
doTest(fileName); doTest(fileName);
} }
@TestMetadata("recursiveSuspend.kt")
public void testRecursiveSuspend() throws Exception {
String fileName = KotlinTestUtils.navigationMetadata("compiler/testData/codegen/box/coroutines/recursiveSuspend.kt");
doTest(fileName);
}
@TestMetadata("returnByLabel.kt") @TestMetadata("returnByLabel.kt")
public void testReturnByLabel() throws Exception { public void testReturnByLabel() throws Exception {
String fileName = KotlinTestUtils.navigationMetadata("compiler/testData/codegen/box/coroutines/returnByLabel.kt"); String fileName = KotlinTestUtils.navigationMetadata("compiler/testData/codegen/box/coroutines/returnByLabel.kt");
@@ -5114,6 +5114,12 @@ public class LightAnalysisModeTestGenerated extends AbstractLightAnalysisModeTes
doTest(fileName); doTest(fileName);
} }
@TestMetadata("recursiveSuspend.kt")
public void testRecursiveSuspend() throws Exception {
String fileName = KotlinTestUtils.navigationMetadata("compiler/testData/codegen/box/coroutines/recursiveSuspend.kt");
doTest(fileName);
}
@TestMetadata("returnByLabel.kt") @TestMetadata("returnByLabel.kt")
public void testReturnByLabel() throws Exception { public void testReturnByLabel() throws Exception {
String fileName = KotlinTestUtils.navigationMetadata("compiler/testData/codegen/box/coroutines/returnByLabel.kt"); String fileName = KotlinTestUtils.navigationMetadata("compiler/testData/codegen/box/coroutines/returnByLabel.kt");
@@ -28,7 +28,8 @@ class CoroutineFunctionTransformer(private val program: JsProgram, private val f
private val innerFunction = function.getInnerFunction() private val innerFunction = function.getInnerFunction()
private val functionWithBody = innerFunction ?: function private val functionWithBody = innerFunction ?: function
private val body = functionWithBody.body private val body = functionWithBody.body
private val localVariables = (function.collectLocalVariables() + functionWithBody.collectLocalVariables()).toMutableSet() private val localVariables = (function.collectLocalVariables() + functionWithBody.collectLocalVariables() -
functionWithBody.parameters.last().name).toMutableSet()
private val className = function.scope.parent.declareFreshName("Coroutine\$${name ?: "anonymous"}") private val className = function.scope.parent.declareFreshName("Coroutine\$${name ?: "anonymous"}")
fun transform(): List<JsStatement> { fun transform(): List<JsStatement> {
@@ -61,29 +62,20 @@ class CoroutineFunctionTransformer(private val program: JsProgram, private val f
if (context.metadata.hasReceiver) { if (context.metadata.hasReceiver) {
constructor.parameters += JsParameter(context.receiverFieldName) constructor.parameters += JsParameter(context.receiverFieldName)
} }
constructor.parameters += function.parameters.map { JsParameter(it.name) } val parameters = function.parameters + innerFunction?.parameters.orEmpty()
if (innerFunction != null) { constructor.parameters += parameters.map { JsParameter(it.name) }
constructor.parameters += innerFunction.parameters.map { JsParameter(it.name) } val lastParameter = parameters.lastOrNull()?.name
}
val lastParameter = function.parameters.lastOrNull()?.name
val controllerName = if (context.metadata.hasController) { val controllerName = if (context.metadata.hasController) {
function.scope.declareFreshName("controller").apply { constructor.parameters += JsParameter(this) } function.scope.declareFreshName("controller").apply {
constructor.parameters.add(constructor.parameters.lastIndex, JsParameter(this))
}
} }
else { else {
null null
} }
val interceptorRef = if (context.metadata.isLambda) { val interceptorRef = lastParameter!!.makeRef()
val interceptorName = function.scope.declareFreshName("interceptor")
constructor.parameters += JsParameter(interceptorName)
interceptorName.makeRef()
}
else {
lastParameter!!.makeRef()
}
val parameterNames = (function.parameters.map { it.name } + innerFunction?.parameters?.map { it.name }.orEmpty()).toSet() val parameterNames = (function.parameters.map { it.name } + innerFunction?.parameters?.map { it.name }.orEmpty()).toSet()
constructor.body.statements.run { constructor.body.statements.run {
@@ -156,20 +148,14 @@ class CoroutineFunctionTransformer(private val program: JsProgram, private val f
if (context.metadata.hasReceiver) { if (context.metadata.hasReceiver) {
instantiation.arguments += JsLiteral.THIS instantiation.arguments += JsLiteral.THIS
} }
instantiation.arguments += function.parameters.map { it.name.makeRef() } val parameters = function.parameters + innerFunction?.parameters.orEmpty()
if (innerFunction != null) { instantiation.arguments += parameters.dropLast(1).map { it.name.makeRef() }
instantiation.arguments += innerFunction.parameters.map { it.name.makeRef() }
}
if (function.coroutineMetadata!!.hasController) { if (function.coroutineMetadata!!.hasController) {
instantiation.arguments += JsLiteral.THIS instantiation.arguments += JsLiteral.THIS
} }
if (context.metadata.isLambda) { instantiation.arguments += parameters.last().name.makeRef()
val interceptorParamName = functionWithBody.scope.declareFreshName("interceptor")
functionWithBody.parameters += JsParameter(interceptorParamName)
instantiation.arguments += interceptorParamName.makeRef()
}
val suspendedName = functionWithBody.scope.declareFreshName("suspended") val suspendedName = functionWithBody.scope.declareFreshName("suspended")
functionWithBody.parameters += JsParameter(suspendedName) functionWithBody.parameters += JsParameter(suspendedName)
@@ -5829,6 +5829,12 @@ public class JsCodegenBoxTestGenerated extends AbstractJsCodegenBoxTest {
doTest(fileName); doTest(fileName);
} }
@TestMetadata("recursiveSuspend.kt")
public void testRecursiveSuspend() throws Exception {
String fileName = KotlinTestUtils.navigationMetadata("compiler/testData/codegen/box/coroutines/recursiveSuspend.kt");
doTest(fileName);
}
@TestMetadata("returnByLabel.kt") @TestMetadata("returnByLabel.kt")
public void testReturnByLabel() throws Exception { public void testReturnByLabel() throws Exception {
String fileName = KotlinTestUtils.navigationMetadata("compiler/testData/codegen/box/coroutines/returnByLabel.kt"); String fileName = KotlinTestUtils.navigationMetadata("compiler/testData/codegen/box/coroutines/returnByLabel.kt");
@@ -137,7 +137,7 @@ private fun translateFunctionCall(
if (resolvedCall.resultingDescriptor.isSuspend) { if (resolvedCall.resultingDescriptor.isSuspend) {
if (context.isInStateMachine) { if (context.isInStateMachine) {
context.currentBlock.statements += JsAstUtils.asSyntheticStatement(callExpression.apply { isSuspend = true }) context.currentBlock.statements += JsAstUtils.asSyntheticStatement(callExpression.apply { isSuspend = true })
val coroutineRef = TranslationUtils.translateContinuationArgument(context, resolvedCall) val coroutineRef = TranslationUtils.translateContinuationArgument(context)
return context.defineTemporary(JsNameRef("\$\$coroutineResult\$\$", coroutineRef).apply { return context.defineTemporary(JsNameRef("\$\$coroutineResult\$\$", coroutineRef).apply {
sideEffects = SideEffectKind.DEPENDS_ON_STATE sideEffects = SideEffectKind.DEPENDS_ON_STATE
coroutineResult = true coroutineResult = true
@@ -21,7 +21,6 @@ import org.jetbrains.annotations.NotNull;
import org.jetbrains.annotations.Nullable; import org.jetbrains.annotations.Nullable;
import org.jetbrains.kotlin.descriptors.*; import org.jetbrains.kotlin.descriptors.*;
import org.jetbrains.kotlin.descriptors.annotations.Annotations; import org.jetbrains.kotlin.descriptors.annotations.Annotations;
import org.jetbrains.kotlin.descriptors.impl.AnonymousFunctionDescriptor;
import org.jetbrains.kotlin.descriptors.impl.LocalVariableDescriptor; import org.jetbrains.kotlin.descriptors.impl.LocalVariableDescriptor;
import org.jetbrains.kotlin.descriptors.impl.TypeAliasConstructorDescriptor; import org.jetbrains.kotlin.descriptors.impl.TypeAliasConstructorDescriptor;
import org.jetbrains.kotlin.incremental.components.NoLookupLocation; import org.jetbrains.kotlin.incremental.components.NoLookupLocation;
@@ -109,7 +108,7 @@ public class TranslationContext {
} }
if (declarationDescriptor instanceof FunctionDescriptor) { if (declarationDescriptor instanceof FunctionDescriptor) {
FunctionDescriptor function = (FunctionDescriptor) declarationDescriptor; FunctionDescriptor function = (FunctionDescriptor) declarationDescriptor;
if (function.isSuspend() && !(function instanceof AnonymousFunctionDescriptor)) { if (function.isSuspend()) {
ClassDescriptor continuationDescriptor = ClassDescriptor continuationDescriptor =
DescriptorUtilKt.findContinuationClassDescriptor(getCurrentModule(), NoLookupLocation.FROM_BACKEND); DescriptorUtilKt.findContinuationClassDescriptor(getCurrentModule(), NoLookupLocation.FROM_BACKEND);
@@ -27,7 +27,6 @@ import org.jetbrains.kotlin.js.translate.expression.translateAndAliasParameters
import org.jetbrains.kotlin.js.translate.expression.translateFunction import org.jetbrains.kotlin.js.translate.expression.translateFunction
import org.jetbrains.kotlin.js.translate.expression.wrapWithInlineMetadata import org.jetbrains.kotlin.js.translate.expression.wrapWithInlineMetadata
import org.jetbrains.kotlin.js.translate.general.TranslatorVisitor import org.jetbrains.kotlin.js.translate.general.TranslatorVisitor
import org.jetbrains.kotlin.js.translate.reference.ReferenceTranslator
import org.jetbrains.kotlin.js.translate.utils.* import org.jetbrains.kotlin.js.translate.utils.*
import org.jetbrains.kotlin.psi.* import org.jetbrains.kotlin.psi.*
import org.jetbrains.kotlin.resolve.descriptorUtil.isExtensionProperty import org.jetbrains.kotlin.resolve.descriptorUtil.isExtensionProperty
@@ -99,17 +98,11 @@ abstract class AbstractDeclarationVisitor : TranslatorVisitor<Unit>() {
context: TranslationContext context: TranslationContext
): JsExpression { ): JsExpression {
val function = context.getFunctionObject(descriptor) val function = context.getFunctionObject(descriptor)
var innerContext = context.newDeclaration(descriptor).translateAndAliasParameters(descriptor, function.parameters) val innerContext = context.newDeclaration(descriptor).translateAndAliasParameters(descriptor, function.parameters)
if (descriptor.isSuspend) { if (descriptor.isSuspend) {
if (descriptor.requiresStateMachineTransformation(context)) { if (descriptor.requiresStateMachineTransformation(context)) {
function.fillCoroutineMetadata(context, descriptor, hasController = false, isLambda = false) function.fillCoroutineMetadata(context, descriptor, hasController = false, isLambda = false)
innerContext = innerContext.innerContextWithAliased(descriptor, JsAstUtils.stateMachineReceiver())
}
else {
val continuationRef = ReferenceTranslator.translateAsValueReference(
innerContext.continuationParameterDescriptor!!, innerContext)
innerContext = innerContext.innerContextWithAliased(descriptor, continuationRef)
} }
} }
@@ -29,6 +29,8 @@ import org.jetbrains.kotlin.js.translate.context.TranslationContext
import org.jetbrains.kotlin.js.translate.reference.CallExpressionTranslator.shouldBeInlined import org.jetbrains.kotlin.js.translate.reference.CallExpressionTranslator.shouldBeInlined
import org.jetbrains.kotlin.js.translate.utils.BindingUtils import org.jetbrains.kotlin.js.translate.utils.BindingUtils
import org.jetbrains.kotlin.js.translate.utils.FunctionBodyTranslator.translateFunctionBody import org.jetbrains.kotlin.js.translate.utils.FunctionBodyTranslator.translateFunctionBody
import org.jetbrains.kotlin.js.translate.utils.JsAstUtils
import org.jetbrains.kotlin.js.translate.utils.requiresStateMachineTransformation
import org.jetbrains.kotlin.psi.KtDeclarationWithBody import org.jetbrains.kotlin.psi.KtDeclarationWithBody
import org.jetbrains.kotlin.resolve.DescriptorUtils import org.jetbrains.kotlin.resolve.DescriptorUtils
import org.jetbrains.kotlin.resolve.descriptorUtil.hasDefaultValue import org.jetbrains.kotlin.resolve.descriptorUtil.hasDefaultValue
@@ -66,6 +68,12 @@ fun TranslationContext.translateAndAliasParameters(
if (continuationDescriptor != null) { if (continuationDescriptor != null) {
val jsParameter = JsParameter(getNameForDescriptor(continuationDescriptor)) val jsParameter = JsParameter(getNameForDescriptor(continuationDescriptor))
targetList += jsParameter targetList += jsParameter
aliases[continuationDescriptor] = if (!descriptor.requiresStateMachineTransformation(this)) {
JsAstUtils.pureFqn(jsParameter.name, null)
}
else {
JsAstUtils.stateMachineReceiver()
}
} }
return this.innerContextWithDescriptorsAliased(aliases) return this.innerContextWithDescriptorsAliased(aliases)
@@ -45,14 +45,8 @@ class LiteralFunctionTranslator(context: TranslationContext) : AbstractTranslato
val lambda = invokingContext.getFunctionObject(descriptor) val lambda = invokingContext.getFunctionObject(descriptor)
val aliases = mutableMapOf<DeclarationDescriptor, JsExpression>()
if (descriptor.isCoroutineLambda) {
aliases.put(descriptor, JsAstUtils.stateMachineReceiver())
}
val functionContext = invokingContext val functionContext = invokingContext
.newFunctionBodyWithUsageTracker(lambda, descriptor) .newFunctionBodyWithUsageTracker(lambda, descriptor)
.innerContextWithDescriptorsAliased(aliases)
.translateAndAliasParameters(descriptor, lambda.parameters) .translateAndAliasParameters(descriptor, lambda.parameters)
descriptor.valueParameters.forEach { descriptor.valueParameters.forEach {
@@ -76,8 +70,7 @@ class LiteralFunctionTranslator(context: TranslationContext) : AbstractTranslato
} }
lambdaCreator.name.staticRef = lambdaCreator lambdaCreator.name.staticRef = lambdaCreator
lambdaCreator.fillCoroutineMetadata(invokingContext, descriptor) lambdaCreator.fillCoroutineMetadata(invokingContext, descriptor)
return lambdaCreator.withCapturedParameters(descriptor, descriptor.wrapContextForCoroutineIfNecessary(functionContext), return lambdaCreator.withCapturedParameters(descriptor, functionContext, invokingContext)
invokingContext)
} }
lambda.name = invokingContext.getInnerNameForDescriptor(descriptor) lambda.name = invokingContext.getInnerNameForDescriptor(descriptor)
@@ -112,15 +105,6 @@ class LiteralFunctionTranslator(context: TranslationContext) : AbstractTranslato
} }
} }
private fun CallableMemberDescriptor.wrapContextForCoroutineIfNecessary(context: TranslationContext): TranslationContext {
return if (isCoroutineLambda) {
context.innerContextWithDescriptorsAliased(mapOf(this to JsAstUtils.stateMachineReceiver()))
}
else {
context
}
}
fun JsFunction.withCapturedParameters( fun JsFunction.withCapturedParameters(
descriptor: CallableMemberDescriptor, descriptor: CallableMemberDescriptor,
context: TranslationContext, context: TranslationContext,
@@ -155,7 +155,7 @@ class CallArgumentTranslator private constructor(
val callableDescriptor = resolvedCall.resultingDescriptor val callableDescriptor = resolvedCall.resultingDescriptor
if (callableDescriptor is FunctionDescriptor && callableDescriptor.isSuspend) { if (callableDescriptor is FunctionDescriptor && callableDescriptor.isSuspend) {
val facadeName = context().getNameForDescriptor(TranslationUtils.getCoroutineProperty(context(), "facade")) val facadeName = context().getNameForDescriptor(TranslationUtils.getCoroutineProperty(context(), "facade"))
result.add(JsAstUtils.pureFqn(facadeName, TranslationUtils.translateContinuationArgument(context(), resolvedCall))) result.add(JsAstUtils.pureFqn(facadeName, TranslationUtils.translateContinuationArgument(context())))
} }
removeLastUndefinedArguments(result) removeLastUndefinedArguments(result)
@@ -24,7 +24,6 @@ import org.jetbrains.kotlin.descriptors.impl.LocalVariableAccessorDescriptor;
import org.jetbrains.kotlin.descriptors.impl.LocalVariableDescriptor; import org.jetbrains.kotlin.descriptors.impl.LocalVariableDescriptor;
import org.jetbrains.kotlin.incremental.components.NoLookupLocation; import org.jetbrains.kotlin.incremental.components.NoLookupLocation;
import org.jetbrains.kotlin.js.backend.ast.*; import org.jetbrains.kotlin.js.backend.ast.*;
import org.jetbrains.kotlin.js.backend.ast.JsBinaryOperator;
import org.jetbrains.kotlin.js.translate.context.Namer; import org.jetbrains.kotlin.js.translate.context.Namer;
import org.jetbrains.kotlin.js.translate.context.TemporaryConstVariable; import org.jetbrains.kotlin.js.translate.context.TemporaryConstVariable;
import org.jetbrains.kotlin.js.translate.context.TranslationContext; import org.jetbrains.kotlin.js.translate.context.TranslationContext;
@@ -38,7 +37,6 @@ import org.jetbrains.kotlin.psi.*;
import org.jetbrains.kotlin.resolve.BindingContext; import org.jetbrains.kotlin.resolve.BindingContext;
import org.jetbrains.kotlin.resolve.BindingContextUtils; import org.jetbrains.kotlin.resolve.BindingContextUtils;
import org.jetbrains.kotlin.resolve.DescriptorUtils; import org.jetbrains.kotlin.resolve.DescriptorUtils;
import org.jetbrains.kotlin.resolve.calls.model.ResolvedCall;
import org.jetbrains.kotlin.resolve.descriptorUtil.DescriptorUtilsKt; import org.jetbrains.kotlin.resolve.descriptorUtil.DescriptorUtilsKt;
import org.jetbrains.kotlin.resolve.inline.InlineUtil; import org.jetbrains.kotlin.resolve.inline.InlineUtil;
import org.jetbrains.kotlin.serialization.deserialization.FindClassInModuleKt; import org.jetbrains.kotlin.serialization.deserialization.FindClassInModuleKt;
@@ -81,7 +79,7 @@ public final class TranslationUtils {
} }
@NotNull @NotNull
public static String getAccessorFunctionName(@NotNull FunctionDescriptor descriptor) { private static String getAccessorFunctionName(@NotNull FunctionDescriptor descriptor) {
boolean isGetter = descriptor instanceof PropertyGetterDescriptor || descriptor instanceof LocalVariableAccessorDescriptor.Getter; boolean isGetter = descriptor instanceof PropertyGetterDescriptor || descriptor instanceof LocalVariableAccessorDescriptor.Getter;
return isGetter ? "get" : "set"; return isGetter ? "get" : "set";
} }
@@ -227,12 +225,6 @@ public final class TranslationUtils {
return Translation.translateAsExpression(left, context, block); return Translation.translateAsExpression(left, context, block);
} }
@NotNull
public static JsExpression translateRightExpression(@NotNull TranslationContext context,
@NotNull KtBinaryExpression expression) {
return translateRightExpression(context, expression, context.dynamicContext().jsBlock());
}
@NotNull @NotNull
public static JsExpression translateRightExpression( public static JsExpression translateRightExpression(
@NotNull TranslationContext context, @NotNull TranslationContext context,
@@ -350,19 +342,13 @@ public final class TranslationUtils {
} }
@NotNull @NotNull
public static JsExpression translateContinuationArgument(@NotNull TranslationContext context, @NotNull ResolvedCall<?> resolvedCall) { public static JsExpression translateContinuationArgument(@NotNull TranslationContext context) {
CallableDescriptor continuationDescriptor = CallableDescriptor continuationDescriptor = getEnclosingContinuationParameter(context);
context.bindingContext().get(BindingContext.ENCLOSING_SUSPEND_FUNCTION_FOR_SUSPEND_FUNCTION_CALL, resolvedCall.getCall());
if (continuationDescriptor == null) {
continuationDescriptor = getEnclosingContinuationParameter(context);
}
return ReferenceTranslator.translateAsValueReference(continuationDescriptor, context); return ReferenceTranslator.translateAsValueReference(continuationDescriptor, context);
} }
@NotNull @NotNull
public static VariableDescriptor getEnclosingContinuationParameter(@NotNull TranslationContext context) { private static VariableDescriptor getEnclosingContinuationParameter(@NotNull TranslationContext context) {
VariableDescriptor result = context.getContinuationParameterDescriptor(); VariableDescriptor result = context.getContinuationParameterDescriptor();
if (result == null) { if (result == null) {
assert context.getParent() != null; assert context.getParent() != null;
@@ -395,13 +381,6 @@ public final class TranslationUtils {
.iterator().next(); .iterator().next();
} }
@NotNull
public static FunctionDescriptor getCoroutineResumeFunction(@NotNull TranslationContext context) {
return getCoroutineBaseClass(context).getUnsubstitutedMemberScope()
.getContributedFunctions(Name.identifier("resume"), NoLookupLocation.FROM_DESERIALIZATION)
.iterator().next();
}
public static boolean isOverridableFunctionWithDefaultParameters(@NotNull FunctionDescriptor descriptor) { public static boolean isOverridableFunctionWithDefaultParameters(@NotNull FunctionDescriptor descriptor) {
return DescriptorUtilsKt.hasOrInheritsParametersWithDefaultValue(descriptor) && return DescriptorUtilsKt.hasOrInheritsParametersWithDefaultValue(descriptor) &&
!(descriptor instanceof ConstructorDescriptor) && !(descriptor instanceof ConstructorDescriptor) &&