JVM_IR: apply TailCallOptimizationLowering to all suspend functions

Even if a function is known to be tail call because it's a compiler
generated bridge, the tail return might still need to be added in case
of Unit return type.
This commit is contained in:
pyos
2020-03-09 14:33:59 +01:00
committed by Ilmir Usmanov
parent dc388f3f3a
commit 735fae0e5a
8 changed files with 82 additions and 21 deletions
@@ -9,10 +9,9 @@ import org.jetbrains.kotlin.backend.common.FileLoweringPass
import org.jetbrains.kotlin.backend.common.ir.isSuspend
import org.jetbrains.kotlin.backend.common.phaser.makeIrFilePhase
import org.jetbrains.kotlin.backend.jvm.JvmBackendContext
import org.jetbrains.kotlin.backend.jvm.codegen.hasContinuation
import org.jetbrains.kotlin.ir.IrStatement
import org.jetbrains.kotlin.ir.declarations.IrFile
import org.jetbrains.kotlin.ir.declarations.IrFunction
import org.jetbrains.kotlin.ir.declarations.IrSimpleFunction
import org.jetbrains.kotlin.ir.expressions.*
import org.jetbrains.kotlin.ir.expressions.impl.IrReturnImpl
import org.jetbrains.kotlin.ir.expressions.impl.IrTypeOperatorCallImpl
@@ -34,8 +33,8 @@ internal val tailCallOptimizationPhase = makeIrFilePhase(
private class TailCallOptimizationLowering(private val context: JvmBackendContext) : FileLoweringPass {
override fun lower(irFile: IrFile) {
irFile.transformChildren(object : IrElementTransformer<TailCallOptimizationData?> {
override fun visitFunction(declaration: IrFunction, data: TailCallOptimizationData?) =
super.visitFunction(declaration, TailCallOptimizationData(declaration))
override fun visitSimpleFunction(declaration: IrSimpleFunction, data: TailCallOptimizationData?) =
super.visitSimpleFunction(declaration, if (declaration.isSuspend) TailCallOptimizationData(declaration) else null)
override fun visitCall(expression: IrCall, data: TailCallOptimizationData?): IrExpression {
val transformed = super.visitCall(expression, data) as IrExpression
@@ -52,33 +51,31 @@ private class TailCallOptimizationLowering(private val context: JvmBackendContex
)
}
private class TailCallOptimizationData(val function: IrFunction) {
private class TailCallOptimizationData(val function: IrSimpleFunction) {
val returnsUnit = function.returnType.isUnit()
val tailCalls = mutableSetOf<IrCall>()
// Collect all tail calls, including those nested in `when`s, which are not arguments to `return`s.
private fun findCallsOnTailPositionWithoutImmediateReturn(statement: IrStatement, immediateReturn: Boolean = false) {
private fun IrStatement.findCallsOnTailPositionWithoutImmediateReturn(immediateReturn: Boolean = false) {
when {
statement is IrCall && statement.isSuspend && !immediateReturn && (returnsUnit || statement.type == function.returnType) ->
tailCalls += statement
statement is IrBlock ->
statement.statements.findTailCall(returnsUnit)?.let(::findCallsOnTailPositionWithoutImmediateReturn)
statement is IrWhen ->
statement.branches.forEach { findCallsOnTailPositionWithoutImmediateReturn(it.result) }
statement is IrReturn ->
findCallsOnTailPositionWithoutImmediateReturn(statement.value, immediateReturn = true)
statement is IrTypeOperatorCall && statement.operator == IrTypeOperator.IMPLICIT_COERCION_TO_UNIT ->
findCallsOnTailPositionWithoutImmediateReturn(statement.argument)
this is IrCall && isSuspend && !immediateReturn && (returnsUnit || type == function.returnType) ->
tailCalls += this
this is IrBlock ->
statements.findTailCall(returnsUnit)?.findCallsOnTailPositionWithoutImmediateReturn()
this is IrWhen ->
branches.forEach { it.result.findCallsOnTailPositionWithoutImmediateReturn() }
this is IrReturn ->
value.findCallsOnTailPositionWithoutImmediateReturn(immediateReturn = true)
this is IrTypeOperatorCall && operator == IrTypeOperator.IMPLICIT_COERCION_TO_UNIT ->
argument.findCallsOnTailPositionWithoutImmediateReturn()
// TODO: Support binary logical operations and elvis, though. KT-23826 and KT-23825
}
}
init {
if (function.hasContinuation()) {
when (val body = function.body) {
is IrBlockBody -> body.statements.findTailCall(returnsUnit)?.let(::findCallsOnTailPositionWithoutImmediateReturn)
is IrExpressionBody -> findCallsOnTailPositionWithoutImmediateReturn(body.expression)
}
when (val body = function.body) {
is IrBlockBody -> body.statements.findTailCall(returnsUnit)?.findCallsOnTailPositionWithoutImmediateReturn()
is IrExpressionBody -> body.expression.findCallsOnTailPositionWithoutImmediateReturn()
}
}
}