[IR] Simplify IrOverridingUtil

Previously, it had some internal state that was saved between calls
and needed to be cleared.

This state was eliminated, which makes invariants clearer and using
easier.

^KT-61934
This commit is contained in:
Pavel Kunyavskiy
2023-09-14 12:23:40 +02:00
parent 0f4c550ee0
commit fcca73e900
3 changed files with 73 additions and 118 deletions
@@ -5,18 +5,15 @@
package org.jetbrains.kotlin.ir.overrides package org.jetbrains.kotlin.ir.overrides
import org.jetbrains.kotlin.descriptors.CallableMemberDescriptor
import org.jetbrains.kotlin.descriptors.DescriptorVisibilities import org.jetbrains.kotlin.descriptors.DescriptorVisibilities
import org.jetbrains.kotlin.descriptors.DescriptorVisibility import org.jetbrains.kotlin.descriptors.DescriptorVisibility
import org.jetbrains.kotlin.descriptors.Modality import org.jetbrains.kotlin.descriptors.Modality
import org.jetbrains.kotlin.ir.declarations.* import org.jetbrains.kotlin.ir.declarations.*
import org.jetbrains.kotlin.ir.linkage.partial.IrUnimplementedOverridesStrategy
import org.jetbrains.kotlin.ir.symbols.IrPropertySymbol import org.jetbrains.kotlin.ir.symbols.IrPropertySymbol
import org.jetbrains.kotlin.ir.symbols.IrSimpleFunctionSymbol import org.jetbrains.kotlin.ir.symbols.IrSimpleFunctionSymbol
import org.jetbrains.kotlin.ir.symbols.IrSymbol import org.jetbrains.kotlin.ir.symbols.IrSymbol
import org.jetbrains.kotlin.ir.types.* import org.jetbrains.kotlin.ir.types.*
import org.jetbrains.kotlin.ir.util.collectAndFilterRealOverrides import org.jetbrains.kotlin.ir.util.collectAndFilterRealOverrides
import org.jetbrains.kotlin.ir.util.isReal
import org.jetbrains.kotlin.ir.util.render import org.jetbrains.kotlin.ir.util.render
import org.jetbrains.kotlin.resolve.OverridingUtil.OverrideCompatibilityInfo import org.jetbrains.kotlin.resolve.OverridingUtil.OverrideCompatibilityInfo
import org.jetbrains.kotlin.types.AbstractTypeChecker import org.jetbrains.kotlin.types.AbstractTypeChecker
@@ -37,14 +34,8 @@ class IrOverridingUtil(
private val externalOverridabilityConditions: List<IrExternalOverridabilityCondition>, private val externalOverridabilityConditions: List<IrExternalOverridabilityCondition>,
) { ) {
private val overrideChecker = IrOverrideChecker(typeSystem, externalOverridabilityConditions) private val overrideChecker = IrOverrideChecker(typeSystem, externalOverridabilityConditions)
private val originals = mutableMapOf<IrOverridableMember, IrOverridableMember>()
private val IrOverridableMember.original get() = originals[this] ?: error("No original for ${this.render()}")
private val originalSuperTypes = mutableMapOf<IrOverridableMember, IrType>()
fun clear() { private data class FakeOverride(val override: IrOverridableMember, val original: IrOverridableMember)
originals.clear()
originalSuperTypes.clear()
}
private var IrOverridableMember.overriddenSymbols: List<IrSymbol> private var IrOverridableMember.overriddenSymbols: List<IrSymbol>
get() = when (this) { get() = when (this) {
@@ -81,19 +72,18 @@ class IrOverridingUtil(
.map { .map {
val overriddenMember = it as IrOverridableMember val overriddenMember = it as IrOverridableMember
val fakeOverride = fakeOverrideBuilder.fakeOverrideMember(superType, overriddenMember, clazz) val fakeOverride = fakeOverrideBuilder.fakeOverrideMember(superType, overriddenMember, clazz)
originals[fakeOverride] = overriddenMember FakeOverride(fakeOverride, overriddenMember)
originalSuperTypes[fakeOverride] = superType
fakeOverride
} }
} }
val allFromSuperByName = allFromSuper.groupBy { it.name } val allFromSuperByName = allFromSuper.groupBy { it.override.name }
allFromSuperByName.forEach { group -> allFromSuperByName.forEach { group ->
generateOverridesInFunctionGroup( generateOverridesInFunctionGroup(
group.value, group.value,
fromCurrent.filter { it.name == group.key && !it.isStaticMember }, fromCurrent.filter { it.name == group.key && !it.isStaticMember },
clazz, oldSignatures clazz,
oldSignatures
) )
} }
} }
@@ -116,13 +106,11 @@ class IrOverridingUtil(
} }
.map { overriddenMember -> .map { overriddenMember ->
val fakeOverride = fakeOverrideBuilder.fakeOverrideMember(superType, overriddenMember, clazz) val fakeOverride = fakeOverrideBuilder.fakeOverrideMember(superType, overriddenMember, clazz)
originals[fakeOverride] = overriddenMember FakeOverride(fakeOverride, overriddenMember)
originalSuperTypes[fakeOverride] = superType
fakeOverride
} }
} }
val unoverriddenSuperMembersGroupedByName = unoverriddenSuperMembers.groupBy { it.name } val unoverriddenSuperMembersGroupedByName = unoverriddenSuperMembers.groupBy { it.override.name }
val fakeOverrides = mutableListOf<IrOverridableMember>() val fakeOverrides = mutableListOf<IrOverridableMember>()
for (group in unoverriddenSuperMembersGroupedByName.values) { for (group in unoverriddenSuperMembersGroupedByName.values) {
createAndBindFakeOverrides(clazz, group, fakeOverrides, compatibilityMode) createAndBindFakeOverrides(clazz, group, fakeOverrides, compatibilityMode)
@@ -141,7 +129,7 @@ class IrOverridingUtil(
} }
private fun generateOverridesInFunctionGroup( private fun generateOverridesInFunctionGroup(
membersFromSupertypes: List<IrOverridableMember>, membersFromSupertypes: List<FakeOverride>,
membersFromCurrent: List<IrOverridableMember>, membersFromCurrent: List<IrOverridableMember>,
current: IrClass, current: IrClass,
compatibilityMode: Boolean compatibilityMode: Boolean
@@ -160,18 +148,19 @@ class IrOverridingUtil(
private fun extractAndBindOverridesForMember( private fun extractAndBindOverridesForMember(
fromCurrent: IrOverridableMember, fromCurrent: IrOverridableMember,
descriptorsFromSuper: Collection<IrOverridableMember> membersFromSuper: List<FakeOverride>
): Collection<IrOverridableMember> { ): List<FakeOverride> {
val bound = ArrayList<IrOverridableMember>(descriptorsFromSuper.size) val bound = ArrayList<FakeOverride>(membersFromSuper.size)
val overridden = mutableSetOf<IrOverridableMember>() val overridden = mutableSetOf<FakeOverride>()
for (fromSupertype in descriptorsFromSuper) { for (fromSupertype in membersFromSuper) {
// Note: We do allow overriding multiple FOs at once one of which is `isInline=true`. // Note: We do allow overriding multiple FOs at once one of which is `isInline=true`.
when (overrideChecker.isOverridableBy(fromSupertype, fromCurrent, checkIsInlineFlag = true).result) { when (overrideChecker.isOverridableBy(fromSupertype.override, fromCurrent, checkIsInlineFlag = true).result) {
OverrideCompatibilityInfo.Result.OVERRIDABLE -> { OverrideCompatibilityInfo.Result.OVERRIDABLE -> {
val isVisibleFake = fromSupertype.visibility != DescriptorVisibilities.INVISIBLE_FAKE val isVisibleFake = fromSupertype.override.visibility != DescriptorVisibilities.INVISIBLE_FAKE
if (isVisibleFake && isVisibleForOverride(fromCurrent, fromSupertype.original)) if (isVisibleFake && isVisibleForOverride(fromCurrent, fromSupertype.original)) {
overridden += fromSupertype overridden += fromSupertype
}
bound += fromSupertype bound += fromSupertype
} }
OverrideCompatibilityInfo.Result.CONFLICT -> { OverrideCompatibilityInfo.Result.CONFLICT -> {
@@ -186,16 +175,36 @@ class IrOverridingUtil(
return bound return bound
} }
// Based on findMemberWithMaxVisibility from VisibilityUtil.kt.
private fun findMemberWithMaxVisibility(members: Collection<FakeOverride>): FakeOverride {
assert(members.isNotEmpty())
var member: FakeOverride? = null
for (candidate in members) {
if (member == null) {
member = candidate
continue
}
val result = DescriptorVisibilities.compare(member.override.visibility, candidate.override.visibility)
if (result != null && result < 0) {
member = candidate
}
}
return member ?: error("Could not find a visible member")
}
private fun createAndBindFakeOverrides( private fun createAndBindFakeOverrides(
current: IrClass, current: IrClass,
notOverridden: Collection<IrOverridableMember>, notOverridden: Collection<FakeOverride>,
addedFakeOverrides: MutableList<IrOverridableMember>, addedFakeOverrides: MutableList<IrOverridableMember>,
compatibilityMode: Boolean compatibilityMode: Boolean
) { ) {
val fromSuper = notOverridden.toMutableSet() val fromSuper = notOverridden.toMutableSet()
while (fromSuper.isNotEmpty()) { while (fromSuper.isNotEmpty()) {
val notOverriddenFromSuper: IrOverridableMember = findMemberWithMaxVisibility(filterOutCustomizedFakeOverrides(fromSuper)) val notOverriddenFromSuper = findMemberWithMaxVisibility(filterOutCustomizedFakeOverrides(fromSuper))
val overridables: Collection<IrOverridableMember> = extractMembersOverridableInBothWays( val overridables = extractMembersOverridableInBothWays(
notOverriddenFromSuper, notOverriddenFromSuper,
fromSuper fromSuper
) )
@@ -209,22 +218,22 @@ class IrOverridingUtil(
* then leave only true ones. Rationale: They should point to non-abstract callable members in one of super classes, so * then leave only true ones. Rationale: They should point to non-abstract callable members in one of super classes, so
* effectively they are implemented in the current class. * effectively they are implemented in the current class.
*/ */
private fun filterOutCustomizedFakeOverrides(overridableMembers: Collection<IrOverridableMember>): Collection<IrOverridableMember> { private fun filterOutCustomizedFakeOverrides(overridableMembers: Collection<FakeOverride>): Collection<FakeOverride> {
if (overridableMembers.size < 2) return overridableMembers if (overridableMembers.size < 2) return overridableMembers
val (trueFakeOverrides, customizedFakeOverrides) = overridableMembers.partition { it.origin == IrDeclarationOrigin.FAKE_OVERRIDE } val (trueFakeOverrides, customizedFakeOverrides) = overridableMembers.partition { it.override.origin == IrDeclarationOrigin.FAKE_OVERRIDE }
return trueFakeOverrides.ifEmpty { customizedFakeOverrides } return trueFakeOverrides.ifEmpty { customizedFakeOverrides }
} }
private fun determineModalityForFakeOverride( private fun determineModalityForFakeOverride(
members: Collection<IrOverridableMember>, members: List<FakeOverride>,
current: IrClass current: IrClass
): Modality { ): Modality {
// Optimization: avoid creating hash sets in frequent cases when modality can be computed trivially // Optimization: avoid creating hash sets in frequent cases when modality can be computed trivially
var hasOpen = false var hasOpen = false
var hasAbstract = false var hasAbstract = false
for (member in members) { for (member in members) {
when (member.modality) { when (member.override.modality) {
Modality.FINAL -> return Modality.FINAL Modality.FINAL -> return Modality.FINAL
Modality.SEALED -> throw IllegalStateException("Member cannot have SEALED modality: $member") Modality.SEALED -> throw IllegalStateException("Member cannot have SEALED modality: $member")
Modality.OPEN -> hasOpen = true Modality.OPEN -> hasOpen = true
@@ -246,54 +255,20 @@ class IrOverridingUtil(
} }
val realOverrides = members val realOverrides = members
.map { originals[it]!! } .map { it.original }
.collectAndFilterRealOverrides() .collectAndFilterRealOverrides()
return getMinimalModality(realOverrides, transformAbstractToClassModality, current.modality) return getMinimalModality(realOverrides, transformAbstractToClassModality, current.modality)
} }
private fun areEquivalent(a: IrOverridableMember, b: IrOverridableMember) = (a == b)
fun overrides(f: IrOverridableMember, g: IrOverridableMember): Boolean {
if (f != g && areEquivalent(f.original, g.original)) return true
for (overriddenFunction in getOverriddenDeclarations(f)) {
if (areEquivalent(g.original, overriddenFunction)) return true
}
return false
}
private fun getOverriddenDeclarations(member: IrOverridableMember): Set<IrOverridableMember> {
val result = mutableSetOf<IrOverridableMember>()
collectOverriddenDeclarations(member.original, result)
return result
}
private fun collectOverriddenDeclarations(
member: IrOverridableMember,
result: MutableSet<IrOverridableMember>
) {
if (member.isReal) {
result.add(member)
} else {
check(member.overriddenSymbols.isNotEmpty()) { "No overridden descriptors found for (fake override) $member" }
for (overridden in member.original.overriddenSymbols.map { it.owner as IrOverridableMember }) {
val original = overridden.original
collectOverriddenDeclarations(original, result)
result.add(original)
}
}
}
private fun getMinimalModality( private fun getMinimalModality(
descriptors: Collection<IrOverridableMember>, members: Collection<IrOverridableMember>,
transformAbstractToClassModality: Boolean, transformAbstractToClassModality: Boolean,
classModality: Modality classModality: Modality
): Modality { ): Modality {
var result = Modality.ABSTRACT var result = Modality.ABSTRACT
for (descriptor in descriptors) { for (member in members) {
val effectiveModality = val effectiveModality =
if (transformAbstractToClassModality && descriptor.modality === Modality.ABSTRACT) classModality else descriptor.modality if (transformAbstractToClassModality && member.modality === Modality.ABSTRACT) classModality else member.modality
if (effectiveModality < result) { if (effectiveModality < result) {
result = effectiveModality result = effectiveModality
} }
@@ -317,22 +292,22 @@ class IrOverridingUtil(
} }
private fun createAndBindFakeOverride( private fun createAndBindFakeOverride(
overridables: Collection<IrOverridableMember>, overridables: List<FakeOverride>,
currentClass: IrClass, currentClass: IrClass,
addedFakeOverrides: MutableList<IrOverridableMember>, addedFakeOverrides: MutableList<IrOverridableMember>,
compatibilityMode: Boolean compatibilityMode: Boolean
) { ) {
val effectiveOverridden = overridables.filter { it.isVisibleInClass(currentClass) } val effectiveOverridden = overridables.filter { it.override.isVisibleInClass(currentClass) }
// The descriptor based algorithm goes further building invisible fakes here, // The descriptor based algorithm goes further building invisible fakes here,
// but we don't use invisible fakes in IR // but we don't use invisible fakes in IR
if (effectiveOverridden.isEmpty()) return if (effectiveOverridden.isEmpty()) return
val modality = determineModalityForFakeOverride(effectiveOverridden, currentClass) val modality = determineModalityForFakeOverride(effectiveOverridden, currentClass)
val visibility = findMemberWithMaxVisibility(effectiveOverridden).visibility val visibility = findMemberWithMaxVisibility(effectiveOverridden).override.visibility
val mostSpecific = selectMostSpecificMember(effectiveOverridden) val mostSpecific = selectMostSpecificMember(effectiveOverridden)
val fakeOverride = mostSpecific.apply { val fakeOverride = mostSpecific.override.apply {
when (this) { when (this) {
is IrPropertyWithLateBinding -> { is IrPropertyWithLateBinding -> {
this.visibility = visibility this.visibility = visibility
@@ -352,7 +327,7 @@ class IrOverridingUtil(
require( require(
fakeOverride.overriddenSymbols.isNotEmpty() fakeOverride.overriddenSymbols.isNotEmpty()
) { "Overridden symbols should be set for " + CallableMemberDescriptor.Kind.FAKE_OVERRIDE } ) { "Overridden symbols should be set for fake override ${fakeOverride.render()}" }
addedFakeOverrides.add(fakeOverride) addedFakeOverrides.add(fakeOverride)
fakeOverrideBuilder.linkFakeOverride(fakeOverride, compatibilityMode) fakeOverrideBuilder.linkFakeOverride(fakeOverride, compatibilityMode)
@@ -429,13 +404,13 @@ class IrOverridingUtil(
} }
private fun isMoreSpecificThenAllOf( private fun isMoreSpecificThenAllOf(
candidate: IrOverridableMember, candidate: FakeOverride,
descriptors: Collection<IrOverridableMember> overrides: Collection<FakeOverride>
): Boolean { ): Boolean {
// NB subtyping relation in Kotlin is not transitive in presence of flexible types: // NB subtyping relation in Kotlin is not transitive in presence of flexible types:
// String? <: String! <: String, but not String? <: String // String? <: String! <: String, but not String? <: String
for (descriptor in descriptors) { for (override in overrides) {
if (!isMoreSpecific(candidate, descriptor)) { if (!isMoreSpecific(candidate.override, override.override)) {
return false return false
} }
} }
@@ -443,27 +418,27 @@ class IrOverridingUtil(
} }
private fun selectMostSpecificMember( private fun selectMostSpecificMember(
overridables: Collection<IrOverridableMember> overridables: Collection<FakeOverride>
): IrOverridableMember { ): FakeOverride {
require(!overridables.isEmpty()) { "Should have at least one overridable descriptor" } require(!overridables.isEmpty()) { "Should have at least one overridable member" }
if (overridables.size == 1) { if (overridables.size == 1) {
return overridables.first() return overridables.first()
} }
val candidates = mutableListOf<IrOverridableMember>() val candidates = mutableListOf<FakeOverride>()
var transitivelyMostSpecific = overridables.first() var transitivelyMostSpecific = overridables.first()
val transitivelyMostSpecificDescriptor = transitivelyMostSpecific val transitivelyMostSpecificMember = transitivelyMostSpecific
for (overridable in overridables) { for (overridable in overridables) {
if (isMoreSpecificThenAllOf(overridable, overridables) if (isMoreSpecificThenAllOf(overridable, overridables)
) { ) {
candidates.add(overridable) candidates.add(overridable)
} }
if (isMoreSpecific( if (isMoreSpecific(
overridable, overridable.override,
transitivelyMostSpecificDescriptor transitivelyMostSpecificMember.override
) )
&& !isMoreSpecific( && !isMoreSpecific(
transitivelyMostSpecificDescriptor, transitivelyMostSpecificMember.override,
overridable overridable.override
) )
) { ) {
transitivelyMostSpecific = overridable transitivelyMostSpecific = overridable
@@ -474,9 +449,9 @@ class IrOverridingUtil(
} else if (candidates.size == 1) { } else if (candidates.size == 1) {
return candidates.first() return candidates.first()
} }
var firstNonFlexible: IrOverridableMember? = null var firstNonFlexible: FakeOverride? = null
for (candidate in candidates) { for (candidate in candidates) {
if (candidate.returnType !is IrDynamicType) { if (candidate.override.returnType !is IrDynamicType) {
firstNonFlexible = candidate firstNonFlexible = candidate
break break
} }
@@ -485,10 +460,10 @@ class IrOverridingUtil(
} }
private fun extractMembersOverridableInBothWays( private fun extractMembersOverridableInBothWays(
overrider: IrOverridableMember, overrider: FakeOverride,
extractFrom: MutableCollection<IrOverridableMember> extractFrom: MutableSet<FakeOverride>
): Collection<IrOverridableMember> { ): List<FakeOverride> {
val overridable = arrayListOf<IrOverridableMember>() val overridable = arrayListOf<FakeOverride>()
overridable.add(overrider) overridable.add(overrider)
val iterator = extractFrom.iterator() val iterator = extractFrom.iterator()
while (iterator.hasNext()) { while (iterator.hasNext()) {
@@ -497,7 +472,7 @@ class IrOverridingUtil(
iterator.remove() iterator.remove()
continue continue
} }
val finalResult = overrideChecker.getBothWaysOverridability(overrider, candidate) val finalResult = overrideChecker.getBothWaysOverridability(overrider.override, candidate.override)
if (finalResult == OverrideCompatibilityInfo.Result.OVERRIDABLE) { if (finalResult == OverrideCompatibilityInfo.Result.OVERRIDABLE) {
overridable.add(candidate) overridable.add(candidate)
iterator.remove() iterator.remove()
@@ -15,25 +15,6 @@ fun isVisibleForOverride(overriding: IrOverridableMember, fromSuper: IrOverridab
return fromSuper.isVisibleInClass(overriding.parentAsClass) return fromSuper.isVisibleInClass(overriding.parentAsClass)
} }
// Based on findMemberWithMaxVisibility from VisibilityUtil.kt.
fun findMemberWithMaxVisibility(members: Collection<IrOverridableMember>): IrOverridableMember {
assert(members.isNotEmpty())
var member: IrOverridableMember? = null
for (candidate in members) {
if (member == null) {
member = candidate
continue
}
val result = DescriptorVisibilities.compare(member.visibility, candidate.visibility)
if (result != null && result < 0) {
member = candidate
}
}
return member ?: error("Could not find a visible member")
}
fun IrDeclarationWithVisibility.isEffectivelyPrivate(): Boolean { fun IrDeclarationWithVisibility.isEffectivelyPrivate(): Boolean {
fun DescriptorVisibility.isNonPrivate(): Boolean = fun DescriptorVisibility.isNonPrivate(): Boolean =
this == DescriptorVisibilities.PUBLIC this == DescriptorVisibilities.PUBLIC
@@ -281,7 +281,6 @@ class FakeOverrideBuilder(
fun provideFakeOverrides(klass: IrClass, compatibilityMode: CompatibilityMode) { fun provideFakeOverrides(klass: IrClass, compatibilityMode: CompatibilityMode) {
buildFakeOverrideChainsForClass(klass, compatibilityMode) buildFakeOverrideChainsForClass(klass, compatibilityMode)
irOverridingUtil.clear()
haveFakeOverrides.add(klass) haveFakeOverrides.add(klass)
} }