Simplify and optimize JvmLateinitLowering

This commit is contained in:
mcpiroman
2023-03-04 14:18:16 +01:00
committed by Alexander Udalov
parent 8315eeaf92
commit d87468ef39
@@ -43,18 +43,24 @@ class JvmLateinitLowering(
) : FileLoweringPass { ) : FileLoweringPass {
override fun lower(irFile: IrFile) { override fun lower(irFile: IrFile) {
irFile.transformChildrenVoid(Transformer(context)) val transformer = Transformer(context)
irFile.transformChildrenVoid(transformer)
for (variable in transformer.lateinitVariables) {
variable.isLateinit = false
}
} }
private class Transformer(private val backendContext: JvmBackendContext) : IrElementTransformerVoid() { private class Transformer(private val backendContext: JvmBackendContext) : IrElementTransformerVoid() {
private val backingVariables = HashMap<IrVariable, IrVariable>() val lateinitVariables = mutableListOf<IrVariable>()
override fun visitField(declaration: IrField): IrStatement { override fun visitField(declaration: IrField): IrStatement {
if (declaration.isLateinitBackingField()) { if (declaration.isLateinitBackingField()) {
assert(declaration.initializer == null) { assert(declaration.initializer == null) {
"lateinit property backing field should not have an initializer:\n${declaration.dump()}" "lateinit property backing field should not have an initializer:\n${declaration.dump()}"
} }
return getOrBuildLateinitBackingField(declaration)
declaration.type = declaration.type.makeNullable()
} }
declaration.transformChildrenVoid() declaration.transformChildrenVoid()
@@ -63,30 +69,23 @@ class JvmLateinitLowering(
override fun visitVariable(declaration: IrVariable): IrStatement { override fun visitVariable(declaration: IrVariable): IrStatement {
declaration.transformChildrenVoid(this) declaration.transformChildrenVoid(this)
if (!declaration.isLateinit) return declaration
return getOrBuildLateinitBackingVar(declaration)
}
private fun getOrBuildLateinitBackingVar(declaration: IrVariable): IrVariable = if (declaration.isLateinit) {
backingVariables.getOrPut(declaration) { declaration.type = declaration.type.makeNullable()
buildVariable( declaration.isVar = true
declaration.parent, declaration.initializer =
declaration.startOffset, IrConstImpl.constNull(declaration.startOffset, declaration.endOffset, backendContext.irBuiltIns.nothingNType)
declaration.endOffset,
declaration.origin, lateinitVariables += declaration
declaration.name,
declaration.type.makeNullable(),
isVar = true,
).also {
it.initializer =
IrConstImpl.constNull(declaration.startOffset, declaration.endOffset, backendContext.irBuiltIns.nothingNType)
}
} }
return declaration
}
override fun visitSimpleFunction(declaration: IrSimpleFunction): IrStatement { override fun visitSimpleFunction(declaration: IrSimpleFunction): IrStatement {
val property = declaration.correspondingPropertySymbol?.owner val property = declaration.correspondingPropertySymbol?.owner
if (property != null && property.isRealLateinit() && declaration == property.getter) { if (property != null && property.isRealLateinit() && declaration == property.getter) {
transformGetter(getOrBuildLateinitBackingField(property.backingField!!), declaration) transformGetter(property.backingField!!, declaration)
return declaration return declaration
} }
@@ -99,59 +98,28 @@ class JvmLateinitLowering(
if (irValue !is IrVariable || !irValue.isLateinit) { if (irValue !is IrVariable || !irValue.isLateinit) {
return expression return expression
} }
val irBackingVar = backingVariables[irValue]
?: throw AssertionError("Lateinit variable reference before use: ${expression.dump()}")
return backendContext.createIrBuilder( return backendContext.createIrBuilder(
(irBackingVar.parent as IrSymbolOwner).symbol, (irValue.parent as IrSymbolOwner).symbol,
expression.startOffset, expression.startOffset,
expression.endOffset expression.endOffset
).run { ).run {
irIfThenElse( irIfThenElse(
expression.type, expression.type,
irEqualsNull(irGet(irBackingVar)), irEqualsNull(irGet(irValue)),
backendContext.throwUninitializedPropertyAccessException(this, irBackingVar.name.asString()), backendContext.throwUninitializedPropertyAccessException(this, irValue.name.asString()),
irGet(irBackingVar) irGet(irValue)
) )
} }
} }
override fun visitSetValue(expression: IrSetValue): IrExpression {
expression.transformChildrenVoid(this)
val irValue = expression.symbol.owner
if (irValue !is IrVariable || !irValue.isLateinit) {
return expression
}
val irBackingVar = backingVariables[expression.symbol.owner]
?: throw AssertionError("Lateinit variable reference before use: ${expression.dump()}")
return with(expression) {
IrSetValueImpl(startOffset, endOffset, type, irBackingVar.symbol, value, origin)
}
}
override fun visitGetField(expression: IrGetField): IrExpression { override fun visitGetField(expression: IrGetField): IrExpression {
expression.transformChildrenVoid(this) expression.transformChildrenVoid(this)
val irField = expression.symbol.owner val irField = expression.symbol.owner
if (!irField.isLateinitBackingField()) { if (irField.isLateinitBackingField()) {
return expression expression.type = expression.type.makeNullable()
}
val newField = getOrBuildLateinitBackingField(irField)
return with(expression) {
IrGetFieldImpl(startOffset, endOffset, newField.symbol, newField.type, receiver, origin, superQualifierSymbol)
}
}
override fun visitSetField(expression: IrSetField): IrExpression {
expression.transformChildrenVoid(this)
val irField = expression.symbol.owner
if (!irField.isLateinitBackingField()) {
return expression
}
val newField = getOrBuildLateinitBackingField(irField)
return with(expression) {
IrSetFieldImpl(startOffset, endOffset, newField.symbol, receiver, value, type, origin, superQualifierSymbol)
} }
return expression
} }
private fun IrField.isLateinitBackingField(): Boolean { private fun IrField.isLateinitBackingField(): Boolean {
@@ -177,10 +145,9 @@ class JvmLateinitLowering(
} }
val backingField = property.backingField val backingField = property.backingField
?: throw AssertionError("Lateinit property is supposed to have a backing field") ?: throw AssertionError("Lateinit property is supposed to have a backing field")
val newField = getOrBuildLateinitBackingField(backingField)
backendContext.createIrBuilder(it.symbol, expression.startOffset, expression.endOffset).run { backendContext.createIrBuilder(it.symbol, expression.startOffset, expression.endOffset).run {
irNotEquals( irNotEquals(
irGetField(it.dispatchReceiver, newField), irGetField(it.dispatchReceiver, backingField),
irNull() irNull()
) )
} }
@@ -198,7 +165,7 @@ class JvmLateinitLowering(
val irBuilder = backendContext.createIrBuilder(getter.symbol, startOffset, endOffset) val irBuilder = backendContext.createIrBuilder(getter.symbol, startOffset, endOffset)
irBuilder.run { irBuilder.run {
val resultVar = scope.createTmpVariable( val resultVar = scope.createTmpVariable(
irGetField(getter.dispatchReceiverParameter?.let { irGet(it) }, backingField) irGetField(getter.dispatchReceiverParameter?.let { irGet(it) }, backingField, backingField.type.makeNullable())
) )
resultVar.parent = getter resultVar.parent = getter
statements.add(resultVar) statements.add(resultVar)
@@ -215,22 +182,5 @@ class JvmLateinitLowering(
private fun IrBuilderWithScope.throwUninitializedPropertyAccessException(name: String) = private fun IrBuilderWithScope.throwUninitializedPropertyAccessException(name: String) =
backendContext.throwUninitializedPropertyAccessException(this, name) backendContext.throwUninitializedPropertyAccessException(this, name)
private fun getOrBuildLateinitBackingField(originalField: IrField): IrField =
if (originalField.type.isMarkedNullable())
originalField
else
backendContext.mapping.lateInitFieldToNullableField.getOrPut(originalField) {
backendContext.irFactory.buildField {
updateFrom(originalField)
type = originalField.type.makeNullable()
name = originalField.name
}.apply {
parent = originalField.parent
correspondingPropertySymbol = originalField.correspondingPropertySymbol
annotations = originalField.annotations
}
}
} }
} }