PSI2IR: Generate adapted callable references
Callable reference is "adapted" if it requires some adaptation to an
expected function type - e.g., when a reference to
```
fun foo(vararg xs: Int): Int
```
is used where `(Int, Int, Int) -> Int` is expected.
For such callable references we generate the following IR (in
pseudo-Kotlin):
```
{
fun foo'(p0: Int, p1: Int, p2: Int): Int {
return [| foo(p0, p1, p2) |]
}
::foo'
}
```
where `[| foo(p0, p1, p2) |]` is calling function `foo` with arguments
`p0`, `p1`, and `p2`, as they were mapped by callable reference
resolution.
This commit is contained in:
@@ -134,8 +134,11 @@ val IrType.isBoxedArray: Boolean
|
||||
fun IrType.getArrayElementType(irBuiltIns: IrBuiltIns): IrType =
|
||||
if (isBoxedArray)
|
||||
((this as IrSimpleType).arguments.single() as IrTypeProjection).type
|
||||
else
|
||||
irBuiltIns.primitiveArrayElementTypes.getValue(this.classOrNull!!)
|
||||
else {
|
||||
val classifier = this.classOrNull!!
|
||||
irBuiltIns.primitiveArrayElementTypes[classifier]
|
||||
?: throw AssertionError("Primitive array expected: $classifier")
|
||||
}
|
||||
|
||||
val IrStatementOrigin?.isLambda: Boolean
|
||||
get() = this == IrStatementOrigin.LAMBDA || this == IrStatementOrigin.ANONYMOUS_FUNCTION
|
||||
|
||||
+269
-3
@@ -16,19 +16,30 @@
|
||||
|
||||
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.impl.IrFunctionImpl
|
||||
import org.jetbrains.kotlin.ir.declarations.impl.IrValueParameterImpl
|
||||
import org.jetbrains.kotlin.ir.descriptors.WrappedSimpleFunctionDescriptor
|
||||
import org.jetbrains.kotlin.ir.descriptors.WrappedValueParameterDescriptor
|
||||
import org.jetbrains.kotlin.ir.expressions.*
|
||||
import org.jetbrains.kotlin.ir.expressions.impl.*
|
||||
import org.jetbrains.kotlin.ir.symbols.IrConstructorSymbol
|
||||
import org.jetbrains.kotlin.ir.symbols.IrFunctionSymbol
|
||||
import org.jetbrains.kotlin.ir.symbols.IrVariableSymbol
|
||||
import org.jetbrains.kotlin.ir.util.referenceClassifier
|
||||
import org.jetbrains.kotlin.ir.util.referenceFunction
|
||||
import org.jetbrains.kotlin.psi.KtCallableReferenceExpression
|
||||
import org.jetbrains.kotlin.psi.KtClassLiteralExpression
|
||||
import org.jetbrains.kotlin.psi.KtElement
|
||||
import org.jetbrains.kotlin.ir.util.withScope
|
||||
import org.jetbrains.kotlin.name.Name
|
||||
import org.jetbrains.kotlin.psi.*
|
||||
import org.jetbrains.kotlin.psi.psiUtil.endOffset
|
||||
import org.jetbrains.kotlin.psi.psiUtil.startOffsetSkippingComments
|
||||
import org.jetbrains.kotlin.psi2ir.intermediate.CallBuilder
|
||||
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
|
||||
|
||||
@@ -61,6 +72,15 @@ class ReflectionReferencesGenerator(statementGenerator: StatementGenerator) : St
|
||||
|
||||
val callBuilder = unwrapCallableDescriptorAndTypeArguments(resolvedCall, context.extensions.samConversion)
|
||||
|
||||
if (resolvedCall.valueArguments.isNotEmpty()) {
|
||||
val adaptedCallableReference = generateAdaptedCallableReference(ktCallableReference, callBuilder)
|
||||
if (adaptedCallableReference.hasResultAdaptation ||
|
||||
!isTrivialArgumentAdaptation(adaptedCallableReference.irAdapteeCall)
|
||||
) {
|
||||
return adaptedCallableReference.irReferenceExpression
|
||||
}
|
||||
}
|
||||
|
||||
return statementGenerator.generateCallReceiver(
|
||||
ktCallableReference,
|
||||
resolvedDescriptor,
|
||||
@@ -79,6 +99,252 @@ class ReflectionReferencesGenerator(statementGenerator: StatementGenerator) : St
|
||||
}
|
||||
}
|
||||
|
||||
private fun isTrivialArgumentAdaptation(irAdapteeCall: IrFunctionAccessExpression): Boolean {
|
||||
for (i in 0 until irAdapteeCall.valueArgumentsCount) {
|
||||
val irValueArgument = irAdapteeCall.getValueArgument(i) ?: return false
|
||||
if (irValueArgument is IrVararg) {
|
||||
val irVarargElements = irValueArgument.elements
|
||||
if (irVarargElements.size != 1 || irVarargElements[0] !is IrSpreadElement) return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
private class AdaptedCallableReference(
|
||||
val irReferenceExpression: IrExpression,
|
||||
val irAdapteeCall: IrFunctionAccessExpression,
|
||||
val hasResultAdaptation: Boolean
|
||||
)
|
||||
|
||||
private fun generateAdaptedCallableReference(
|
||||
ktCallableReference: KtCallableReferenceExpression,
|
||||
callBuilder: CallBuilder
|
||||
): AdaptedCallableReference {
|
||||
val adapteeDescriptor = callBuilder.descriptor
|
||||
if (adapteeDescriptor !is FunctionDescriptor) {
|
||||
throw AssertionError("Function descriptor expected in adapted callable reference: $adapteeDescriptor")
|
||||
}
|
||||
|
||||
val startOffset = ktCallableReference.startOffsetSkippingComments
|
||||
val endOffset = ktCallableReference.endOffset
|
||||
|
||||
val adapteeSymbol = context.symbolTable.referenceFunction(adapteeDescriptor.original)
|
||||
|
||||
val ktFunctionalType = getTypeInferredByFrontendOrFail(ktCallableReference)
|
||||
val irFunctionalType = ktFunctionalType.toIrType()
|
||||
|
||||
val ktFunctionalTypeArguments = ktFunctionalType.arguments
|
||||
val ktExpectedReturnType = ktFunctionalTypeArguments.last().type
|
||||
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)
|
||||
|
||||
irAdapterFun.body = IrBlockBodyImpl(startOffset, endOffset).apply {
|
||||
if (KotlinBuiltIns.isUnit(ktExpectedReturnType))
|
||||
statements.add(irCall)
|
||||
else
|
||||
statements.add(IrReturnImpl(startOffset, endOffset, context.irBuiltIns.nothingType, irAdapterFun.symbol, irCall))
|
||||
}
|
||||
|
||||
val irBlock = IrBlockImpl(startOffset, endOffset, irFunctionalType).apply {
|
||||
statements.add(irAdapterFun)
|
||||
statements.add(
|
||||
IrFunctionReferenceImpl(
|
||||
startOffset, endOffset,
|
||||
irFunctionalType,
|
||||
irAdapterFun.symbol,
|
||||
0,
|
||||
irAdapterFun.valueParameters.size
|
||||
)
|
||||
)
|
||||
}
|
||||
|
||||
return AdaptedCallableReference(
|
||||
irBlock, irAdapteeCallInner,
|
||||
KotlinBuiltIns.isUnit(ktExpectedReturnType) && !KotlinBuiltIns.isUnit(adapteeDescriptor.returnType!!)
|
||||
)
|
||||
}
|
||||
|
||||
private fun createAdapteeCall(
|
||||
startOffset: Int,
|
||||
endOffset: Int,
|
||||
ktCallableReference: KtCallableReferenceExpression,
|
||||
adapteeSymbol: IrFunctionSymbol,
|
||||
callBuilder: CallBuilder,
|
||||
irAdapterFun: IrSimpleFunction
|
||||
): Pair<IrExpression, IrFunctionAccessExpression> {
|
||||
val resolvedCall = callBuilder.original
|
||||
val resolvedDescriptor = resolvedCall.resultingDescriptor
|
||||
|
||||
var irAdapteeCall: IrFunctionAccessExpression? = null
|
||||
val irCall = statementGenerator.generateCallReceiver(
|
||||
ktCallableReference,
|
||||
resolvedDescriptor,
|
||||
resolvedCall.dispatchReceiver, resolvedCall.extensionReceiver,
|
||||
isSafe = false
|
||||
).call { dispatchReceiverValue, extensionReceiverValue ->
|
||||
val irType = resolvedDescriptor.returnType!!.toIrType()
|
||||
|
||||
val irAdapteeCallInner =
|
||||
if (resolvedDescriptor is ConstructorDescriptor)
|
||||
IrConstructorCallImpl.fromSymbolDescriptor(
|
||||
startOffset, endOffset, irType,
|
||||
adapteeSymbol as IrConstructorSymbol
|
||||
)
|
||||
else
|
||||
IrCallImpl(
|
||||
startOffset, endOffset, irType,
|
||||
adapteeSymbol,
|
||||
origin = null, superQualifierSymbol = null
|
||||
)
|
||||
|
||||
context.callToSubstitutedDescriptorMap[irAdapteeCallInner] = resolvedDescriptor
|
||||
|
||||
irAdapteeCallInner.dispatchReceiver = dispatchReceiverValue?.load()
|
||||
irAdapteeCallInner.extensionReceiver = extensionReceiverValue?.load()
|
||||
|
||||
irAdapteeCallInner.putTypeArguments(callBuilder.typeArguments) { it.toIrType() }
|
||||
|
||||
putAdaptedValueArguments(startOffset, endOffset, irAdapteeCallInner, irAdapterFun, resolvedCall)
|
||||
|
||||
irAdapteeCall = irAdapteeCallInner
|
||||
|
||||
irAdapteeCallInner
|
||||
}
|
||||
|
||||
return Pair(irCall, irAdapteeCall!!)
|
||||
}
|
||||
|
||||
private fun putAdaptedValueArguments(
|
||||
startOffset: Int,
|
||||
endOffset: Int,
|
||||
irAdapteeCall: IrFunctionAccessExpression,
|
||||
irAdapterFun: IrSimpleFunction,
|
||||
resolvedCall: ResolvedCall<*>
|
||||
) {
|
||||
val adaptedArguments = resolvedCall.valueArguments
|
||||
if (adaptedArguments.isEmpty()) {
|
||||
throw AssertionError("Callable reference with adapted arguments expected: ${resolvedCall.call.callElement.text}")
|
||||
}
|
||||
|
||||
for ((valueParameter, valueArgument) in adaptedArguments) {
|
||||
irAdapteeCall.putValueArgument(
|
||||
valueParameter.index,
|
||||
adaptResolvedValueArgument(startOffset, endOffset, valueArgument, irAdapterFun, valueParameter)
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
private fun adaptResolvedValueArgument(
|
||||
startOffset: Int,
|
||||
endOffset: Int,
|
||||
resolvedValueArgument: ResolvedValueArgument,
|
||||
irAdapterFun: IrSimpleFunction,
|
||||
valueParameter: ValueParameterDescriptor
|
||||
): IrExpression? =
|
||||
when (resolvedValueArgument) {
|
||||
is DefaultValueArgument ->
|
||||
null
|
||||
is VarargValueArgument ->
|
||||
IrVarargImpl(
|
||||
startOffset, endOffset,
|
||||
valueParameter.type.toIrType(), valueParameter.varargElementType!!.toIrType(),
|
||||
resolvedValueArgument.arguments.map {
|
||||
adaptValueArgument(startOffset, endOffset, it, irAdapterFun)
|
||||
}
|
||||
)
|
||||
is ExpressionValueArgument -> {
|
||||
val valueArgument = resolvedValueArgument.valueArgument!!
|
||||
|
||||
adaptValueArgument(startOffset, endOffset, valueArgument, irAdapterFun) as IrExpression
|
||||
}
|
||||
else ->
|
||||
throw AssertionError("Unexpected ResolvedValueArgument: $resolvedValueArgument")
|
||||
}
|
||||
|
||||
private fun adaptValueArgument(
|
||||
startOffset: Int,
|
||||
endOffset: Int,
|
||||
valueArgument: ValueArgument,
|
||||
irAdapterFun: IrSimpleFunction
|
||||
): IrVarargElement =
|
||||
when (valueArgument) {
|
||||
is FakeImplicitSpreadValueArgumentForCallableReference ->
|
||||
IrSpreadElementImpl(
|
||||
startOffset, endOffset,
|
||||
adaptValueArgument(startOffset, endOffset, valueArgument.expression, irAdapterFun) as IrExpression
|
||||
)
|
||||
|
||||
is FakePositionalValueArgumentForCallableReference -> {
|
||||
val irAdapterParameter = irAdapterFun.valueParameters[valueArgument.index]
|
||||
IrGetValueImpl(startOffset, endOffset, irAdapterParameter.type, irAdapterParameter.symbol)
|
||||
}
|
||||
|
||||
else ->
|
||||
throw AssertionError("Unexpected ValueArgument: $valueArgument")
|
||||
}
|
||||
|
||||
private fun createAdapterFun(
|
||||
startOffset: Int,
|
||||
endOffset: Int,
|
||||
adapteeDescriptor: FunctionDescriptor,
|
||||
ktExpectedParameterTypes: List<KotlinType>,
|
||||
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
|
||||
) { irAdapterSymbol ->
|
||||
IrFunctionImpl(
|
||||
startOffset, endOffset, adapterOrigin,
|
||||
irAdapterSymbol,
|
||||
adapteeDescriptor.name,
|
||||
Visibilities.LOCAL,
|
||||
Modality.FINAL,
|
||||
ktExpectedReturnType.toIrType(),
|
||||
isInline = adapteeDescriptor.isInline, // TODO ?
|
||||
isExternal = false,
|
||||
isTailrec = false,
|
||||
isSuspend = adapteeDescriptor.isSuspend, // TODO ?
|
||||
isOperator = adapteeDescriptor.isOperator, // TODO ?
|
||||
isExpect = false,
|
||||
isFakeOverride = false
|
||||
).also { irAdapterFun ->
|
||||
adapterFunctionDescriptor.bind(irAdapterFun)
|
||||
|
||||
context.symbolTable.withScope(adapterFunctionDescriptor) {
|
||||
irAdapterFun.metadata = MetadataSource.Function(adapteeDescriptor)
|
||||
|
||||
irAdapterFun.dispatchReceiverParameter = null
|
||||
irAdapterFun.extensionReceiverParameter = null
|
||||
|
||||
ktExpectedParameterTypes.mapIndexedTo(irAdapterFun.valueParameters) { index, ktExpectedParameterType ->
|
||||
val adapterValueParameterDescriptor = WrappedValueParameterDescriptor()
|
||||
context.symbolTable.declareValueParameter(
|
||||
startOffset, endOffset, adapterOrigin, adapterValueParameterDescriptor,
|
||||
ktExpectedParameterType.toIrType()
|
||||
) { irAdapterParameterSymbol ->
|
||||
IrValueParameterImpl(
|
||||
startOffset, endOffset, adapterOrigin,
|
||||
irAdapterParameterSymbol,
|
||||
Name.identifier("p$index"),
|
||||
index,
|
||||
ktExpectedParameterType.toIrType(),
|
||||
varargElementType = null, isCrossinline = false, isNoinline = false
|
||||
).also { irAdapterValueParameter ->
|
||||
adapterValueParameterDescriptor.bind(irAdapterValueParameter)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fun generateCallableReference(
|
||||
ktElement: KtElement,
|
||||
type: KotlinType,
|
||||
|
||||
@@ -377,20 +377,18 @@ class DumpTreeFromSourceLineVisitor(
|
||||
|
||||
internal fun IrMemberAccessExpression.getValueParameterNamesForDebug(): List<String> {
|
||||
val expectedCount = valueArgumentsCount
|
||||
return if (this is IrDeclarationReference && symbol.isBound) {
|
||||
if (symbol.isBound) {
|
||||
val owner = symbol.owner
|
||||
if (owner is IrFunction) {
|
||||
(0 until expectedCount).map {
|
||||
return (0 until expectedCount).map {
|
||||
if (it < owner.valueParameters.size)
|
||||
owner.valueParameters[it].name.asString()
|
||||
else
|
||||
"${it + 1}"
|
||||
}
|
||||
} else {
|
||||
getPlaceholderParameterNames(expectedCount)
|
||||
}
|
||||
} else
|
||||
getPlaceholderParameterNames(expectedCount)
|
||||
}
|
||||
return getPlaceholderParameterNames(expectedCount)
|
||||
}
|
||||
|
||||
internal fun getPlaceholderParameterNames(expectedCount: Int) =
|
||||
|
||||
Reference in New Issue
Block a user