JVM_IR get rid of intermediate patch parent phases

This commit is contained in:
Dmitry Petrov
2022-01-26 17:29:25 +03:00
committed by Space
parent 6c6534ee20
commit 08a946fd3f
13 changed files with 179 additions and 41 deletions
@@ -35,6 +35,21 @@ abstract class IrElementTransformerVoidWithContext : IrElementTransformerVoid()
protected open fun createScope(declaration: IrSymbolOwner): ScopeWithIr = protected open fun createScope(declaration: IrSymbolOwner): ScopeWithIr =
ScopeWithIr(Scope(declaration.symbol), declaration) ScopeWithIr(Scope(declaration.symbol), declaration)
protected fun unsafeEnterScope(declaration: IrSymbolOwner) {
scopeStack.push(createScope(declaration))
}
protected fun unsafeLeaveScope() {
scopeStack.pop()
}
protected inline fun <T> withinScope(declaration: IrSymbolOwner, fn: () -> T): T {
unsafeEnterScope(declaration)
val result = fn()
unsafeLeaveScope()
return result
}
final override fun visitFile(declaration: IrFile): IrFile { final override fun visitFile(declaration: IrFile): IrFile {
scopeStack.push(createScope(declaration)) scopeStack.push(createScope(declaration))
val result = visitFileNew(declaration) val result = visitFileNew(declaration)
@@ -51,13 +51,17 @@ open class TailrecLowering(val context: BackendContext) : BodyLoweringPass {
element.acceptChildrenVoid(this) element.acceptChildrenVoid(this)
} }
override fun visitFunction(declaration: IrFunction) { override fun visitSimpleFunction(declaration: IrSimpleFunction) {
declaration.acceptChildrenVoid(this) declaration.acceptChildrenVoid(this)
lowerTailRecursionCalls(declaration) if (declaration.isTailrec) {
lowerTailRecursionCalls(declaration)
}
} }
}) })
lowerTailRecursionCalls(container) if (container is IrSimpleFunction && container.isTailrec) {
lowerTailRecursionCalls(container)
}
} }
} }
@@ -110,6 +114,9 @@ private fun TailrecLowering.lowerTailRecursionCalls(irFunction: IrFunction) {
} }
} }
}.statements }.statements
// TODO BodyTransformer creates temporary variables with wrong parents in nested functions
oldBody.patchDeclarationParents(irFunction)
} }
private class BodyTransformer( private class BodyTransformer(
@@ -19,20 +19,34 @@ import org.jetbrains.kotlin.descriptors.DescriptorVisibility
import org.jetbrains.kotlin.ir.IrElement import org.jetbrains.kotlin.ir.IrElement
import org.jetbrains.kotlin.ir.declarations.* import org.jetbrains.kotlin.ir.declarations.*
import org.jetbrains.kotlin.ir.symbols.IrValueSymbol import org.jetbrains.kotlin.ir.symbols.IrValueSymbol
import org.jetbrains.kotlin.ir.util.PatchDeclarationParentsVisitor import org.jetbrains.kotlin.ir.util.*
import org.jetbrains.kotlin.ir.util.isAnonymousObject
import org.jetbrains.kotlin.ir.util.parentAsClass
import org.jetbrains.kotlin.ir.visitors.IrElementVisitorVoid 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.ir.visitors.acceptVoid
import org.jetbrains.kotlin.load.java.JavaDescriptorVisibilities import org.jetbrains.kotlin.load.java.JavaDescriptorVisibilities
import org.jetbrains.kotlin.name.NameUtils import org.jetbrains.kotlin.name.NameUtils
private fun makePatchParentsPhase(number: Int): NamedCompilerPhase<CommonBackendContext, IrFile> = makeIrFilePhase( private var patchParentPhases = 0
{ PatchDeclarationParents() },
name = "PatchParents$number", private fun makePatchParentsPhase(): NamedCompilerPhase<CommonBackendContext, IrFile> {
description = "Patch parent references in IrFile, pass $number", val number = patchParentPhases++
) return makeIrFilePhase(
{ PatchDeclarationParents() },
name = "PatchParents$number",
description = "Patch parent references in IrFile, pass $number",
)
}
private var checkParentPhases = 0
private fun makeCheckParentsPhase(): NamedCompilerPhase<CommonBackendContext, IrFile> {
val number = checkParentPhases++
return makeIrFilePhase(
{ CheckDeclarationParents() },
name = "CheckParents$number",
description = "Check parent references in IrFile, pass $number",
)
}
private class PatchDeclarationParents : FileLoweringPass { private class PatchDeclarationParents : FileLoweringPass {
override fun lower(irFile: IrFile) { override fun lower(irFile: IrFile) {
@@ -40,6 +54,39 @@ private class PatchDeclarationParents : FileLoweringPass {
} }
} }
private class CheckDeclarationParents : FileLoweringPass {
override fun lower(irFile: IrFile) {
val checker = CheckDeclarationParentsVisitor()
irFile.acceptVoid(checker)
if (checker.errors.isNotEmpty()) {
val expectedParents = LinkedHashSet<IrDeclarationParent>()
throw AssertionError(
buildString {
append("Declarations with wrong parent: ")
append(checker.errors.size)
append("\n")
checker.errors.forEach {
append("declaration: ")
append(it.declaration.render())
append("\n\t")
append(it.declaration)
append("\nexpectedParent: ")
append(it.expectedParent.render())
append("\nactualParent: ")
append(it.actualParent?.render())
append("\n")
expectedParents.add(it.expectedParent)
}
append("\nExpected parents:\n")
expectedParents.forEach {
append(it.dump())
}
}
)
}
}
}
private val validateIrBeforeLowering = makeCustomPhase( private val validateIrBeforeLowering = makeCustomPhase(
::validateIr, ::validateIr,
name = "ValidateIrBeforeLowering", name = "ValidateIrBeforeLowering",
@@ -291,7 +338,7 @@ private val jvmFilePhases = listOf(
singleAbstractMethodPhase, singleAbstractMethodPhase,
jvmInlineClassPhase, jvmInlineClassPhase,
tailrecPhase, tailrecPhase,
makePatchParentsPhase(1), // makePatchParentsPhase(),
enumWhenPhase, enumWhenPhase,
singletonReferencesPhase, singletonReferencesPhase,
@@ -300,7 +347,7 @@ private val jvmFilePhases = listOf(
returnableBlocksPhase, returnableBlocksPhase,
sharedVariablesPhase, sharedVariablesPhase,
localDeclarationsPhase, localDeclarationsPhase,
makePatchParentsPhase(2), // makePatchParentsPhase(),
jvmLocalClassExtractionPhase, jvmLocalClassExtractionPhase,
staticCallableReferencePhase, staticCallableReferencePhase,
@@ -316,6 +363,7 @@ private val jvmFilePhases = listOf(
defaultArgumentInjectorPhase, defaultArgumentInjectorPhase,
defaultArgumentCleanerPhase, defaultArgumentCleanerPhase,
// makePatchParentsPhase(),
interfacePhase, interfacePhase,
inheritedDefaultMethodsOnClassesPhase, inheritedDefaultMethodsOnClassesPhase,
replaceDefaultImplsOverriddenSymbolsPhase, replaceDefaultImplsOverriddenSymbolsPhase,
@@ -331,7 +379,7 @@ private val jvmFilePhases = listOf(
innerClassesMemberBodyPhase, innerClassesMemberBodyPhase,
innerClassConstructorCallsPhase, innerClassConstructorCallsPhase,
makePatchParentsPhase(3), // makePatchParentsPhase(),
enumClassPhase, enumClassPhase,
objectClassPhase, objectClassPhase,
@@ -357,7 +405,7 @@ private val jvmFilePhases = listOf(
renameFieldsPhase, renameFieldsPhase,
fakeInliningLocalVariablesLowering, fakeInliningLocalVariablesLowering,
makePatchParentsPhase(4) // makePatchParentsPhase()
) )
val jvmLoweringPhases = NamedCompilerPhase( val jvmLoweringPhases = NamedCompilerPhase(
@@ -22,6 +22,7 @@ import org.jetbrains.kotlin.ir.expressions.*
import org.jetbrains.kotlin.ir.expressions.impl.IrConstructorCallImpl import org.jetbrains.kotlin.ir.expressions.impl.IrConstructorCallImpl
import org.jetbrains.kotlin.ir.expressions.impl.IrGetValueImpl import org.jetbrains.kotlin.ir.expressions.impl.IrGetValueImpl
import org.jetbrains.kotlin.ir.expressions.impl.IrTypeOperatorCallImpl import org.jetbrains.kotlin.ir.expressions.impl.IrTypeOperatorCallImpl
import org.jetbrains.kotlin.ir.util.patchDeclarationParents
import org.jetbrains.kotlin.ir.util.transformInPlace import org.jetbrains.kotlin.ir.util.transformInPlace
internal val anonymousObjectSuperConstructorPhase = makeIrFilePhase( internal val anonymousObjectSuperConstructorPhase = makeIrFilePhase(
@@ -124,8 +125,10 @@ private class AnonymousObjectSuperConstructorLowering(val context: JvmBackendCon
// Avoid complex expressions between `new` and `<init>`, as the inliner gets confused if // Avoid complex expressions between `new` and `<init>`, as the inliner gets confused if
// an argument to `<init>` is an anonymous object. Put them in variables instead. // an argument to `<init>` is an anonymous object. Put them in variables instead.
// See KT-21781 for an example; in short, it looks like `object : S({ ... })` in an inline function. // See KT-21781 for an example; in short, it looks like `object : S({ ... })` in an inline function.
for ((i, argument) in newArguments.withIndex()) for ((i, argument) in newArguments.withIndex()) {
argument.patchDeclarationParents(currentDeclarationParent)
putValueArgument(i + objectConstructorCall.valueArgumentsCount, irGet(irTemporary(argument))) putValueArgument(i + objectConstructorCall.valueArgumentsCount, irGet(irTemporary(argument)))
}
} }
} }
} }
@@ -17,10 +17,7 @@ import org.jetbrains.kotlin.config.JVMAssertionsMode
import org.jetbrains.kotlin.ir.IrElement import org.jetbrains.kotlin.ir.IrElement
import org.jetbrains.kotlin.ir.IrStatement import org.jetbrains.kotlin.ir.IrStatement
import org.jetbrains.kotlin.ir.builders.* import org.jetbrains.kotlin.ir.builders.*
import org.jetbrains.kotlin.ir.declarations.IrClass import org.jetbrains.kotlin.ir.declarations.*
import org.jetbrains.kotlin.ir.declarations.IrField
import org.jetbrains.kotlin.ir.declarations.IrFile
import org.jetbrains.kotlin.ir.declarations.IrFunction
import org.jetbrains.kotlin.ir.expressions.IrCall import org.jetbrains.kotlin.ir.expressions.IrCall
import org.jetbrains.kotlin.ir.expressions.IrExpression import org.jetbrains.kotlin.ir.expressions.IrExpression
import org.jetbrains.kotlin.ir.expressions.impl.IrCompositeImpl import org.jetbrains.kotlin.ir.expressions.impl.IrCompositeImpl
@@ -45,16 +42,25 @@ private class AssertionLowering(private val context: JvmBackendContext) :
// assertions when compiled with -Xassertions=jvm. // assertions when compiled with -Xassertions=jvm.
class ClassInfo(val irClass: IrClass, val topLevelClass: IrClass, var assertionsDisabledField: IrField? = null) class ClassInfo(val irClass: IrClass, val topLevelClass: IrClass, var assertionsDisabledField: IrField? = null)
private val scopeOwnerStack = java.util.ArrayDeque<IrDeclaration>()
override fun lower(irFile: IrFile) { override fun lower(irFile: IrFile) {
// In legacy mode we treat assertions as inline function calls // In legacy mode we treat assertions as inline function calls
if (context.state.assertionsMode != JVMAssertionsMode.LEGACY) if (context.state.assertionsMode != JVMAssertionsMode.LEGACY)
irFile.transformChildren(this, null) irFile.transformChildren(this, null)
} }
override fun visitDeclaration(declaration: IrDeclarationBase, data: ClassInfo?): IrStatement {
scopeOwnerStack.push(declaration)
val result = super.visitDeclaration(declaration, data)
scopeOwnerStack.pop()
return result
}
override fun visitClass(declaration: IrClass, data: ClassInfo?): IrStatement { override fun visitClass(declaration: IrClass, data: ClassInfo?): IrStatement {
val info = ClassInfo(declaration, data?.topLevelClass ?: declaration) val info = ClassInfo(declaration, data?.topLevelClass ?: declaration)
super.visitClass(declaration, info) visitDeclaration(declaration, info)
// Note that it's necessary to add this member at the beginning of the class, before all user-visible // Note that it's necessary to add this member at the beginning of the class, before all user-visible
// initializers, which may contain assertions. At the same time, assertions are supposed to be enabled // initializers, which may contain assertions. At the same time, assertions are supposed to be enabled
@@ -77,7 +83,7 @@ private class AssertionLowering(private val context: JvmBackendContext) :
if (mode == JVMAssertionsMode.ALWAYS_DISABLE) if (mode == JVMAssertionsMode.ALWAYS_DISABLE)
return IrCompositeImpl(expression.startOffset, expression.endOffset, context.irBuiltIns.unitType) return IrCompositeImpl(expression.startOffset, expression.endOffset, context.irBuiltIns.unitType)
context.createIrBuilder(expression.symbol).run { context.createIrBuilder(scopeOwnerStack.peek().symbol).run {
at(expression) at(expression)
val assertCondition = expression.getValueArgument(0)!! val assertCondition = expression.getValueArgument(0)!!
val lambdaArgument = if (function.valueParameters.size == 2) expression.getValueArgument(1) else null val lambdaArgument = if (function.valueParameters.size == 2) expression.getValueArgument(1) else null
@@ -1021,7 +1021,7 @@ internal class FunctionReferenceLowering(private val context: JvmBackendContext)
inlinedAdapterBlock.statements.add(inlinedAdapterResult) inlinedAdapterBlock.statements.add(inlinedAdapterResult)
callee.body = null callee.body = null
return inlinedAdapterBlock return inlinedAdapterBlock.patchDeclarationParents(invokeMethod)
} }
private fun buildOverride(superFunction: IrSimpleFunction, newReturnType: IrType = superFunction.returnType): IrSimpleFunction = private fun buildOverride(superFunction: IrSimpleFunction, newReturnType: IrType = superFunction.returnType): IrSimpleFunction =
@@ -184,7 +184,6 @@ internal class InterfaceLowering(val context: JvmBackendContext) : IrElementTran
irClass.declarations.remove(field) irClass.declarations.remove(field)
defaultImplsIrClass.declarations.add(0, field) defaultImplsIrClass.declarations.add(0, field)
field.parent = defaultImplsIrClass field.parent = defaultImplsIrClass
field.initializer?.patchDeclarationParents(defaultImplsIrClass)
} }
} }
@@ -78,7 +78,9 @@ private class JvmInlineClassLowering(private val context: JvmBackendContext) : F
declaration.transformDeclarationsFlat { memberDeclaration -> declaration.transformDeclarationsFlat { memberDeclaration ->
if (memberDeclaration is IrFunction) { if (memberDeclaration is IrFunction) {
transformFunctionFlat(memberDeclaration) withinScope(memberDeclaration) {
transformFunctionFlat(memberDeclaration)
}
} else { } else {
memberDeclaration.accept(this, null) memberDeclaration.accept(this, null)
null null
@@ -134,7 +136,10 @@ private class JvmInlineClassLowering(private val context: JvmBackendContext) : F
} }
private fun transformSimpleFunctionFlat(function: IrSimpleFunction, replacement: IrSimpleFunction): List<IrDeclaration> { private fun transformSimpleFunctionFlat(function: IrSimpleFunction, replacement: IrSimpleFunction): List<IrDeclaration> {
replacement.valueParameters.forEach { it.transformChildrenVoid() } replacement.valueParameters.forEach {
it.transformChildrenVoid()
it.defaultValue?.patchDeclarationParents(replacement)
}
allScopes.push(createScope(function)) allScopes.push(createScope(function))
replacement.body = function.body?.transform(this, null)?.patchDeclarationParents(replacement) replacement.body = function.body?.transform(this, null)?.patchDeclarationParents(replacement)
allScopes.pop() allScopes.pop()
@@ -78,17 +78,17 @@ private class MainMethodGenerationLowering(private val context: JvmBackendContex
irClass.functions.find { it.isMainMethod() }?.let { mainMethod -> irClass.functions.find { it.isMainMethod() }?.let { mainMethod ->
if (mainMethod.isSuspend) { if (mainMethod.isSuspend) {
irClass.generateMainMethod { args -> irClass.generateMainMethod { newMain, args ->
+irRunSuspend(mainMethod, args) +irRunSuspend(mainMethod, args, newMain)
} }
} }
return return
} }
irClass.functions.find { it.isParameterlessMainMethod() }?.let { parameterlessMainMethod -> irClass.functions.find { it.isParameterlessMainMethod() }?.let { parameterlessMainMethod ->
irClass.generateMainMethod { irClass.generateMainMethod { newMain, _ ->
if (parameterlessMainMethod.isSuspend) { if (parameterlessMainMethod.isSuspend) {
+irRunSuspend(parameterlessMainMethod, null) +irRunSuspend(parameterlessMainMethod, null, newMain)
} else { } else {
+irCall(parameterlessMainMethod) +irCall(parameterlessMainMethod)
} }
@@ -104,7 +104,7 @@ private class MainMethodGenerationLowering(private val context: JvmBackendContex
name.asString() == "main" name.asString() == "main"
private fun IrSimpleFunction.isMainMethod(): Boolean { private fun IrSimpleFunction.isMainMethod(): Boolean {
if (getJvmNameFromAnnotation() ?: name.asString() != "main") return false if ((getJvmNameFromAnnotation() ?: name.asString()) != "main") return false
if (!returnType.isUnit()) return false if (!returnType.isUnit()) return false
val parameter = allParameters.singleOrNull() ?: return false val parameter = allParameters.singleOrNull() ?: return false
@@ -119,7 +119,7 @@ private class MainMethodGenerationLowering(private val context: JvmBackendContex
} }
} }
private fun IrClass.generateMainMethod(makeBody: IrBlockBodyBuilder.(IrValueParameter) -> Unit) = private fun IrClass.generateMainMethod(makeBody: IrBlockBodyBuilder.(IrSimpleFunction, IrValueParameter) -> Unit) =
addFunction { addFunction {
name = Name.identifier("main") name = Name.identifier("main")
visibility = DescriptorVisibilities.PUBLIC visibility = DescriptorVisibilities.PUBLIC
@@ -131,10 +131,14 @@ private class MainMethodGenerationLowering(private val context: JvmBackendContex
name = Name.identifier("args") name = Name.identifier("args")
type = context.irBuiltIns.arrayClass.typeWith(context.irBuiltIns.stringType) type = context.irBuiltIns.arrayClass.typeWith(context.irBuiltIns.stringType)
} }
body = context.createIrBuilder(symbol).irBlockBody { makeBody(args) } body = context.createIrBuilder(symbol).irBlockBody { makeBody(this@apply, args) }
} }
private fun IrBuilderWithScope.irRunSuspend(target: IrSimpleFunction, args: IrValueParameter?): IrExpression { private fun IrBuilderWithScope.irRunSuspend(
target: IrSimpleFunction,
args: IrValueParameter?,
newMain: IrSimpleFunction
): IrExpression {
val backendContext = this@MainMethodGenerationLowering.context val backendContext = this@MainMethodGenerationLowering.context
return irBlock { return irBlock {
val wrapperConstructor = backendContext.irFactory.buildClass { val wrapperConstructor = backendContext.irFactory.buildClass {
@@ -152,7 +156,7 @@ private class MainMethodGenerationLowering(private val context: JvmBackendContex
wrapper.superTypes += lambdaSuperClass.defaultType wrapper.superTypes += lambdaSuperClass.defaultType
wrapper.superTypes += functionClass.typeWith(backendContext.irBuiltIns.anyNType) wrapper.superTypes += functionClass.typeWith(backendContext.irBuiltIns.anyNType)
wrapper.parent = target.parent wrapper.parent = newMain
val stringArrayType = backendContext.irBuiltIns.arrayClass.typeWith(backendContext.irBuiltIns.stringType) val stringArrayType = backendContext.irBuiltIns.arrayClass.typeWith(backendContext.irBuiltIns.stringType)
val argsField = args?.let { val argsField = args?.let {
@@ -127,7 +127,6 @@ private class MappedEnumWhenLowering(context: JvmBackendContext) : EnumWhenLower
super.visitClassNew(declaration) super.visitClassNew(declaration)
for ((enum, mapping) in mappingState.mappings) { for ((enum, mapping) in mappingState.mappings) {
val builder = context.createIrBuilder(mappingState.mappingsClass.symbol)
val enumValues = enum.functions.single { val enumValues = enum.functions.single {
it.name.toString() == "values" it.name.toString() == "values"
&& it.dispatchReceiverParameter == null && it.dispatchReceiverParameter == null
@@ -136,6 +135,7 @@ private class MappedEnumWhenLowering(context: JvmBackendContext) : EnumWhenLower
&& it.returnType.isBoxedArray && it.returnType.isBoxedArray
&& it.returnType.getArrayElementType(context.irBuiltIns).classOrNull == enum.symbol && it.returnType.getArrayElementType(context.irBuiltIns).classOrNull == enum.symbol
} }
val builder = context.createIrBuilder(mapping.field.symbol)
mapping.field.initializer = builder.irExprBody(builder.irBlock { mapping.field.initializer = builder.irExprBody(builder.irBlock {
val enumSize = irCall(refArraySize).apply { dispatchReceiver = irCall(enumValues) } val enumSize = irCall(refArraySize).apply { dispatchReceiver = irCall(enumValues) }
val result = irTemporary(irCall(intArrayConstructor).apply { putValueArgument(0, enumSize) }) val result = irTemporary(irCall(intArrayConstructor).apply { putValueArgument(0, enumSize) })
@@ -238,7 +238,7 @@ internal class PropertyReferenceLowering(val context: JvmBackendContext) : IrEle
if (data.kProperties.isNotEmpty()) { if (data.kProperties.isNotEmpty()) {
declaration.declarations.add(0, data.kPropertiesField.apply { declaration.declarations.add(0, data.kPropertiesField.apply {
parent = declaration parent = declaration
initializer = context.createJvmIrBuilder(declaration.symbol).run { initializer = context.createJvmIrBuilder(data.kPropertiesField.symbol).run {
val initializers = data.kProperties.values.sortedBy { it.index }.map { it.initializer } val initializers = data.kProperties.values.sortedBy { it.index }.map { it.initializer }
irExprBody(irArrayOf(kPropertiesFieldType, initializers)) irExprBody(irArrayOf(kPropertiesFieldType, initializers))
} }
@@ -76,6 +76,8 @@ private class ScriptsToClassesLowering(val context: JvmBackendContext, val inner
for ((irScript, irScriptClass) in scriptsToClasses) { for ((irScript, irScriptClass) in scriptsToClasses) {
finalizeScriptClass(irScriptClass, irScript, symbolRemapper) finalizeScriptClass(irScriptClass, irScript, symbolRemapper)
// TODO fix parents in script classes
irScriptClass.patchDeclarationParents(irScript.parent)
} }
} }
@@ -6,14 +6,12 @@
package org.jetbrains.kotlin.ir.util package org.jetbrains.kotlin.ir.util
import org.jetbrains.kotlin.ir.IrElement import org.jetbrains.kotlin.ir.IrElement
import org.jetbrains.kotlin.ir.declarations.IrDeclaration import org.jetbrains.kotlin.ir.declarations.*
import org.jetbrains.kotlin.ir.declarations.IrDeclarationParent
import org.jetbrains.kotlin.ir.declarations.IrPackageFragment
import org.jetbrains.kotlin.ir.declarations.IrDeclarationBase
import org.jetbrains.kotlin.ir.visitors.IrElementVisitorVoid 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.ir.visitors.acceptVoid
import java.util.* import java.util.*
import kotlin.collections.ArrayList
fun <T : IrElement> T.patchDeclarationParents(initialParent: IrDeclarationParent? = null) = fun <T : IrElement> T.patchDeclarationParents(initialParent: IrDeclarationParent? = null) =
apply { apply {
@@ -57,3 +55,54 @@ class PatchDeclarationParentsVisitor() : IrElementVisitorVoid {
declaration.parent = declarationParentsStack.peekFirst() declaration.parent = declarationParentsStack.peekFirst()
} }
} }
class CheckDeclarationParentsVisitor() : IrElementVisitorVoid {
constructor(containingDeclaration: IrDeclarationParent) : this() {
declarationParentsStack.push(containingDeclaration)
}
private val declarationParentsStack = ArrayDeque<IrDeclarationParent>()
class Data(val declaration: IrDeclaration, val expectedParent: IrDeclarationParent, val actualParent: IrDeclarationParent?)
val errors = ArrayList<Data>()
override fun visitElement(element: IrElement) {
element.acceptChildrenVoid(this)
}
override fun visitPackageFragment(declaration: IrPackageFragment) {
declarationParentsStack.push(declaration)
super.visitPackageFragment(declaration)
declarationParentsStack.pop()
}
override fun visitDeclaration(declaration: IrDeclarationBase) {
checkParent(declaration)
if (declaration is IrDeclarationParent) {
declarationParentsStack.push(declaration)
}
super.visitDeclaration(declaration)
if (declaration is IrDeclarationParent) {
declarationParentsStack.pop()
}
}
private fun checkParent(declaration: IrDeclaration) {
val expectedParent = declarationParentsStack.peekFirst()
try {
val actualParent = declaration.parent
if (actualParent != expectedParent) {
errors.add(Data(declaration, expectedParent, actualParent))
}
} catch (e: Exception) {
errors.add(Data(declaration, expectedParent, null))
}
}
}