[JS IR BE] Fix callable reference to make it possible to be bound

This commit is contained in:
Roman Artemev
2018-05-31 19:55:43 +03:00
committed by Roman Artemev
parent 11c330effd
commit eeb16a38e8
2 changed files with 129 additions and 77 deletions
@@ -5,10 +5,8 @@
package org.jetbrains.kotlin.ir.backend.js.lower package org.jetbrains.kotlin.ir.backend.js.lower
import org.jetbrains.kotlin.backend.common.DeclarationContainerLoweringPass
import org.jetbrains.kotlin.backend.common.FileLoweringPass import org.jetbrains.kotlin.backend.common.FileLoweringPass
import org.jetbrains.kotlin.backend.common.lower.copyAsValueParameter import org.jetbrains.kotlin.backend.common.lower.copyAsValueParameter
import org.jetbrains.kotlin.backend.common.runOnFilePostfix
import org.jetbrains.kotlin.descriptors.CallableDescriptor import org.jetbrains.kotlin.descriptors.CallableDescriptor
import org.jetbrains.kotlin.descriptors.ClassConstructorDescriptor import org.jetbrains.kotlin.descriptors.ClassConstructorDescriptor
import org.jetbrains.kotlin.descriptors.ValueParameterDescriptor import org.jetbrains.kotlin.descriptors.ValueParameterDescriptor
@@ -23,40 +21,47 @@ import org.jetbrains.kotlin.ir.declarations.*
import org.jetbrains.kotlin.ir.declarations.impl.IrValueParameterImpl import org.jetbrains.kotlin.ir.declarations.impl.IrValueParameterImpl
import org.jetbrains.kotlin.ir.expressions.* import org.jetbrains.kotlin.ir.expressions.*
import org.jetbrains.kotlin.ir.expressions.impl.IrCallImpl import org.jetbrains.kotlin.ir.expressions.impl.IrCallImpl
import org.jetbrains.kotlin.ir.symbols.IrFunctionSymbol
import org.jetbrains.kotlin.ir.symbols.IrSimpleFunctionSymbol import org.jetbrains.kotlin.ir.symbols.IrSimpleFunctionSymbol
import org.jetbrains.kotlin.ir.symbols.IrValueParameterSymbol import org.jetbrains.kotlin.ir.symbols.IrValueParameterSymbol
import org.jetbrains.kotlin.ir.symbols.IrValueSymbol import org.jetbrains.kotlin.ir.symbols.IrValueSymbol
import org.jetbrains.kotlin.ir.util.transformFlat
import org.jetbrains.kotlin.ir.visitors.* import org.jetbrains.kotlin.ir.visitors.*
import org.jetbrains.kotlin.resolve.descriptorUtil.fqNameSafe import org.jetbrains.kotlin.resolve.descriptorUtil.fqNameSafe
import org.jetbrains.kotlin.types.KotlinType
class CallableReferenceLowering(val context: JsIrBackendContext) : FileLoweringPass, DeclarationContainerLoweringPass { class CallableReferenceLowering(val context: JsIrBackendContext) : FileLoweringPass {
private val callableToGetterFunction = mutableMapOf<CallableDescriptor, IrFunction>() private data class CallableReferenceKey(
val declaration: IrFunction,
val hasDispatchReference: Boolean,
val hasExtensionReceiver: Boolean
)
// TODO: replace descriptor usage with symbol instead private val callableToGetterFunction = mutableMapOf<CallableReferenceKey, IrFunction>()
private val collectedReferenceMap = mutableMapOf<IrDeclaration, IrCallableReference>() private val collectedReferenceMap = mutableMapOf<CallableReferenceKey, IrCallableReference>()
private val callableNameConst = JsIrBuilder.buildString(context.irBuiltIns.string, Namer.KCALLABLE_NAME) private val callableNameConst = JsIrBuilder.buildString(context.irBuiltIns.string, Namer.KCALLABLE_NAME)
private val getterConst = JsIrBuilder.buildString(context.irBuiltIns.string, Namer.KPROPERTY_GET) private val getterConst = JsIrBuilder.buildString(context.irBuiltIns.string, Namer.KPROPERTY_GET)
private val setterConst = JsIrBuilder.buildString(context.irBuiltIns.string, Namer.KPROPERTY_SET) private val setterConst = JsIrBuilder.buildString(context.irBuiltIns.string, Namer.KPROPERTY_SET)
private val newDeclarations = mutableListOf<IrDeclaration>()
override fun lower(irFile: IrFile) { override fun lower(irFile: IrFile) {
irFile.acceptVoid(CallableReferenceCollector()) irFile.acceptVoid(CallableReferenceCollector())
runOnFilePostfix(irFile) buildClosures()
irFile.transformChildrenVoid(CallableReferenceTransformer()) irFile.transformChildrenVoid(CallableReferenceTransformer())
irFile.declarations += newDeclarations
} }
private fun makeCallableKey(declaration: IrFunction, reference: IrCallableReference) =
CallableReferenceKey(declaration, reference.dispatchReceiver != null, reference.extensionReceiver != null)
inner class CallableReferenceCollector : IrElementVisitorVoid { inner class CallableReferenceCollector : IrElementVisitorVoid {
override fun visitFunctionReference(expression: IrFunctionReference) { override fun visitFunctionReference(expression: IrFunctionReference) {
collectedReferenceMap[expression.symbol.owner] = expression collectedReferenceMap[makeCallableKey(expression.symbol.owner, expression)] = expression
} }
override fun visitPropertyReference(expression: IrPropertyReference) { override fun visitPropertyReference(expression: IrPropertyReference) {
//Note: The getter is taken because the `invoke()` function of the resulted reference has to be corresponding getter call //Note: The getter is taken because the `invoke()` function of the resulted reference has to be corresponding getter call
collectedReferenceMap[expression.getter!!.owner] = expression collectedReferenceMap[makeCallableKey(expression.getter!!.owner, expression)] = expression
} }
override fun visitElement(element: IrElement) { override fun visitElement(element: IrElement) {
@@ -64,28 +69,38 @@ class CallableReferenceLowering(val context: JsIrBackendContext) : FileLoweringP
} }
} }
override fun lower(irDeclarationContainer: IrDeclarationContainer) { private fun buildClosures() {
irDeclarationContainer.declarations.transformFlat { d -> for (v in collectedReferenceMap.values) {
collectedReferenceMap[d]?.accept(object : IrElementVisitor<List<IrDeclaration>, IrFunction> { newDeclarations += v.accept(object : IrElementVisitor<List<IrDeclaration>, Nothing?> {
override fun visitElement(element: IrElement, data: IrFunction) = error("Unreachable execution") override fun visitElement(element: IrElement, data: Nothing?) = error("Unreachable execution")
override fun visitFunctionReference(expression: IrFunctionReference, data: IrFunction) = override fun visitFunctionReference(expression: IrFunctionReference, data: Nothing?) =
lowerKFunctionReference(data, expression) lowerKFunctionReference(expression.symbol.owner, expression)
override fun visitPropertyReference(expression: IrPropertyReference, data: IrFunction) = override fun visitPropertyReference(expression: IrPropertyReference, data: Nothing?) =
lowerKPropertyReference(data, expression) lowerKPropertyReference(expression.getter!!.owner, expression)
}, d as IrFunction) }, null)
} }
} }
inner class CallableReferenceTransformer : IrElementTransformerVoid() { inner class CallableReferenceTransformer : IrElementTransformerVoid() {
override fun visitCallableReference(expression: IrCallableReference) = callableToGetterFunction[expression.descriptor]?.let { override fun visitFunctionReference(expression: IrFunctionReference): IrExpression {
redirectToFunction(expression, it) return callableToGetterFunction[makeCallableKey(expression.symbol.owner, expression)]?.let {
} ?: expression redirectToFunction(expression, it)
} ?: expression
}
override fun visitPropertyReference(expression: IrPropertyReference): IrExpression {
return callableToGetterFunction[makeCallableKey(expression.getter!!.owner, expression)]?.let {
redirectToFunction(expression, it)
} ?: expression
}
private fun redirectToFunction(callable: IrCallableReference, newTarget: IrFunction) = private fun redirectToFunction(callable: IrCallableReference, newTarget: IrFunction) =
IrCallImpl(callable.startOffset, callable.endOffset, newTarget.symbol, callable.origin).apply { IrCallImpl(callable.startOffset, callable.endOffset, newTarget.symbol, callable.origin).apply {
copyTypeArgumentsFrom(callable) copyTypeArgumentsFrom(callable)
var index = 0 var index = 0
callable.dispatchReceiver?.let { putValueArgument(index++, it) }
callable.extensionReceiver?.let { putValueArgument(index++, it) }
for (i in 0 until callable.valueArgumentsCount) { for (i in 0 until callable.valueArgumentsCount) {
val arg = callable.getValueArgument(i) val arg = callable.getValueArgument(i)
if (arg != null) { if (arg != null) {
@@ -95,7 +110,7 @@ class CallableReferenceLowering(val context: JsIrBackendContext) : FileLoweringP
} }
} }
private fun createClosureGetterName(descriptor: CallableDescriptor) = createHelperFunctionName(descriptor, "KReferenceGet") private fun createFunctionClosureGetterName(descriptor: CallableDescriptor) = createHelperFunctionName(descriptor, "KReferenceGet")
private fun createPropertyClosureGetterName(descriptor: CallableDescriptor) = createHelperFunctionName(descriptor, "KPropertyGet") private fun createPropertyClosureGetterName(descriptor: CallableDescriptor) = createHelperFunctionName(descriptor, "KPropertyGet")
private fun createClosureInstanceName(descriptor: CallableDescriptor) = createHelperFunctionName(descriptor, "KReferenceClosure") private fun createClosureInstanceName(descriptor: CallableDescriptor) = createHelperFunctionName(descriptor, "KReferenceClosure")
@@ -134,8 +149,8 @@ class CallableReferenceLowering(val context: JsIrBackendContext) : FileLoweringP
// KFunctionN<Foo, T2, ..., TN, TReturn>, arguments.size = N + 1 // KFunctionN<Foo, T2, ..., TN, TReturn>, arguments.size = N + 1
val refGetFunction = buildGetFunction(declaration, functionReference.type, createClosureGetterName(declaration.descriptor)) val refGetFunction = buildGetFunction(declaration, functionReference, createFunctionClosureGetterName(declaration.descriptor))
val refClosureFunction = buildClosureFunction(declaration, refGetFunction) val refClosureFunction = buildClosureFunction(declaration, refGetFunction, functionReference)
val additionalDeclarations = generateGetterBodyWithGuard(refGetFunction) { val additionalDeclarations = generateGetterBodyWithGuard(refGetFunction) {
val irClosureReference = JsIrBuilder.buildFunctionReference(functionReference.type, refClosureFunction.symbol) val irClosureReference = JsIrBuilder.buildFunctionReference(functionReference.type, refClosureFunction.symbol)
@@ -151,9 +166,10 @@ class CallableReferenceLowering(val context: JsIrBackendContext) : FileLoweringP
Pair(listOf(irVar, irSetName), irVarSymbol) Pair(listOf(irVar, irSetName), irVarSymbol)
} }
callableToGetterFunction[functionReference.descriptor] = refGetFunction
return additionalDeclarations + listOf(declaration, refGetFunction) callableToGetterFunction[makeCallableKey(declaration, functionReference)] = refGetFunction
return additionalDeclarations + listOf(refGetFunction)
} }
private fun lowerKPropertyReference(getterDeclaration: IrFunction, propertyReference: IrPropertyReference): List<IrDeclaration> { private fun lowerKPropertyReference(getterDeclaration: IrFunction, propertyReference: IrPropertyReference): List<IrDeclaration> {
@@ -177,10 +193,10 @@ class CallableReferenceLowering(val context: JsIrBackendContext) : FileLoweringP
// } // }
val getterName = createPropertyClosureGetterName(propertyReference.descriptor) val getterName = createPropertyClosureGetterName(propertyReference.descriptor)
val refGetFunction = buildGetFunction(propertyReference.getter!!.owner, propertyReference.type, getterName) val refGetFunction = buildGetFunction(propertyReference.getter!!.owner, propertyReference, getterName)
val getterFunction = propertyReference.getter?.let { buildClosureFunction(it.owner, refGetFunction) }!! val getterFunction = propertyReference.getter?.let { buildClosureFunction(it.owner, refGetFunction, propertyReference) }!!
val setterFunction = propertyReference.setter?.let { buildClosureFunction(it.owner, refGetFunction) } val setterFunction = propertyReference.setter?.let { buildClosureFunction(it.owner, refGetFunction, propertyReference) }
val additionalDeclarations = generateGetterBodyWithGuard(refGetFunction) { val additionalDeclarations = generateGetterBodyWithGuard(refGetFunction) {
val statements = mutableListOf<IrStatement>() val statements = mutableListOf<IrStatement>()
@@ -191,38 +207,35 @@ class CallableReferenceLowering(val context: JsIrBackendContext) : FileLoweringP
statements += JsIrBuilder.buildVar(irVarSymbol, irGetReference) statements += JsIrBuilder.buildVar(irVarSymbol, irGetReference)
JsIrBuilder.buildCall(context.intrinsics.jsSetJSField.symbol).run { statements += JsIrBuilder.buildCall(context.intrinsics.jsSetJSField.symbol).apply {
putValueArgument(0, JsIrBuilder.buildGetValue(irVarSymbol)) putValueArgument(0, JsIrBuilder.buildGetValue(irVarSymbol))
putValueArgument(1, getterConst) putValueArgument(1, getterConst)
putValueArgument(2, JsIrBuilder.buildGetValue(irVarSymbol)) putValueArgument(2, JsIrBuilder.buildGetValue(irVarSymbol))
statements += this
} }
if (setterFunction != null) { if (setterFunction != null) {
val setterFunctionType = context.builtIns.getFunction(setterFunction.valueParameters.size + 1) val setterFunctionType = context.builtIns.getFunction(setterFunction.valueParameters.size + 1)
val irSetReference = JsIrBuilder.buildFunctionReference(setterFunctionType.defaultType, setterFunction.symbol) val irSetReference = JsIrBuilder.buildFunctionReference(setterFunctionType.defaultType, setterFunction.symbol)
JsIrBuilder.buildCall(context.intrinsics.jsSetJSField.symbol).run { statements += JsIrBuilder.buildCall(context.intrinsics.jsSetJSField.symbol).apply {
putValueArgument(0, JsIrBuilder.buildGetValue(irVarSymbol)) putValueArgument(0, JsIrBuilder.buildGetValue(irVarSymbol))
putValueArgument(1, setterConst) putValueArgument(1, setterConst)
putValueArgument(2, irSetReference) putValueArgument(2, irSetReference)
statements += this
} }
} }
// TODO: fill other fields of callable reference (returnType, parameters, isFinal, etc.) // TODO: fill other fields of callable reference (returnType, parameters, isFinal, etc.)
JsIrBuilder.buildCall(context.intrinsics.jsSetJSField.symbol).run { statements += JsIrBuilder.buildCall(context.intrinsics.jsSetJSField.symbol).apply {
putValueArgument(0, JsIrBuilder.buildGetValue(irVarSymbol)) putValueArgument(0, JsIrBuilder.buildGetValue(irVarSymbol))
putValueArgument(1, callableNameConst) putValueArgument(1, callableNameConst)
putValueArgument(2, JsIrBuilder.buildString(context.irBuiltIns.string, getReferenceName(propertyReference.descriptor))) putValueArgument(2, JsIrBuilder.buildString(context.irBuiltIns.string, getReferenceName(propertyReference.descriptor)))
statements += this
} }
Pair(statements, irVarSymbol) Pair(statements, irVarSymbol)
} }
callableToGetterFunction[propertyReference.descriptor] = refGetFunction callableToGetterFunction[makeCallableKey(getterDeclaration, propertyReference)] = refGetFunction
return additionalDeclarations + listOf(getterDeclaration, refGetFunction) return additionalDeclarations + listOf(refGetFunction)
} }
private fun generateGetterBodyWithGuard( private fun generateGetterBodyWithGuard(
@@ -265,6 +278,7 @@ class CallableReferenceLowering(val context: JsIrBackendContext) : FileLoweringP
returnValue = JsIrBuilder.buildGetValue(varSymbol) returnValue = JsIrBuilder.buildGetValue(varSymbol)
returnStatements = emptyList() returnStatements = emptyList()
} }
statements += JsIrBuilder.buildReturn(getterFunction.symbol, returnValue) statements += JsIrBuilder.buildReturn(getterFunction.symbol, returnValue)
getterFunction.body = JsIrBuilder.buildBlockBody(statements) getterFunction.body = JsIrBuilder.buildBlockBody(statements)
@@ -272,17 +286,23 @@ class CallableReferenceLowering(val context: JsIrBackendContext) : FileLoweringP
} }
private fun generateSignatureForClosure( private fun generateSignatureForClosure(
callable: IrFunctionSymbol, callable: IrFunction,
getter: IrSimpleFunctionSymbol, getter: IrFunction,
closure: IrSimpleFunctionSymbol closure: IrSimpleFunctionSymbol,
reference: IrCallableReference
): List<IrValueParameterSymbol> { ): List<IrValueParameterSymbol> {
val result = mutableListOf<IrValueParameterSymbol>() val result = mutableListOf<IrValueParameterSymbol>()
callable.owner.dispatchReceiverParameter?.run { result.add(JsSymbolBuilder.buildValueParameter(closure, result.size, type)) } if (callable.dispatchReceiverParameter != null && reference.dispatchReceiver == null) {
callable.owner.extensionReceiverParameter?.run { result.add(JsSymbolBuilder.buildValueParameter(closure, result.size, type)) } result.add(JsSymbolBuilder.buildValueParameter(closure, result.size, callable.dispatchReceiverParameter!!.type))
}
for (i in getter.owner.valueParameters.size until callable.owner.valueParameters.size) { if (callable.extensionReceiverParameter != null && reference.extensionReceiver == null) {
val param = callable.owner.valueParameters[i] result.add(JsSymbolBuilder.buildValueParameter(closure, result.size, callable.extensionReceiverParameter!!.type))
}
for (i in getter.valueParameters.size until callable.valueParameters.size) {
val param = callable.valueParameters[i]
val paramName = param.name.run { if (!isSpecial) identifier else null } val paramName = param.name.run { if (!isSpecial) identifier else null }
result += JsSymbolBuilder.buildValueParameter(closure, result.size, param.type, paramName) result += JsSymbolBuilder.buildValueParameter(closure, result.size, param.type, paramName)
} }
@@ -290,17 +310,33 @@ class CallableReferenceLowering(val context: JsIrBackendContext) : FileLoweringP
return result return result
} }
private fun buildGetFunction(declaration: IrFunction, callableType: KotlinType, getterName: String): IrSimpleFunction { private fun buildGetFunction(declaration: IrFunction, reference: IrCallableReference, getterName: String): IrSimpleFunction {
val callableType = reference.type
val closureParams = callableType.arguments.dropLast(1) // drop return type val closureParams = callableType.arguments.dropLast(1) // drop return type
var kFunctionValueParamsCount = closureParams.size var kFunctionValueParamsCount = closureParams.size
if (declaration.dispatchReceiverParameter != null) kFunctionValueParamsCount--
if (declaration.extensionReceiverParameter != null) kFunctionValueParamsCount-- if (declaration.dispatchReceiverParameter != null && reference.dispatchReceiver == null) {
kFunctionValueParamsCount--
}
if (declaration.extensionReceiverParameter != null && reference.extensionReceiver == null) {
kFunctionValueParamsCount--
}
assert(kFunctionValueParamsCount >= 0) assert(kFunctionValueParamsCount >= 0)
// The `getter` function takes only closure parameters // The `getter` function takes only closure parameters
val getterValueParameters = declaration.valueParameters.dropLast(kFunctionValueParamsCount) val receivers = mutableListOf<IrValueParameter>()
if (reference.dispatchReceiver != null) {
receivers += declaration.dispatchReceiverParameter!!
}
if (reference.extensionReceiver != null) {
receivers += declaration.extensionReceiverParameter!!
}
val getterValueParameters = receivers + declaration.valueParameters.dropLast(kFunctionValueParamsCount)
val refGetSymbol = JsSymbolBuilder.buildSimpleFunction(declaration.descriptor.containingDeclaration, getterName).apply { val refGetSymbol = JsSymbolBuilder.buildSimpleFunction(declaration.descriptor.containingDeclaration, getterName).apply {
initialize( initialize(
@@ -317,39 +353,57 @@ class CallableReferenceLowering(val context: JsIrBackendContext) : FileLoweringP
} }
} }
private fun buildClosureFunction(declaration: IrFunction, refGetFunction: IrSimpleFunction): IrFunction { private fun buildClosureFunction(
declaration: IrFunction,
refGetFunction: IrSimpleFunction,
reference: IrCallableReference
): IrFunction {
val closureName = createClosureInstanceName(declaration.descriptor) val closureName = createClosureInstanceName(declaration.descriptor)
val refClosureSymbol = JsSymbolBuilder.buildSimpleFunction(refGetFunction.descriptor, closureName) val refClosureSymbol = JsSymbolBuilder.buildSimpleFunction(refGetFunction.descriptor, closureName)
// the params which are passed to closure // the params which are passed to closure
val closureParamSymbols = generateSignatureForClosure(declaration.symbol, refGetFunction.symbol, refClosureSymbol) val closureParamSymbols = generateSignatureForClosure(declaration, refGetFunction, refClosureSymbol, reference)
val closureParamDescriptors = closureParamSymbols.map { it.descriptor as ValueParameterDescriptor } val closureParamDescriptors = closureParamSymbols.map { it.descriptor as ValueParameterDescriptor }
refClosureSymbol.initialize(valueParameters = closureParamDescriptors, type = declaration.returnType) refClosureSymbol.initialize(valueParameters = closureParamDescriptors, type = declaration.returnType)
return JsIrBuilder.buildFunction(refClosureSymbol).apply { val closureFunction = JsIrBuilder.buildFunction(refClosureSymbol)
for (it in closureParamSymbols) {
valueParameters += JsIrBuilder.buildValueParameter(it)
}
val irCall = JsIrBuilder.buildCall(declaration.symbol) for (it in closureParamSymbols) {
closureFunction.valueParameters += JsIrBuilder.buildValueParameter(it)
var p = 0
declaration.dispatchReceiverParameter?.run { irCall.dispatchReceiver = JsIrBuilder.buildGetValue(closureParamSymbols[p++]) }
declaration.extensionReceiverParameter?.run { irCall.extensionReceiver = JsIrBuilder.buildGetValue(closureParamSymbols[p++]) }
var j = 0
for (v in refGetFunction.valueParameters) {
irCall.putValueArgument(j++, JsIrBuilder.buildGetValue(v.symbol))
}
for (i in p until closureParamSymbols.size) {
irCall.putValueArgument(j++, JsIrBuilder.buildGetValue(closureParamSymbols[i]))
}
val irClosureReturn = JsIrBuilder.buildReturn(symbol, irCall)
body = JsIrBuilder.buildBlockBody(listOf(irClosureReturn))
} }
val irCall = JsIrBuilder.buildCall(declaration.symbol)
var cp = 0
var gp = 0
if (declaration.dispatchReceiverParameter != null) {
val dispatchReceiverDeclaration =
if (reference.dispatchReceiver != null) refGetFunction.valueParameters[gp++].symbol else closureParamSymbols[cp++]
irCall.dispatchReceiver = JsIrBuilder.buildGetValue(dispatchReceiverDeclaration)
}
if (declaration.extensionReceiverParameter != null) {
val extensionReceiverDeclaration =
if (reference.extensionReceiver != null) refGetFunction.valueParameters[gp++].symbol else closureParamSymbols[cp++]
irCall.extensionReceiver = JsIrBuilder.buildGetValue(extensionReceiverDeclaration)
}
var j = 0
for (i in gp until refGetFunction.valueParameters.size) {
irCall.putValueArgument(j++, JsIrBuilder.buildGetValue(refGetFunction.valueParameters[i].symbol))
}
for (i in cp until closureParamSymbols.size) {
irCall.putValueArgument(j++, JsIrBuilder.buildGetValue(closureParamSymbols[i]))
}
val irClosureReturn = JsIrBuilder.buildReturn(closureFunction.symbol, irCall)
closureFunction.body = JsIrBuilder.buildBlockBody(listOf(irClosureReturn))
return closureFunction
} }
} }
@@ -40,8 +40,6 @@ object Namer {
val OUTER_NAME = "\$outer" val OUTER_NAME = "\$outer"
val UNREACHABLE_NAME = "\$unreachable" val UNREACHABLE_NAME = "\$unreachable"
val OUTER_FIELD_NAME = "\$outer"
val DELEGATE = "\$delegate" val DELEGATE = "\$delegate"
val ROOT_PACKAGE = "_" val ROOT_PACKAGE = "_"