Disable tail-call optimization for suspend functions with Unit return type
if it overrides functions with another return type.
Otherwise, we cannot determine on call site that the function returns Unit
and cannot { POP, PUSH Unit } in order to avoid the situation when callee's
continuation resumes with non-unit result. The observed behavior is that
suspend function, which should return Unit, suddenly returns other value.
#KT-35262: Fixed
This commit is contained in:
@@ -475,7 +475,8 @@ class CoroutineCodegenForLambda private constructor(
|
||||
shouldPreserveClassInitialization = constructorCallNormalizationMode.shouldPreserveClassInitialization,
|
||||
containingClassInternalName = v.thisName,
|
||||
isForNamedFunction = false,
|
||||
languageVersionSettings = languageVersionSettings
|
||||
languageVersionSettings = languageVersionSettings,
|
||||
disableTailCallOptimizationForFunctionReturningUnit = false
|
||||
)
|
||||
return if (forInline) AddEndLabelMethodVisitor(
|
||||
MethodNodeCopyingMethodVisitor(
|
||||
|
||||
+28
-20
@@ -63,6 +63,10 @@ class CoroutineTransformerMethodVisitor(
|
||||
// These two are needed to report diagnostics about suspension points inside critical section
|
||||
private val element: KtElement,
|
||||
private val diagnostics: DiagnosticSink,
|
||||
// Since tail-call optimization of functions with Unit return type relies on ability of call-site to recognize them,
|
||||
// in order to ignore return value and push Unit, when we cannot ensure this ability, for example, when the function overrides function,
|
||||
// returning Any, we need to disable tail-call optimization for these functions.
|
||||
private val disableTailCallOptimizationForFunctionReturningUnit: Boolean,
|
||||
// It's only matters for named functions, may differ from '!isStatic(access)' in case of DefaultImpls
|
||||
private val needDispatchReceiver: Boolean = false,
|
||||
// May differ from containingClassInternalName in case of DefaultImpls
|
||||
@@ -111,7 +115,8 @@ class CoroutineTransformerMethodVisitor(
|
||||
val examiner = MethodNodeExaminer(
|
||||
languageVersionSettings,
|
||||
containingClassInternalName,
|
||||
methodNode
|
||||
methodNode,
|
||||
disableTailCallOptimizationForFunctionReturningUnit
|
||||
)
|
||||
if (examiner.allSuspensionPointsAreTailCalls(suspensionPoints)) {
|
||||
examiner.replacePopsBeforeSafeUnitInstancesWithCoroutineSuspendedChecks()
|
||||
@@ -899,7 +904,8 @@ class CoroutineTransformerMethodVisitor(
|
||||
private class MethodNodeExaminer(
|
||||
val languageVersionSettings: LanguageVersionSettings,
|
||||
val containingClassInternalName: String,
|
||||
val methodNode: MethodNode
|
||||
val methodNode: MethodNode,
|
||||
disableTailCallOptimizationForFunctionReturningUnit: Boolean
|
||||
) {
|
||||
private val sourceFrames: Array<Frame<SourceValue>?> =
|
||||
MethodTransformer.analyze(containingClassInternalName, methodNode, IgnoringCopyOperationSourceInterpreter())
|
||||
@@ -912,25 +918,27 @@ private class MethodNodeExaminer(
|
||||
private val meaningfulPredecessorsCache = hashMapOf<AbstractInsnNode, List<AbstractInsnNode>>()
|
||||
|
||||
init {
|
||||
// retrieve all POP insns
|
||||
val pops = methodNode.instructions.asSequence().filter { it.opcode == Opcodes.POP }
|
||||
// for each of them check that all successors are PUSH Unit
|
||||
val popsBeforeUnitInstances = pops.map { it to it.meaningfulSuccessors() }
|
||||
.filter { (_, succs) -> succs.all { it.isUnitInstance() } }
|
||||
.map { it.first }.toList()
|
||||
for (pop in popsBeforeUnitInstances) {
|
||||
val units = pop.meaningfulSuccessors()
|
||||
val allUnitsAreSafe = units.all { unit ->
|
||||
// check no other predecessor exists
|
||||
unit.meaningfulPredecessors().all { it in popsBeforeUnitInstances } &&
|
||||
// check they have only returns among successors
|
||||
unit.meaningfulSuccessors().all { it.opcode == Opcodes.ARETURN }
|
||||
if (!disableTailCallOptimizationForFunctionReturningUnit) {
|
||||
// retrieve all POP insns
|
||||
val pops = methodNode.instructions.asSequence().filter { it.opcode == Opcodes.POP }
|
||||
// for each of them check that all successors are PUSH Unit
|
||||
val popsBeforeUnitInstances = pops.map { it to it.meaningfulSuccessors() }
|
||||
.filter { (_, succs) -> succs.all { it.isUnitInstance() } }
|
||||
.map { it.first }.toList()
|
||||
for (pop in popsBeforeUnitInstances) {
|
||||
val units = pop.meaningfulSuccessors()
|
||||
val allUnitsAreSafe = units.all { unit ->
|
||||
// check no other predecessor exists
|
||||
unit.meaningfulPredecessors().all { it in popsBeforeUnitInstances } &&
|
||||
// check they have only returns among successors
|
||||
unit.meaningfulSuccessors().all { it.opcode == Opcodes.ARETURN }
|
||||
}
|
||||
if (!allUnitsAreSafe) continue
|
||||
// save them all to the properties
|
||||
popsBeforeSafeUnitInstances += pop
|
||||
safeUnitInstances += units
|
||||
units.flatMapTo(areturnsAfterSafeUnitInstances) { it.meaningfulSuccessors() }
|
||||
}
|
||||
if (!allUnitsAreSafe) continue
|
||||
// save them all to the properties
|
||||
popsBeforeSafeUnitInstances += pop
|
||||
safeUnitInstances += units
|
||||
units.flatMapTo(areturnsAfterSafeUnitInstances) { it.meaningfulSuccessors() }
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
+19
-1
@@ -19,6 +19,7 @@ import org.jetbrains.kotlin.psi.KtFunction
|
||||
import org.jetbrains.kotlin.psi.psiUtil.getElementTextWithContext
|
||||
import org.jetbrains.kotlin.resolve.jvm.diagnostics.OtherOrigin
|
||||
import org.jetbrains.kotlin.resolve.jvm.jvmSignature.JvmMethodSignature
|
||||
import org.jetbrains.kotlin.types.typeUtil.isUnit
|
||||
import org.jetbrains.kotlin.utils.addToStdlib.safeAs
|
||||
import org.jetbrains.kotlin.utils.sure
|
||||
import org.jetbrains.org.objectweb.asm.MethodVisitor
|
||||
@@ -100,10 +101,27 @@ open class SuspendFunctionGenerationStrategy(
|
||||
shouldPreserveClassInitialization = constructorCallNormalizationMode.shouldPreserveClassInitialization,
|
||||
needDispatchReceiver = originalSuspendDescriptor.dispatchReceiverParameter != null,
|
||||
internalNameForDispatchReceiver = containingClassInternalNameOrNull(),
|
||||
languageVersionSettings = languageVersionSettings
|
||||
languageVersionSettings = languageVersionSettings,
|
||||
disableTailCallOptimizationForFunctionReturningUnit = originalSuspendDescriptor.returnType?.isUnit() == true &&
|
||||
originalSuspendDescriptor.overriddenDescriptors.isNotEmpty() &&
|
||||
!originalSuspendDescriptor.allOverriddenFunctionsReturnUnit()
|
||||
)
|
||||
}
|
||||
|
||||
private fun FunctionDescriptor.allOverriddenFunctionsReturnUnit(): Boolean {
|
||||
val visited = mutableSetOf<FunctionDescriptor>()
|
||||
|
||||
fun bfs(descriptor: FunctionDescriptor): Boolean {
|
||||
if (!visited.add(descriptor)) return true
|
||||
if (descriptor.original.returnType?.isUnit() != true) return false
|
||||
for (parent in descriptor.overriddenDescriptors) {
|
||||
if (!bfs(parent)) return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
return bfs(this)
|
||||
}
|
||||
|
||||
private fun containingClassInternalNameOrNull() =
|
||||
originalSuspendDescriptor.containingDeclaration.safeAs<ClassDescriptor>()?.let(state.typeMapper::mapClass)?.internalName
|
||||
|
||||
|
||||
+17
-4
@@ -17,11 +17,20 @@ import org.jetbrains.kotlin.codegen.optimization.transformer.MethodTransformer
|
||||
import org.jetbrains.kotlin.config.isReleaseCoroutines
|
||||
import org.jetbrains.kotlin.descriptors.ClassDescriptor
|
||||
import org.jetbrains.kotlin.descriptors.DeclarationDescriptor
|
||||
import org.jetbrains.kotlin.load.java.JvmAnnotationNames
|
||||
import org.jetbrains.kotlin.load.kotlin.FileBasedKotlinClass
|
||||
import org.jetbrains.kotlin.load.kotlin.header.KotlinClassHeader
|
||||
import org.jetbrains.kotlin.load.kotlin.header.ReadKotlinClassHeaderAnnotationVisitor
|
||||
import org.jetbrains.kotlin.metadata.ProtoBuf
|
||||
import org.jetbrains.kotlin.metadata.deserialization.*
|
||||
import org.jetbrains.kotlin.metadata.jvm.deserialization.JvmProtoBufUtil
|
||||
import org.jetbrains.kotlin.name.FqName
|
||||
import org.jetbrains.kotlin.psi.KtElement
|
||||
import org.jetbrains.kotlin.resolve.jvm.diagnostics.JvmDeclarationOrigin
|
||||
import org.jetbrains.kotlin.serialization.deserialization.getClassId
|
||||
import org.jetbrains.kotlin.serialization.deserialization.getName
|
||||
import org.jetbrains.kotlin.utils.addToStdlib.cast
|
||||
import org.jetbrains.org.objectweb.asm.MethodVisitor
|
||||
import org.jetbrains.org.objectweb.asm.Opcodes
|
||||
import org.jetbrains.org.objectweb.asm.*
|
||||
import org.jetbrains.org.objectweb.asm.tree.*
|
||||
import org.jetbrains.org.objectweb.asm.tree.analysis.Frame
|
||||
import org.jetbrains.org.objectweb.asm.tree.analysis.SourceInterpreter
|
||||
@@ -108,7 +117,8 @@ class CoroutineTransformer(
|
||||
languageVersionSettings = state.languageVersionSettings,
|
||||
shouldPreserveClassInitialization = state.constructorCallNormalizationMode.shouldPreserveClassInitialization,
|
||||
containingClassInternalName = classBuilder.thisName,
|
||||
isForNamedFunction = false
|
||||
isForNamedFunction = false,
|
||||
disableTailCallOptimizationForFunctionReturningUnit = false
|
||||
)
|
||||
)
|
||||
|
||||
@@ -137,6 +147,8 @@ class CoroutineTransformer(
|
||||
ArrayUtil.toStringArray(node.exceptions)
|
||||
)
|
||||
) {
|
||||
// If the node already has state-machine, it is safer to generate state-machine.
|
||||
val disableTailCallOptimization = methods.find { it.name == name && it.desc == node.desc }?.let { isStateMachine(it) } ?: false
|
||||
val stateMachineBuilder = surroundNoinlineCallsWithMarkers(
|
||||
node,
|
||||
CoroutineTransformerMethodVisitor(
|
||||
@@ -149,7 +161,8 @@ class CoroutineTransformer(
|
||||
containingClassInternalName = classBuilder.thisName,
|
||||
isForNamedFunction = true,
|
||||
needDispatchReceiver = true,
|
||||
internalNameForDispatchReceiver = classBuilder.thisName
|
||||
internalNameForDispatchReceiver = classBuilder.thisName,
|
||||
disableTailCallOptimizationForFunctionReturningUnit = disableTailCallOptimization
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user