KT-1436 Allow break/continue in inlined lambdas

This commit is contained in:
Pavel Mikhailovskii
2022-08-11 00:12:43 +02:00
committed by teamcity
parent ba7df005a1
commit 8ba80b4b7b
52 changed files with 1851 additions and 88 deletions
@@ -10,6 +10,7 @@ import org.jetbrains.kotlin.ir.declarations.IrDeclarationParent
import org.jetbrains.kotlin.ir.declarations.IrTypeParametersContainer
import org.jetbrains.kotlin.ir.declarations.copyAttributes
import org.jetbrains.kotlin.ir.expressions.IrConstructorCall
import org.jetbrains.kotlin.ir.expressions.IrLoop
import org.jetbrains.kotlin.ir.expressions.IrTypeOperatorCall
import org.jetbrains.kotlin.ir.expressions.impl.IrTypeOperatorCallImpl
import org.jetbrains.kotlin.ir.symbols.IrClassifierSymbol
@@ -73,7 +73,7 @@ private class DirectInvokeLowering(private val context: JvmBackendContext) : Fil
return context.createIrBuilder(scope.scopeOwnerSymbol).run {
at(expression)
irBlock {
val arguments = function.explicitParameters.withIndex().map { (index, parameter) ->
val arguments = function.explicitParameters.mapIndexed { index, parameter ->
val argument = expression.getValueArgument(index)!!
IrVariableImpl(
argument.startOffset, argument.endOffset, IrDeclarationOrigin.DEFINED, IrVariableSymbolImpl(), parameter.name,
@@ -39,10 +39,15 @@ import org.jetbrains.kotlin.resolve.descriptorUtil.getSuperClassOrAny
import org.jetbrains.kotlin.types.KotlinType
import org.jetbrains.kotlin.types.typeUtil.isUnit
interface LoopResolver {
fun getLoop(expression: KtExpression): IrLoop?
}
class BodyGenerator(
val scopeOwnerSymbol: IrSymbol,
override val context: GeneratorContext
) : GeneratorWithScope {
override val context: GeneratorContext,
private val parentLoopResolver: LoopResolver?
) : GeneratorWithScope, LoopResolver {
val scopeOwner: DeclarationDescriptor get() = scopeOwnerSymbol.descriptor
@@ -196,8 +201,9 @@ class BodyGenerator(
loopTable[expression] = irLoop
}
fun getLoop(expression: KtExpression): IrLoop? =
loopTable[expression]
override fun getLoop(expression: KtExpression): IrLoop? {
return loopTable[expression] ?: parentLoopResolver?.getLoop(expression)
}
fun generatePrimaryConstructorBody(ktClassOrObject: KtPureClassOrObject, irConstructor: IrConstructor): IrBody {
val irBlockBody = context.irFactory.createBlockBody(ktClassOrObject.pureStartOffset, ktClassOrObject.pureEndOffset)
@@ -46,7 +46,7 @@ class DeclarationGenerator(override val context: GeneratorContext) : Generator {
return try {
when (ktDeclaration) {
is KtNamedFunction ->
FunctionGenerator(this).generateFunctionDeclaration(ktDeclaration)
FunctionGenerator(this).generateFunctionDeclaration(ktDeclaration, null)
is KtProperty ->
PropertyGenerator(this).generatePropertyDeclaration(ktDeclaration)
is KtClassOrObject ->
@@ -207,5 +207,5 @@ abstract class DeclarationGeneratorExtension(val declarationGenerator: Declarati
fun KotlinType.toIrType() = with(declarationGenerator) { toIrType() }
}
fun Generator.createBodyGenerator(scopeOwnerSymbol: IrSymbol) =
BodyGenerator(scopeOwnerSymbol, context)
fun Generator.createBodyGenerator(scopeOwnerSymbol: IrSymbol, parentLoopResolver: LoopResolver? = null) =
BodyGenerator(scopeOwnerSymbol, context, parentLoopResolver)
@@ -40,6 +40,7 @@ class FunctionGenerator(declarationGenerator: DeclarationGenerator) : Declaratio
@JvmOverloads
fun generateFunctionDeclaration(
ktFunction: KtNamedFunction,
parentLoopResolver: LoopResolver?,
origin: IrDeclarationOrigin = IrDeclarationOrigin.DEFINED
): IrSimpleFunction =
declareSimpleFunction(
@@ -47,19 +48,21 @@ class FunctionGenerator(declarationGenerator: DeclarationGenerator) : Declaratio
ktFunction.receiverTypeReference,
ktFunction.contextReceivers.mapNotNull { it.typeReference() },
origin,
getOrFail(BindingContext.FUNCTION, ktFunction)
getOrFail(BindingContext.FUNCTION, ktFunction),
parentLoopResolver
) {
ktFunction.bodyExpression?.let { generateFunctionBody(it) }
}
fun generateLambdaFunctionDeclaration(ktFunction: KtFunctionLiteral): IrSimpleFunction {
fun generateLambdaFunctionDeclaration(ktFunction: KtFunctionLiteral, parentLoopResolver: LoopResolver?): IrSimpleFunction {
val lambdaDescriptor = getOrFail(BindingContext.FUNCTION, ktFunction)
return declareSimpleFunction(
ktFunction,
null,
emptyList(),
IrDeclarationOrigin.LOCAL_FUNCTION_FOR_LAMBDA,
lambdaDescriptor
lambdaDescriptor,
parentLoopResolver
) {
generateLambdaBody(ktFunction, lambdaDescriptor)
}
@@ -80,11 +83,12 @@ class FunctionGenerator(declarationGenerator: DeclarationGenerator) : Declaratio
ktContextReceivers: List<KtElement>,
origin: IrDeclarationOrigin,
descriptor: FunctionDescriptor,
parentLoopResolver: LoopResolver?,
generateBody: BodyGenerator.() -> IrBody?
): IrSimpleFunction =
declareSimpleFunctionInner(descriptor, ktFunction, origin).buildWithScope { irFunction ->
generateFunctionParameterDeclarationsAndReturnType(irFunction, ktFunction, ktReceiver, ktContextReceivers)
irFunction.body = createBodyGenerator(irFunction.symbol).generateBody()
irFunction.body = createBodyGenerator(irFunction.symbol, parentLoopResolver).generateBody()
}
private fun declareSimpleFunctionInner(
@@ -16,6 +16,7 @@
package org.jetbrains.kotlin.psi2ir.generators
import org.jetbrains.kotlin.config.LanguageFeature
import org.jetbrains.kotlin.ir.IrStatement
import org.jetbrains.kotlin.ir.declarations.IrDeclarationOrigin
import org.jetbrains.kotlin.ir.expressions.IrStatementOrigin
@@ -24,13 +25,19 @@ import org.jetbrains.kotlin.psi.KtLambdaExpression
import org.jetbrains.kotlin.psi.KtNamedFunction
import org.jetbrains.kotlin.psi.psiUtil.endOffset
import org.jetbrains.kotlin.psi.psiUtil.startOffset
import org.jetbrains.kotlin.resolve.bindingContextUtil.isInlineableFunctionLiteral
class LocalFunctionGenerator(statementGenerator: StatementGenerator) : StatementGeneratorExtension(statementGenerator) {
fun generateLambda(ktLambda: KtLambdaExpression): IrStatement {
val ktFun = ktLambda.functionLiteral
val lambdaExpressionType = getTypeInferredByFrontendOrFail(ktLambda).toIrType()
val irLambdaFunction = FunctionGenerator(context).generateLambdaFunctionDeclaration(ktFun)
val loopResolver = if (context.languageVersionSettings.supportsFeature(LanguageFeature.BreakContinueInInlineLambdas)
&& isInlineableFunctionLiteral(ktLambda, context.bindingContext)
)
statementGenerator.bodyGenerator
else null
val irLambdaFunction = FunctionGenerator(context).generateLambdaFunctionDeclaration(ktFun, loopResolver)
return IrFunctionExpressionImpl(
ktLambda.startOffset, ktLambda.endOffset,
@@ -54,5 +61,12 @@ class LocalFunctionGenerator(statementGenerator: StatementGenerator) : Statement
}
private fun generateFunctionDeclaration(ktFun: KtNamedFunction) =
FunctionGenerator(context).generateFunctionDeclaration(ktFun, IrDeclarationOrigin.LOCAL_FUNCTION)
FunctionGenerator(context).generateFunctionDeclaration(
ktFun,
if (context.languageVersionSettings.supportsFeature(LanguageFeature.BreakContinueInInlineLambdas)
&& isInlineableFunctionLiteral(ktFun, context.bindingContext)
) statementGenerator.bodyGenerator
else null,
IrDeclarationOrigin.LOCAL_FUNCTION
)
}
@@ -168,7 +168,8 @@ class ScriptGenerator(declarationGenerator: DeclarationGenerator) : DeclarationG
is KtScriptInitializer -> {
val irExpressionBody = BodyGenerator(
irScript.symbol,
context
context,
null
).generateExpressionBody(d.body!!)
if (d == ktScript.declarations.last() && descriptor.resultValue != null) {
descriptor.resultValue!!.let { resultDescriptor ->
@@ -194,7 +195,7 @@ class ScriptGenerator(declarationGenerator: DeclarationGenerator) : DeclarationG
is KtDestructuringDeclaration -> {
// copied with modifications from StatementGenerator.visitDestructuringDeclaration
// TODO: consider code deduplication
val bodyGenerator = BodyGenerator(irScript.symbol, context)
val bodyGenerator = BodyGenerator(irScript.symbol, context, null)
val statementGenerator = bodyGenerator.createStatementGenerator()
val irBlock = IrCompositeImpl(
d.startOffsetSkippingComments, d.endOffset,
@@ -5,7 +5,6 @@
package org.jetbrains.kotlin.ir
import org.jetbrains.kotlin.ir.expressions.IrLoop
import org.jetbrains.kotlin.ir.util.DeepCopyIrTreeWithSymbols
import org.jetbrains.kotlin.ir.util.DeepCopySymbolRemapper
import org.jetbrains.kotlin.ir.util.DeepCopyTypeRemapper
@@ -19,12 +18,5 @@ fun <T : IrElement> T.deepCopyWithVariables(): T {
val typesRemapper = DeepCopyTypeRemapper(symbolsRemapper)
return this.transform(
object : DeepCopyIrTreeWithSymbols(symbolsRemapper, typesRemapper) {
override fun getNonTransformedLoop(irLoop: IrLoop): IrLoop {
return irLoop
}
},
null
) as T
return this.transform(DeepCopyIrTreeWithSymbols(symbolsRemapper, typesRemapper), null) as T
}
@@ -738,10 +738,7 @@ open class DeepCopyIrTreeWithSymbols(
private val transformedLoops = HashMap<IrLoop, IrLoop>()
private fun getTransformedLoop(irLoop: IrLoop): IrLoop =
transformedLoops.getOrElse(irLoop) { getNonTransformedLoop(irLoop) }
protected open fun getNonTransformedLoop(irLoop: IrLoop): IrLoop =
throw AssertionError("Outer loop was not transformed: ${irLoop.render()}")
transformedLoops.getOrDefault(irLoop, irLoop)
override fun visitWhileLoop(loop: IrWhileLoop): IrWhileLoop =
IrWhileLoopImpl(loop.startOffset, loop.endOffset, loop.type.remapType(), mapStatementOrigin(loop.origin)).also { newLoop ->