PSI2IR: capture adapted reference receiver to a local val if required

This commit is contained in:
Dmitry Petrov
2020-01-21 14:26:49 +03:00
parent 38b90b7fbd
commit 57bbfbbfcc
9 changed files with 488 additions and 10 deletions
@@ -18,9 +18,7 @@ package org.jetbrains.kotlin.psi2ir.generators
import org.jetbrains.kotlin.builtins.KotlinBuiltIns
import org.jetbrains.kotlin.descriptors.*
import org.jetbrains.kotlin.ir.declarations.IrDeclarationOrigin
import org.jetbrains.kotlin.ir.declarations.IrSimpleFunction
import org.jetbrains.kotlin.ir.declarations.MetadataSource
import org.jetbrains.kotlin.ir.declarations.*
import org.jetbrains.kotlin.ir.declarations.impl.IrFunctionImpl
import org.jetbrains.kotlin.ir.declarations.impl.IrValueParameterImpl
import org.jetbrains.kotlin.ir.descriptors.WrappedSimpleFunctionDescriptor
@@ -42,6 +40,7 @@ import org.jetbrains.kotlin.resolve.BindingContext
import org.jetbrains.kotlin.resolve.calls.model.*
import org.jetbrains.kotlin.types.KotlinType
import org.jetbrains.kotlin.types.expressions.DoubleColonLHS
import org.jetbrains.kotlin.utils.addIfNotNull
class ReflectionReferencesGenerator(statementGenerator: StatementGenerator) : StatementGeneratorExtension(statementGenerator) {
@@ -99,6 +98,8 @@ class ReflectionReferencesGenerator(statementGenerator: StatementGenerator) : St
}
}
private val adapterOrigin = IrDeclarationOrigin.DEFINED // TODO special declaration origin for callable reference adapter?
private fun isTrivialArgumentAdaptation(irAdapteeCall: IrFunctionAccessExpression): Boolean {
for (i in 0 until irAdapteeCall.valueArgumentsCount) {
val irValueArgument = irAdapteeCall.getValueArgument(i) ?: return false
@@ -138,8 +139,12 @@ class ReflectionReferencesGenerator(statementGenerator: StatementGenerator) : St
val ktExpectedParameterTypes = ktFunctionalTypeArguments.take(ktFunctionalTypeArguments.size - 1).map { it.type }
val irAdapterFun = createAdapterFun(startOffset, endOffset, adapteeDescriptor, ktExpectedParameterTypes, ktExpectedReturnType)
val (irCall, irAdapteeCallInner) =
createAdapteeCall(startOffset, endOffset, ktCallableReference, adapteeSymbol, callBuilder, irAdapterFun)
val adapteeCall = createAdapteeCall(startOffset, endOffset, ktCallableReference, adapteeSymbol, callBuilder, irAdapterFun)
val irCall = adapteeCall.callExpression
val irAdapteeCallInner = adapteeCall.innerCallExpression
val tmpDispatchReceiver = adapteeCall.tmpDispatchReceiver
val tmpExtensionReceiver = adapteeCall.tmpExtensionReceiver
irAdapterFun.body = IrBlockBodyImpl(startOffset, endOffset).apply {
if (KotlinBuiltIns.isUnit(ktExpectedReturnType))
@@ -149,6 +154,8 @@ class ReflectionReferencesGenerator(statementGenerator: StatementGenerator) : St
}
val irBlock = IrBlockImpl(startOffset, endOffset, irFunctionalType).apply {
statements.addIfNotNull(tmpDispatchReceiver)
statements.addIfNotNull(tmpExtensionReceiver)
statements.add(irAdapterFun)
statements.add(
IrFunctionReferenceImpl(
@@ -167,6 +174,13 @@ class ReflectionReferencesGenerator(statementGenerator: StatementGenerator) : St
)
}
private class AdapteeCall(
val callExpression: IrExpression,
val innerCallExpression: IrFunctionAccessExpression,
val tmpDispatchReceiver: IrVariable?,
val tmpExtensionReceiver: IrVariable?
)
private fun createAdapteeCall(
startOffset: Int,
endOffset: Int,
@@ -174,11 +188,14 @@ class ReflectionReferencesGenerator(statementGenerator: StatementGenerator) : St
adapteeSymbol: IrFunctionSymbol,
callBuilder: CallBuilder,
irAdapterFun: IrSimpleFunction
): Pair<IrExpression, IrFunctionAccessExpression> {
): AdapteeCall {
val resolvedCall = callBuilder.original
val resolvedDescriptor = resolvedCall.resultingDescriptor
var irAdapteeCall: IrFunctionAccessExpression? = null
var tmpDispatchReceiver: IrVariable? = null
var tmpExtensionReceiver: IrVariable? = null
val irCall = statementGenerator.generateCallReceiver(
ktCallableReference,
resolvedDescriptor,
@@ -202,8 +219,28 @@ class ReflectionReferencesGenerator(statementGenerator: StatementGenerator) : St
context.callToSubstitutedDescriptorMap[irAdapteeCallInner] = resolvedDescriptor
irAdapteeCallInner.dispatchReceiver = dispatchReceiverValue?.load()
irAdapteeCallInner.extensionReceiver = extensionReceiverValue?.load()
val irDispatchReceiver = dispatchReceiverValue?.load()
val irExtensionReceiver = extensionReceiverValue?.load()
if (irDispatchReceiver != null) {
if (irDispatchReceiver.isSafeToUseWithoutCopying()) {
irAdapteeCallInner.dispatchReceiver = irDispatchReceiver
} else {
val irVariable = statementGenerator.scope.createTemporaryVariable(irDispatchReceiver, "this")
irAdapteeCallInner.dispatchReceiver = IrGetValueImpl(startOffset, endOffset, irVariable.symbol)
tmpDispatchReceiver = irVariable
}
}
if (irExtensionReceiver != null) {
if (irExtensionReceiver.isSafeToUseWithoutCopying()) {
irAdapteeCallInner.extensionReceiver = irExtensionReceiver
} else {
val irVariable = statementGenerator.scope.createTemporaryVariable(irExtensionReceiver, "receiver")
irAdapteeCallInner.extensionReceiver = IrGetValueImpl(startOffset, endOffset, irVariable.symbol)
tmpExtensionReceiver = irVariable
}
}
irAdapteeCallInner.putTypeArguments(callBuilder.typeArguments) { it.toIrType() }
@@ -214,9 +251,18 @@ class ReflectionReferencesGenerator(statementGenerator: StatementGenerator) : St
irAdapteeCallInner
}
return Pair(irCall, irAdapteeCall!!)
return AdapteeCall(irCall, irAdapteeCall!!, tmpDispatchReceiver, tmpExtensionReceiver)
}
private fun IrExpression.isSafeToUseWithoutCopying() =
this is IrGetObjectValue ||
this is IrGetEnumValue ||
this is IrConst<*> ||
this is IrGetValue && symbol.isBound && symbol.owner.isImmutable()
private fun IrValueDeclaration.isImmutable() =
this is IrValueParameter || this is IrVariable && !this.isVar
private fun putAdaptedValueArguments(
startOffset: Int,
endOffset: Int,
@@ -294,7 +340,7 @@ class ReflectionReferencesGenerator(statementGenerator: StatementGenerator) : St
ktExpectedReturnType: KotlinType
): IrSimpleFunction {
val adapterFunctionDescriptor = WrappedSimpleFunctionDescriptor()
val adapterOrigin = IrDeclarationOrigin.DEFINED // TODO special declaration origin for callable reference adapter?
return context.symbolTable.declareSimpleFunction(
startOffset, endOffset, adapterOrigin, adapterFunctionDescriptor