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:
Dmitry Petrov
2020-01-15 11:49:50 +03:00
parent 89c832b5a0
commit c5f14a29a4
15 changed files with 729 additions and 19 deletions
@@ -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
@@ -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) =