JVM_IR: Use IrInlineReferenceLocator for inline local var lowering.

This commit is contained in:
Mads Ager
2019-11-29 11:21:08 +01:00
committed by max-kammerer
parent e2a1cb1077
commit 8cd6dc0cd9
2 changed files with 36 additions and 35 deletions
@@ -15,7 +15,7 @@ import org.jetbrains.kotlin.ir.declarations.IrFunction
import org.jetbrains.kotlin.ir.expressions.* import org.jetbrains.kotlin.ir.expressions.*
import org.jetbrains.kotlin.ir.visitors.acceptChildrenVoid import org.jetbrains.kotlin.ir.visitors.acceptChildrenVoid
internal class IrInlineReferenceLocator(private val context: JvmBackendContext) : IrElementVisitorVoidWithContext() { internal open class IrInlineReferenceLocator(private val context: JvmBackendContext) : IrElementVisitorVoidWithContext() {
val inlineReferences = mutableSetOf<IrCallableReference>() val inlineReferences = mutableSetOf<IrCallableReference>()
// For crossinline lambdas, the call site is null as it's probably in a separate class somewhere. // For crossinline lambdas, the call site is null as it's probably in a separate class somewhere.
@@ -36,7 +36,7 @@ internal class IrInlineReferenceLocator(private val context: JvmBackendContext)
continue continue
if (valueArgument is IrPropertyReference) { if (valueArgument is IrPropertyReference) {
inlineReferences.add(valueArgument) handleInlineFunctionCallableReferenceParam(valueArgument)
continue continue
} }
@@ -46,16 +46,24 @@ internal class IrInlineReferenceLocator(private val context: JvmBackendContext)
else -> null else -> null
} ?: continue } ?: continue
inlineReferences.add(reference) handleInlineFunctionCallableReferenceParam(reference)
if (valueArgument is IrBlock && valueArgument.origin.isLambda) { if (valueArgument is IrBlock && valueArgument.origin.isLambda) {
lambdaToCallSite[reference.symbol.owner] = val declaration = if (parameter.isCrossinline) null else currentScope!!.irElement as IrDeclaration
if (parameter.isCrossinline) null else currentScope!!.irElement as IrDeclaration handleInlineFunctionLambdaParam(reference.symbol.owner, function, declaration)
} }
} }
} }
return super.visitFunctionAccess(expression) return super.visitFunctionAccess(expression)
} }
open fun handleInlineFunctionCallableReferenceParam(valueArgument: IrCallableReference) {
inlineReferences.add(valueArgument)
}
open fun handleInlineFunctionLambdaParam(lambda: IrFunction, callee: IrFunction, callSite: IrDeclaration?) {
lambdaToCallSite[lambda] = callSite
}
companion object { companion object {
fun scan(context: JvmBackendContext, element: IrElement) = fun scan(context: JvmBackendContext, element: IrElement) =
IrInlineReferenceLocator(context).apply { element.accept(this, null) } IrInlineReferenceLocator(context).apply { element.accept(this, null) }
@@ -10,15 +10,20 @@ 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
import org.jetbrains.kotlin.backend.jvm.codegen.mapClass import org.jetbrains.kotlin.backend.jvm.codegen.mapClass
import org.jetbrains.kotlin.ir.IrElement import org.jetbrains.kotlin.backend.jvm.ir.IrInlineReferenceLocator
import org.jetbrains.kotlin.ir.builders.* import org.jetbrains.kotlin.ir.builders.createTmpVariable
import org.jetbrains.kotlin.ir.builders.irBlockBody
import org.jetbrains.kotlin.ir.builders.irInt
import org.jetbrains.kotlin.ir.declarations.IrDeclaration
import org.jetbrains.kotlin.ir.declarations.IrDeclarationOrigin import org.jetbrains.kotlin.ir.declarations.IrDeclarationOrigin
import org.jetbrains.kotlin.ir.declarations.IrFile import org.jetbrains.kotlin.ir.declarations.IrFile
import org.jetbrains.kotlin.ir.declarations.IrFunction import org.jetbrains.kotlin.ir.declarations.IrFunction
import org.jetbrains.kotlin.ir.expressions.* import org.jetbrains.kotlin.ir.expressions.IrBlockBody
import org.jetbrains.kotlin.ir.expressions.IrCallableReference
import org.jetbrains.kotlin.ir.expressions.IrExpressionBody
import org.jetbrains.kotlin.ir.util.parentAsClass import org.jetbrains.kotlin.ir.util.parentAsClass
import org.jetbrains.kotlin.ir.visitors.IrElementVisitorVoid
import org.jetbrains.kotlin.ir.visitors.acceptChildrenVoid import org.jetbrains.kotlin.ir.visitors.acceptChildrenVoid
import org.jetbrains.kotlin.ir.visitors.acceptVoid
import org.jetbrains.kotlin.load.java.JvmAbi import org.jetbrains.kotlin.load.java.JvmAbi
internal val fakeInliningLocalVariablesLowering = makeIrFilePhase( internal val fakeInliningLocalVariablesLowering = makeIrFilePhase(
@@ -27,41 +32,29 @@ internal val fakeInliningLocalVariablesLowering = makeIrFilePhase(
description = "Add fake locals to identify the range of inlined functions and lambdas" description = "Add fake locals to identify the range of inlined functions and lambdas"
) )
internal class FakeInliningLocalVariablesLowering(val context: JvmBackendContext) : IrElementVisitorVoid, FileLoweringPass { internal class FakeInliningLocalVariablesLowering(val context: JvmBackendContext) : IrInlineReferenceLocator(context), FileLoweringPass {
override fun lower(irFile: IrFile) { override fun lower(irFile: IrFile) {
irFile.acceptChildrenVoid(this) irFile.acceptVoid(this)
} }
override fun visitElement(element: IrElement) { override fun visitFunctionNew(declaration: IrFunction) {
element.acceptChildrenVoid(this)
}
override fun visitCall(expression: IrCall) {
expression.acceptChildrenVoid(this)
val callee = expression.symbol.owner
if (callee.isInline) {
for (i in 0 until expression.valueArgumentsCount) {
val argument = expression.getValueArgument(i)
if ((argument is IrBlock) && argument.origin == IrStatementOrigin.LAMBDA) {
val lastStatement = argument.statements.last()
if (lastStatement is IrFunctionReference) {
val localFunForLambda = lastStatement.symbol.owner
if (localFunForLambda.origin == IrDeclarationOrigin.LOCAL_FUNCTION_FOR_LAMBDA) {
localFunForLambda.addFakeInliningLocalVariablesForArguments(callee)
}
}
}
}
}
}
override fun visitFunction(declaration: IrFunction) {
declaration.acceptChildrenVoid(this) declaration.acceptChildrenVoid(this)
if (declaration.isInline && !declaration.origin.isSynthetic && declaration.body != null) { if (declaration.isInline && !declaration.origin.isSynthetic && declaration.body != null) {
declaration.addFakeInliningLocalVariables() declaration.addFakeInliningLocalVariables()
} }
} }
override fun handleInlineFunctionCallableReferenceParam(valueArgument: IrCallableReference) {
// Do not record inline function callable reference parameters. They will not be used.
}
override fun handleInlineFunctionLambdaParam(lambda: IrFunction, callee: IrFunction, callSite: IrDeclaration?) {
// Do not record lambda parameters. Instead deal with them now.
if (lambda.origin == IrDeclarationOrigin.LOCAL_FUNCTION_FOR_LAMBDA) {
lambda.addFakeInliningLocalVariablesForArguments(callee)
}
}
private fun IrFunction.addFakeInliningLocalVariables() { private fun IrFunction.addFakeInliningLocalVariables() {
val currentFunctionName = context.methodSignatureMapper.mapFunctionName(this) val currentFunctionName = context.methodSignatureMapper.mapFunctionName(this)
val localName = "${JvmAbi.LOCAL_VARIABLE_NAME_PREFIX_INLINE_FUNCTION}$currentFunctionName" val localName = "${JvmAbi.LOCAL_VARIABLE_NAME_PREFIX_INLINE_FUNCTION}$currentFunctionName"