JVM_IR: Support callable references to suspend functions

This commit is contained in:
Ilmir Usmanov
2019-10-10 18:16:19 +03:00
parent e736b782dd
commit b0a0399dd0
9 changed files with 66 additions and 20 deletions
@@ -169,12 +169,14 @@ fun IrDeclarationContainer.addFunction(
name: String, name: String,
returnType: IrType, returnType: IrType,
modality: Modality = Modality.FINAL, modality: Modality = Modality.FINAL,
isStatic: Boolean = false isStatic: Boolean = false,
isSuspend: Boolean = false
): IrSimpleFunction = ): IrSimpleFunction =
addFunction { addFunction {
this.name = Name.identifier(name) this.name = Name.identifier(name)
this.returnType = returnType this.returnType = returnType
this.modality = modality this.modality = modality
this.isSuspend = isSuspend
}.apply { }.apply {
if (!isStatic) { if (!isStatic) {
dispatchReceiverParameter = parentAsClass.thisReceiver!!.copyTo(this) dispatchReceiverParameter = parentAsClass.thisReceiver!!.copyTo(this)
@@ -222,6 +222,24 @@ class JvmSymbols(
fun getJvmFunctionClass(parameterCount: Int): IrClassSymbol = fun getJvmFunctionClass(parameterCount: Int): IrClassSymbol =
jvmFunctionClasses(parameterCount) jvmFunctionClasses(parameterCount)
private val jvmSuspendFunctionClasses = storageManager.createMemoizedFunction { n: Int ->
createClass(FqName("kotlin.jvm.functions.Function${n + 1}"), ClassKind.INTERFACE) { klass ->
for (i in 1..n) {
klass.addTypeParameter("P$i", irBuiltIns.anyNType, Variance.IN_VARIANCE)
}
val returnType = klass.addTypeParameter("R", irBuiltIns.anyNType, Variance.OUT_VARIANCE)
klass.addFunction("invoke", returnType.defaultType, Modality.ABSTRACT, isSuspend = true).apply {
for (i in 1..n) {
addValueParameter("p$i", klass.typeParameters[i - 1].defaultType)
}
}
}
}
fun getJvmSuspendFunctionClass(parameterCount: Int): IrClassSymbol =
jvmSuspendFunctionClasses(parameterCount)
val functionN: IrClassSymbol = createClass(FqName("kotlin.jvm.functions.FunctionN"), ClassKind.INTERFACE) { klass -> val functionN: IrClassSymbol = createClass(FqName("kotlin.jvm.functions.FunctionN"), ClassKind.INTERFACE) { klass ->
val returnType = klass.addTypeParameter("R", irBuiltIns.anyNType, Variance.OUT_VARIANCE) val returnType = klass.addTypeParameter("R", irBuiltIns.anyNType, Variance.OUT_VARIANCE)
@@ -11,17 +11,16 @@ import org.jetbrains.kotlin.backend.common.ir.copyTo
import org.jetbrains.kotlin.backend.common.ir.copyTypeParametersFrom import org.jetbrains.kotlin.backend.common.ir.copyTypeParametersFrom
import org.jetbrains.kotlin.backend.common.ir.isSuspend import org.jetbrains.kotlin.backend.common.ir.isSuspend
import org.jetbrains.kotlin.backend.jvm.JvmBackendContext import org.jetbrains.kotlin.backend.jvm.JvmBackendContext
import org.jetbrains.kotlin.backend.jvm.JvmLoweredDeclarationOrigin
import org.jetbrains.kotlin.codegen.ClassBuilder import org.jetbrains.kotlin.codegen.ClassBuilder
import org.jetbrains.kotlin.codegen.coroutines.CoroutineTransformerMethodVisitor import org.jetbrains.kotlin.codegen.coroutines.CoroutineTransformerMethodVisitor
import org.jetbrains.kotlin.codegen.coroutines.INVOKE_SUSPEND_METHOD_NAME import org.jetbrains.kotlin.codegen.coroutines.INVOKE_SUSPEND_METHOD_NAME
import org.jetbrains.kotlin.codegen.coroutines.SUSPEND_FUNCTION_COMPLETION_PARAMETER_NAME import org.jetbrains.kotlin.codegen.coroutines.SUSPEND_FUNCTION_COMPLETION_PARAMETER_NAME
import org.jetbrains.kotlin.config.isReleaseCoroutines import org.jetbrains.kotlin.config.isReleaseCoroutines
import org.jetbrains.kotlin.ir.IrStatement
import org.jetbrains.kotlin.ir.UNDEFINED_OFFSET import org.jetbrains.kotlin.ir.UNDEFINED_OFFSET
import org.jetbrains.kotlin.ir.builders.declarations.addValueParameter import org.jetbrains.kotlin.ir.builders.declarations.addValueParameter
import org.jetbrains.kotlin.ir.declarations.IrClass import org.jetbrains.kotlin.ir.declarations.*
import org.jetbrains.kotlin.ir.declarations.IrDeclarationOriginImpl
import org.jetbrains.kotlin.ir.declarations.IrFunction
import org.jetbrains.kotlin.ir.declarations.IrSimpleFunction
import org.jetbrains.kotlin.ir.declarations.impl.IrFunctionImpl import org.jetbrains.kotlin.ir.declarations.impl.IrFunctionImpl
import org.jetbrains.kotlin.ir.expressions.IrCall import org.jetbrains.kotlin.ir.expressions.IrCall
import org.jetbrains.kotlin.ir.expressions.IrExpression import org.jetbrains.kotlin.ir.expressions.IrExpression
@@ -104,6 +103,18 @@ internal fun IrFunction.isInvokeOfSuspendLambda(context: JvmBackendContext): Boo
internal fun IrFunction.isInvokeSuspendOfContinuation(context: JvmBackendContext): Boolean = internal fun IrFunction.isInvokeSuspendOfContinuation(context: JvmBackendContext): Boolean =
name.asString() == INVOKE_SUSPEND_METHOD_NAME && parentAsClass in context.suspendFunctionContinuations.values name.asString() == INVOKE_SUSPEND_METHOD_NAME && parentAsClass in context.suspendFunctionContinuations.values
internal fun IrFunction.isInvokeOfCallableReference(): Boolean = isSuspend && name.asString() == "invoke" && (parent as? IrClass)?.let {
// TODO: Should we use different origin for lowered callable references?
it.origin == JvmLoweredDeclarationOrigin.FUNCTION_REFERENCE_IMPL && it.functions.any { it.name.asString() == "getSignature" }
} == true
internal fun IrFunction.isKnownToBeTailCall(): Boolean =
origin == IrDeclarationOrigin.FUNCTION_FOR_DEFAULT_PARAMETER || origin == JvmLoweredDeclarationOrigin.SYNTHETIC_ACCESSOR ||
isInvokeOfCallableReference()
internal fun IrFunction.shouldNotContainSuspendMarkers(context: JvmBackendContext): Boolean =
isInvokeSuspendOfContinuation(context) || isKnownToBeTailCall()
// Transform `suspend fun foo(params): RetType` into `fun foo(params, $completion: Continuation<RetType>): Any?` // Transform `suspend fun foo(params): RetType` into `fun foo(params, $completion: Continuation<RetType>): Any?`
// the result is called 'view', just to be consistent with old backend. // the result is called 'view', just to be consistent with old backend.
internal fun IrFunction.getOrCreateSuspendFunctionViewIfNeeded(context: JvmBackendContext): IrFunction { internal fun IrFunction.getOrCreateSuspendFunctionViewIfNeeded(context: JvmBackendContext): IrFunction {
@@ -118,7 +129,10 @@ private fun IrFunction.suspendFunctionView(context: JvmBackendContext): IrFuncti
val originalDescriptor = this.descriptor val originalDescriptor = this.descriptor
// For SuspendFunction{N}.invoke we need to generate INVOKEINTERFACE Function{N+1}.invoke(...Ljava/lang/Object;)... // For SuspendFunction{N}.invoke we need to generate INVOKEINTERFACE Function{N+1}.invoke(...Ljava/lang/Object;)...
// instead of INVOKEINTERFACE Function{N+1}.invoke(...Lkotlin/coroutines/Continuation;)... // instead of INVOKEINTERFACE Function{N+1}.invoke(...Lkotlin/coroutines/Continuation;)...
val isInvokeOfNumberedSuspendFunction = (symbol.owner.parent as? IrClass)?.defaultType?.isSuspendFunction() == true val isInvokeOfNumberedSuspendFunction = (parent as? IrClass)?.defaultType?.isSuspendFunction() == true
// And we need to generate this function for callable references
val isBridgeInvokeOfCallableReference = origin == IrDeclarationOrigin.BRIDGE &&
(parent as? IrClass)?.origin == JvmLoweredDeclarationOrigin.FUNCTION_REFERENCE_IMPL
val descriptor = val descriptor =
if (originalDescriptor is DescriptorWithContainerSource && originalDescriptor.containerSource != null) if (originalDescriptor is DescriptorWithContainerSource && originalDescriptor.containerSource != null)
WrappedFunctionDescriptorWithContainerSource(originalDescriptor.containerSource!!) WrappedFunctionDescriptorWithContainerSource(originalDescriptor.containerSource!!)
@@ -139,7 +153,7 @@ private fun IrFunction.suspendFunctionView(context: JvmBackendContext): IrFuncti
valueParameters.mapTo(it.valueParameters) { p -> p.copyTo(it) } valueParameters.mapTo(it.valueParameters) { p -> p.copyTo(it) }
it.addValueParameter( it.addValueParameter(
SUSPEND_FUNCTION_COMPLETION_PARAMETER_NAME, SUSPEND_FUNCTION_COMPLETION_PARAMETER_NAME,
if (isInvokeOfNumberedSuspendFunction) context.irBuiltIns.anyNType if (isInvokeOfNumberedSuspendFunction || isBridgeInvokeOfCallableReference) context.irBuiltIns.anyNType
else context.ir.symbols.continuationClass.createType(false, listOf(makeTypeProjection(returnType, Variance.INVARIANT))) else context.ir.symbols.continuationClass.createType(false, listOf(makeTypeProjection(returnType, Variance.INVARIANT)))
) )
val valueParametersMapping = explicitParameters.zip(it.explicitParameters).toMap() val valueParametersMapping = explicitParameters.zip(it.explicitParameters).toMap()
@@ -153,6 +167,11 @@ private fun IrFunction.suspendFunctionView(context: JvmBackendContext): IrFuncti
expression.run { IrGetValueImpl(startOffset, endOffset, type, newParam.symbol, origin) } expression.run { IrGetValueImpl(startOffset, endOffset, type, newParam.symbol, origin) }
} ?: expression } ?: expression
override fun visitClass(declaration: IrClass): IrStatement {
// Do not cross class boundaries inside functions. Otherwise, callable references will try to access wrong $completion.
return declaration
}
override fun visitCall(expression: IrCall): IrExpression { override fun visitCall(expression: IrCall): IrExpression {
if (!expression.isSuspend) return super.visitCall(expression) if (!expression.isSuspend) return super.visitCall(expression)
return super.visitCall(expression.createSuspendFunctionCallViewIfNeeded(context, it, callerIsInlineLambda = false)) return super.visitCall(expression.createSuspendFunctionCallViewIfNeeded(context, it, callerIsInlineLambda = false))
@@ -356,7 +356,7 @@ class ExpressionCodegen(
} }
expression.descriptor is ConstructorDescriptor -> expression.descriptor is ConstructorDescriptor ->
throw AssertionError("IrCall with ConstructorDescriptor: ${expression.javaClass.simpleName}") throw AssertionError("IrCall with ConstructorDescriptor: ${expression.javaClass.simpleName}")
callee.isSuspend && !irFunction.isInvokeSuspendOfContinuation(classCodegen.context) -> callee.isSuspend && !irFunction.shouldNotContainSuspendMarkers(classCodegen.context) ->
addInlineMarker(mv, isStartNotEnd = true) addInlineMarker(mv, isStartNotEnd = true)
} }
@@ -382,13 +382,13 @@ class ExpressionCodegen(
expression.markLineNumber(true) expression.markLineNumber(true)
// Do not generate redundant markers in continuation class. // Do not generate redundant markers in continuation class.
if (callee.isSuspend && !irFunction.isInvokeSuspendOfContinuation(classCodegen.context)) { if (callee.isSuspend && !irFunction.shouldNotContainSuspendMarkers(classCodegen.context)) {
addSuspendMarker(mv, isStartNotEnd = true) addSuspendMarker(mv, isStartNotEnd = true)
} }
callGenerator.genCall(callable, this, expression) callGenerator.genCall(callable, this, expression)
if (callee.isSuspend && !irFunction.isInvokeSuspendOfContinuation(classCodegen.context)) { if (callee.isSuspend && !irFunction.shouldNotContainSuspendMarkers(classCodegen.context)) {
addSuspendMarker(mv, isStartNotEnd = false) addSuspendMarker(mv, isStartNotEnd = false)
addInlineMarker(mv, isStartNotEnd = false) addInlineMarker(mv, isStartNotEnd = false)
} }
@@ -99,9 +99,6 @@ open class FunctionCodegen(
return signature return signature
} }
private fun IrFunction.isKnownToBeTailCall(): Boolean =
origin == IrDeclarationOrigin.FUNCTION_FOR_DEFAULT_PARAMETER || origin == JvmLoweredDeclarationOrigin.SYNTHETIC_ACCESSOR
private fun calculateMethodFlags(isStatic: Boolean): Int { private fun calculateMethodFlags(isStatic: Boolean): Int {
if (irFunction.origin == IrDeclarationOrigin.FUNCTION_FOR_DEFAULT_PARAMETER) { if (irFunction.origin == IrDeclarationOrigin.FUNCTION_FOR_DEFAULT_PARAMETER) {
return Opcodes.ACC_PUBLIC or Opcodes.ACC_SYNTHETIC.let { return Opcodes.ACC_PUBLIC or Opcodes.ACC_SYNTHETIC.let {
@@ -18,6 +18,7 @@ import org.jetbrains.kotlin.backend.common.pop
import org.jetbrains.kotlin.backend.common.push import org.jetbrains.kotlin.backend.common.push
import org.jetbrains.kotlin.backend.jvm.JvmBackendContext import org.jetbrains.kotlin.backend.jvm.JvmBackendContext
import org.jetbrains.kotlin.backend.jvm.codegen.isInlineIrBlock import org.jetbrains.kotlin.backend.jvm.codegen.isInlineIrBlock
import org.jetbrains.kotlin.backend.jvm.codegen.isInvokeOfCallableReference
import org.jetbrains.kotlin.codegen.coroutines.* import org.jetbrains.kotlin.codegen.coroutines.*
import org.jetbrains.kotlin.config.coroutinesPackageFqName import org.jetbrains.kotlin.config.coroutinesPackageFqName
import org.jetbrains.kotlin.descriptors.Modality import org.jetbrains.kotlin.descriptors.Modality
@@ -431,7 +432,9 @@ private class AddContinuationLowering(private val context: JvmBackendContext) :
override fun visitFunction(declaration: IrFunction) { override fun visitFunction(declaration: IrFunction) {
super.visitFunction(declaration) super.visitFunction(declaration)
if (declaration.isSuspend && declaration !in suspendLambdas && !declaration.isInline) { if (declaration.isSuspend && declaration !in suspendLambdas && !declaration.isInline &&
!declaration.isInvokeOfCallableReference()
) {
result.add(declaration) result.add(declaration)
} }
} }
@@ -469,7 +472,7 @@ private class AddContinuationLowering(private val context: JvmBackendContext) :
override fun visitFunctionReference(expression: IrFunctionReference) { override fun visitFunctionReference(expression: IrFunctionReference) {
expression.acceptChildrenVoid(this) expression.acceptChildrenVoid(this)
if (expression.isSuspend && expression !in inlineLambdas) { if (expression.isSuspend && expression !in inlineLambdas && expression.origin == IrStatementOrigin.LAMBDA) {
suspendLambdas += SuspendLambdaInfo( suspendLambdas += SuspendLambdaInfo(
expression.symbol.owner, expression.symbol.owner,
(expression.type as IrSimpleType).arguments.size - 1, (expression.type as IrSimpleType).arguments.size - 1,
@@ -501,4 +504,4 @@ private class AddContinuationLowering(private val context: JvmBackendContext) :
private class SuspendLambdaInfo(val function: IrFunction, val arity: Int, val reference: IrFunctionReference) { private class SuspendLambdaInfo(val function: IrFunction, val arity: Int, val reference: IrFunctionReference) {
lateinit var constructor: IrConstructor lateinit var constructor: IrConstructor
} }
} }
@@ -10,6 +10,7 @@ import org.jetbrains.kotlin.backend.common.IrElementTransformerVoidWithContext
import org.jetbrains.kotlin.backend.common.IrElementVisitorVoidWithContext import org.jetbrains.kotlin.backend.common.IrElementVisitorVoidWithContext
import org.jetbrains.kotlin.backend.common.ir.copyTo import org.jetbrains.kotlin.backend.common.ir.copyTo
import org.jetbrains.kotlin.backend.common.ir.createImplicitParameterDeclarationWithWrappedDescriptor import org.jetbrains.kotlin.backend.common.ir.createImplicitParameterDeclarationWithWrappedDescriptor
import org.jetbrains.kotlin.backend.common.ir.isSuspend
import org.jetbrains.kotlin.backend.common.lower.createIrBuilder import org.jetbrains.kotlin.backend.common.lower.createIrBuilder
import org.jetbrains.kotlin.backend.common.phaser.makeIrFilePhase import org.jetbrains.kotlin.backend.common.phaser.makeIrFilePhase
import org.jetbrains.kotlin.backend.jvm.JvmBackendContext import org.jetbrains.kotlin.backend.jvm.JvmBackendContext
@@ -96,7 +97,10 @@ internal class CallableReferenceLowering(private val context: JvmBackendContext)
private val ignoredFunctionReferences = mutableSetOf<IrFunctionReference>() private val ignoredFunctionReferences = mutableSetOf<IrFunctionReference>()
private val IrFunctionReference.isIgnored: Boolean private val IrFunctionReference.isIgnored: Boolean
get() = !type.isFunctionOrKFunction() || ignoredFunctionReferences.contains(this) get() = (!type.isFunctionOrKFunction() || ignoredFunctionReferences.contains(this)) && !isSuspendCallableReference()
// TODO: Currently, origin of callable references is null. Do we need to create one?
private fun IrFunctionReference.isSuspendCallableReference(): Boolean = isSuspend && origin == null
override fun lower(irFile: IrFile) { override fun lower(irFile: IrFile) {
ignoredFunctionReferences.addAll(InlineReferenceLocator.scan(context, irFile).inlineReferences) ignoredFunctionReferences.addAll(InlineReferenceLocator.scan(context, irFile).inlineReferences)
@@ -162,7 +166,11 @@ internal class CallableReferenceLowering(private val context: JvmBackendContext)
private val typeArgumentsMap = irFunctionReference.typeSubstitutionMap private val typeArgumentsMap = irFunctionReference.typeSubstitutionMap
private val functionSuperClass = private val functionSuperClass =
samSuperType?.classOrNull ?: context.ir.symbols.getJvmFunctionClass(argumentTypes.size) samSuperType?.classOrNull
?: if (irFunctionReference.isSuspend)
context.ir.symbols.getJvmSuspendFunctionClass(argumentTypes.size)
else
context.ir.symbols.getJvmFunctionClass(argumentTypes.size)
private val superMethod = private val superMethod =
functionSuperClass.functions.single { it.owner.modality == Modality.ABSTRACT } functionSuperClass.functions.single { it.owner.modality == Modality.ABSTRACT }
private val superType = private val superType =
@@ -1,4 +1,3 @@
// IGNORE_BACKEND: JVM_IR
// WITH_RUNTIME // WITH_RUNTIME
// WITH_COROUTINES // WITH_COROUTINES
@@ -1,4 +1,4 @@
// IGNORE_BACKEND: JS, JS_IR, JVM_IR // IGNORE_BACKEND: JS, JS_IR
// WITH_RUNTIME // WITH_RUNTIME
fun box(): String { fun box(): String {