FIR2IR: Rework fake overrides generation

- To discriminate what's already been generated, use the set of declaration
instead of names (it's obviously more correct)
- Make it possible to set more then one overridden (base)
This commit is contained in:
Denis Zharkov
2020-09-24 14:36:40 +03:00
parent b241161c35
commit 5c9187b270
30 changed files with 216 additions and 160 deletions
@@ -283,6 +283,7 @@ internal tailrec fun FirCallableSymbol<*>.deepestOverriddenSymbol(): FirCallable
}
internal tailrec fun FirCallableSymbol<*>.deepestMatchingOverriddenSymbol(root: FirCallableSymbol<*> = this): FirCallableSymbol<*> {
if (isIntersectionOverride) return this
val overriddenSymbol = overriddenSymbol?.takeIf { it.callableId == root.callableId } ?: return this
return overriddenSymbol.deepestMatchingOverriddenSymbol(this)
}
@@ -106,9 +106,9 @@ class Fir2IrConverter(
}
// Add delegated members *before* fake override generations.
// Otherwise, fake overrides for delegated members, which are redundant, will be added.
processedCallableNames += delegatedMemberNames(irClass)
val realDeclarations = delegatedMembers(irClass) + anonymousObject.declarations
with(fakeOverrideGenerator) {
irClass.addFakeOverrides(anonymousObject, processedCallableNames)
irClass.addFakeOverrides(anonymousObject, realDeclarations)
}
return irClass
@@ -124,47 +124,39 @@ class Fir2IrConverter(
if (irConstructor != null) {
irClass.declarations += irConstructor
}
val processedCallableNames = mutableSetOf<Name>()
val allDeclarations = regularClass.declarations.toMutableList()
for (declaration in sortBySynthetic(regularClass.declarations)) {
val irDeclaration = processMemberDeclaration(declaration, regularClass, irClass) ?: continue
when (declaration) {
is FirSimpleFunction -> processedCallableNames += declaration.name
is FirProperty -> processedCallableNames += declaration.name
}
irClass.declarations += irDeclaration
}
// Add delegated members *before* fake override generations.
// Otherwise, fake overrides for delegated members, which are redundant, will be added.
processedCallableNames += delegatedMemberNames(irClass)
allDeclarations += delegatedMembers(irClass)
// Add synthetic members *before* fake override generations.
// Otherwise, redundant members, e.g., synthetic toString _and_ fake override toString, will be added.
if (irConstructor != null && (irClass.isInline || irClass.isData)) {
declarationStorage.enterScope(irConstructor)
val dataClassMembersGenerator = DataClassMembersGenerator(components)
if (irClass.isInline) {
processedCallableNames += dataClassMembersGenerator.generateInlineClassMembers(regularClass, irClass)
allDeclarations += dataClassMembersGenerator.generateInlineClassMembers(regularClass, irClass)
}
if (irClass.isData) {
processedCallableNames += dataClassMembersGenerator.generateDataClassMembers(regularClass, irClass)
allDeclarations += dataClassMembersGenerator.generateDataClassMembers(regularClass, irClass)
}
declarationStorage.leaveScope(irConstructor)
}
with(fakeOverrideGenerator) {
irClass.addFakeOverrides(regularClass, processedCallableNames)
irClass.addFakeOverrides(regularClass, allDeclarations)
}
return irClass
}
private fun delegatedMemberNames(irClass: IrClass): List<Name> {
private fun delegatedMembers(irClass: IrClass): List<FirDeclaration> {
return irClass.declarations.filter {
it.origin == IrDeclarationOrigin.DELEGATED_MEMBER
}.mapNotNull {
when (it) {
is IrSimpleFunction -> it.name
is IrProperty -> it.name
else -> null
}
components.declarationStorage.originalDeclarationForDelegated(it)
}
}
@@ -7,11 +7,15 @@ package org.jetbrains.kotlin.fir.backend.generators
import org.jetbrains.kotlin.descriptors.*
import org.jetbrains.kotlin.descriptors.Visibilities
import org.jetbrains.kotlin.fir.backend.*
import org.jetbrains.kotlin.fir.backend.Fir2IrComponents
import org.jetbrains.kotlin.fir.backend.FirMetadataSource
import org.jetbrains.kotlin.fir.backend.declareThisReceiverParameter
import org.jetbrains.kotlin.fir.backend.toIrType
import org.jetbrains.kotlin.fir.declarations.*
import org.jetbrains.kotlin.fir.declarations.builder.buildSimpleFunction
import org.jetbrains.kotlin.fir.declarations.builder.buildValueParameter
import org.jetbrains.kotlin.fir.declarations.impl.FirDeclarationStatusImpl
import org.jetbrains.kotlin.fir.scopes.unsubstitutedScope
import org.jetbrains.kotlin.fir.symbols.CallableId
import org.jetbrains.kotlin.fir.symbols.impl.FirNamedFunctionSymbol
import org.jetbrains.kotlin.fir.symbols.impl.FirVariableSymbol
@@ -44,10 +48,10 @@ import org.jetbrains.kotlin.descriptors.DescriptorVisibilities
@OptIn(ObsoleteDescriptorBasedAPI::class)
class DataClassMembersGenerator(val components: Fir2IrComponents) {
fun generateInlineClassMembers(klass: FirClass<*>, irClass: IrClass): List<Name> =
fun generateInlineClassMembers(klass: FirClass<*>, irClass: IrClass): List<FirDeclaration> =
MyDataClassMethodsGenerator(irClass, klass.symbol.classId, IrDeclarationOrigin.GENERATED_INLINE_CLASS_MEMBER).generate(klass)
fun generateDataClassMembers(klass: FirClass<*>, irClass: IrClass): List<Name> =
fun generateDataClassMembers(klass: FirClass<*>, irClass: IrClass): List<FirDeclaration> =
MyDataClassMethodsGenerator(irClass, klass.symbol.classId, IrDeclarationOrigin.GENERATED_DATA_CLASS_MEMBER).generate(klass)
fun generateDataClassComponentBody(irFunction: IrFunction, classId: ClassId) =
@@ -120,7 +124,7 @@ class DataClassMembersGenerator(val components: Fir2IrComponents) {
(this.name == hashCodeName && matchesHashCodeSignature) ||
(this.name == toStringName && matchesToStringSignature)
fun generate(klass: FirClass<*>): List<Name> {
fun generate(klass: FirClass<*>): List<FirDeclaration> {
val propertyParametersCount = irClass.primaryConstructor?.explicitParameters?.size ?: 0
val properties = irClass.declarations
.filterIsInstance<IrProperty>()
@@ -130,7 +134,7 @@ class DataClassMembersGenerator(val components: Fir2IrComponents) {
return emptyList()
}
val result = mutableListOf<Name>()
val result = mutableListOf<FirDeclaration>()
val contributedFunctionsInThisType = klass.declarations.mapNotNull {
if (it is FirSimpleFunction && it.matchesDataClassSyntheticMemberSignatures) {
@@ -138,22 +142,30 @@ class DataClassMembersGenerator(val components: Fir2IrComponents) {
} else
null
}
val nonOverridableContributedFunctionsInSupertypes =
klass.collectContributedFunctionsFromSupertypes(components.session) { declaration, map ->
if (declaration is FirSimpleFunction &&
declaration.body != null &&
!Visibilities.isPrivate(declaration.visibility) &&
declaration.modality == Modality.FINAL &&
declaration.matchesDataClassSyntheticMemberSignatures
) {
map.putIfAbsent(declaration.name, declaration)
val contributedFunctionsInSupertypes =
@OptIn(ExperimentalStdlibApi::class)
buildMap<Name, FirSimpleFunction> {
for (name in listOf(equalsName, hashCodeName, toStringName)) {
klass.unsubstitutedScope(components.session, components.scopeSession).processFunctionsByName(name) {
val declaration = it.fir
if (declaration is FirSimpleFunction &&
declaration.matchesDataClassSyntheticMemberSignatures
) {
putIfAbsent(declaration.name, declaration)
}
}
}
}
fun isOverridableDeclaration(name: Name): Boolean {
val declaration = contributedFunctionsInSupertypes[name] ?: return false
return declaration.modality != Modality.FINAL
}
if (!contributedFunctionsInThisType.contains(equalsName) &&
!nonOverridableContributedFunctionsInSupertypes.containsKey(equalsName)
isOverridableDeclaration(equalsName)
) {
result.add(equalsName)
result.add(contributedFunctionsInSupertypes.getValue(equalsName))
val equalsFunction = createSyntheticIrFunction(
equalsName,
components.irBuiltIns.booleanType,
@@ -164,9 +176,9 @@ class DataClassMembersGenerator(val components: Fir2IrComponents) {
}
if (!contributedFunctionsInThisType.contains(hashCodeName) &&
!nonOverridableContributedFunctionsInSupertypes.containsKey(hashCodeName)
isOverridableDeclaration(hashCodeName)
) {
result.add(hashCodeName)
result.add(contributedFunctionsInSupertypes.getValue(hashCodeName))
val hashCodeFunction = createSyntheticIrFunction(
hashCodeName,
components.irBuiltIns.intType,
@@ -176,9 +188,9 @@ class DataClassMembersGenerator(val components: Fir2IrComponents) {
}
if (!contributedFunctionsInThisType.contains(toStringName) &&
!nonOverridableContributedFunctionsInSupertypes.containsKey(toStringName)
isOverridableDeclaration(toStringName)
) {
result.add(toStringName)
result.add(contributedFunctionsInSupertypes.getValue(toStringName))
val toStringFunction = createSyntheticIrFunction(
toStringName,
components.irBuiltIns.stringType,
@@ -6,11 +6,16 @@
package org.jetbrains.kotlin.fir.backend.generators
import org.jetbrains.kotlin.fir.FirSession
import org.jetbrains.kotlin.fir.FirSymbolOwner
import org.jetbrains.kotlin.fir.backend.*
import org.jetbrains.kotlin.fir.declarations.*
import org.jetbrains.kotlin.fir.resolve.ScopeSession
import org.jetbrains.kotlin.fir.scopes.FirTypeScope
import org.jetbrains.kotlin.fir.scopes.getDirectOverriddenFunctions
import org.jetbrains.kotlin.fir.scopes.getDirectOverriddenProperties
import org.jetbrains.kotlin.fir.scopes.impl.FirClassSubstitutionScope
import org.jetbrains.kotlin.fir.scopes.unsubstitutedScope
import org.jetbrains.kotlin.fir.symbols.AbstractFirBasedSymbol
import org.jetbrains.kotlin.fir.symbols.PossiblyFirFakeOverrideSymbol
import org.jetbrains.kotlin.fir.symbols.impl.FirCallableSymbol
import org.jetbrains.kotlin.fir.symbols.impl.FirNamedFunctionSymbol
@@ -23,7 +28,7 @@ import org.jetbrains.kotlin.ir.types.IrSimpleType
import org.jetbrains.kotlin.ir.types.IrType
import org.jetbrains.kotlin.ir.types.IrTypeProjection
import org.jetbrains.kotlin.load.java.JavaDescriptorVisibilities
import org.jetbrains.kotlin.name.Name
import org.jetbrains.kotlin.name.ClassId
class FakeOverrideGenerator(
private val session: FirSession,
@@ -33,8 +38,8 @@ class FakeOverrideGenerator(
private val conversionScope: Fir2IrConversionScope
) {
private val baseFunctionSymbols = mutableMapOf<IrFunction, FirNamedFunctionSymbol>()
private val basePropertySymbols = mutableMapOf<IrProperty, FirPropertySymbol>()
private val baseFunctionSymbols = mutableMapOf<IrFunction, List<FirNamedFunctionSymbol>>()
private val basePropertySymbols = mutableMapOf<IrProperty, List<FirPropertySymbol>>()
private fun IrSimpleFunction.withFunction(f: IrSimpleFunction.() -> Unit): IrSimpleFunction {
return conversionScope.withFunction(this, f)
@@ -58,17 +63,19 @@ class FakeOverrideGenerator(
}
}
fun IrClass.addFakeOverrides(klass: FirClass<*>, processedCallableNames: MutableSet<Name>) {
declarations += getFakeOverrides(klass, processedCallableNames)
fun IrClass.addFakeOverrides(klass: FirClass<*>, declarations: Collection<FirDeclaration>) {
this.declarations += getFakeOverrides(
klass,
declarations
)
}
fun IrClass.getFakeOverrides(klass: FirClass<*>, processedCallableNames: MutableSet<Name>): List<IrDeclaration> {
fun IrClass.getFakeOverrides(klass: FirClass<*>, realDeclarations: Collection<FirDeclaration>): List<IrDeclaration> {
val result = mutableListOf<IrDeclaration>()
val superTypesCallableNames = klass.collectCallableNamesFromSupertypes(session)
val useSiteMemberScope = klass.unsubstitutedScope(session, scopeSession)
val superTypesCallableNames = useSiteMemberScope.getCallableNames()
val realDeclarationSymbols = realDeclarations.filterIsInstance<FirSymbolOwner<*>>().mapTo(mutableSetOf(), FirSymbolOwner<*>::symbol)
for (name in superTypesCallableNames) {
if (name in processedCallableNames) continue
processedCallableNames += name
val isLocal = klass !is FirRegularClass || klass.isLocal
useSiteMemberScope.processFunctionsByName(name) { functionSymbol ->
createFakeOverriddenIfNeeded(
@@ -86,7 +93,10 @@ class FakeOverrideGenerator(
result,
containsErrorTypes = { irFunction ->
irFunction.returnType.containsErrorType() || irFunction.valueParameters.any { it.type.containsErrorType() }
}
},
realDeclarationSymbols,
FirTypeScope::getDirectOverriddenFunctions,
useSiteMemberScope,
)
}
@@ -107,7 +117,10 @@ class FakeOverrideGenerator(
containsErrorTypes = { irProperty ->
irProperty.backingField?.type?.containsErrorType() == true ||
irProperty.getter?.returnType?.containsErrorType() == true
}
},
realDeclarationSymbols,
FirTypeScope::getDirectOverriddenProperties,
useSiteMemberScope,
)
}
}
@@ -122,15 +135,22 @@ class FakeOverrideGenerator(
cachedIrDeclaration: (D) -> I?,
createIrDeclaration: (D, irParent: IrClass, thisReceiverOwner: IrClass?, origin: IrDeclarationOrigin, isLocal: Boolean) -> I,
createFakeOverrideSymbol: (D, S) -> S,
baseSymbols: MutableMap<I, S>,
baseSymbols: MutableMap<I, List<S>>,
result: MutableList<in I>,
containsErrorTypes: (I) -> Boolean
containsErrorTypes: (I) -> Boolean,
realDeclarationSymbols: Set<AbstractFirBasedSymbol<*>>,
computeDirectOverridden: FirTypeScope.(S) -> List<S>,
scope: FirTypeScope,
) where S : FirCallableSymbol<D>, S : PossiblyFirFakeOverrideSymbol<D, S> {
if (originalSymbol !is S) return
if (originalSymbol !is S || originalSymbol in realDeclarationSymbols) return
val originalDeclaration = originalSymbol.fir
val origin = IrDeclarationOrigin.FAKE_OVERRIDE
val baseSymbol = originalSymbol.deepestOverriddenSymbol() as S
if (originalSymbol.isFakeOverride && originalSymbol.callableId.classId == klass.symbol.classId) {
val classId = klass.symbol.classId
if ((originalSymbol.isFakeOverride || originalSymbol.isIntersectionOverride) &&
originalSymbol.callableId.classId == classId
) {
// Substitution case
// NB: see comment above about substituted function' parent
val irDeclaration = cachedIrDeclaration(originalDeclaration)?.takeIf { it.parent == irClass }
@@ -141,7 +161,7 @@ class FakeOverrideGenerator(
isLocal
)
irDeclaration.parent = irClass
baseSymbols[irDeclaration] = baseSymbol
baseSymbols[irDeclaration] = computeBaseSymbols(originalSymbol, baseSymbol, computeDirectOverridden, scope, classId)
result += irDeclaration
} else if (originalDeclaration.allowsToHaveFakeOverrideIn(klass)) {
// Trivial fake override case
@@ -160,58 +180,77 @@ class FakeOverrideGenerator(
return
}
irDeclaration.parent = irClass
baseSymbols[irDeclaration] = baseSymbol
baseSymbols[irDeclaration] = computeBaseSymbols(originalSymbol, baseSymbol, computeDirectOverridden, scope, classId)
result += irDeclaration
}
}
private inline fun <S : FirCallableSymbol<*>> computeBaseSymbols(
symbol: S,
basedSymbol: S,
directOverridden: FirTypeScope.(S) -> List<S>,
scope: FirTypeScope,
containingClassId: ClassId,
): List<S> {
if (!symbol.isIntersectionOverride) return listOf(basedSymbol)
return scope.directOverridden(symbol).map {
@Suppress("UNCHECKED_CAST")
if (it is PossiblyFirFakeOverrideSymbol<*, *> && it.isFakeOverride && it.callableId.classId == containingClassId)
it.overriddenSymbol!! as S
else
it
}
}
fun bindOverriddenSymbols(declarations: List<IrDeclaration>) {
for (declaration in declarations) {
if (declaration.origin != IrDeclarationOrigin.FAKE_OVERRIDE) continue
when (declaration) {
is IrSimpleFunction -> {
val baseSymbol = baseFunctionSymbols[declaration]!!
val overriddenSymbol = declarationStorage.getIrFunctionSymbol(baseSymbol) as IrSimpleFunctionSymbol
val baseSymbols =
baseFunctionSymbols[declaration]!!.map { declarationStorage.getIrFunctionSymbol(it) as IrSimpleFunctionSymbol }
declaration.withFunction {
overriddenSymbols = listOf(overriddenSymbol)
overriddenSymbols = baseSymbols
}
}
is IrProperty -> {
val baseSymbol = basePropertySymbols[declaration]!!
val baseSymbols = basePropertySymbols[declaration]!!
declaration.withProperty {
discardAccessorsAccordingToBaseVisibility(baseSymbol)
setOverriddenSymbolsForAccessors(declarationStorage, declaration.isVar, firOverriddenSymbol = baseSymbol)
discardAccessorsAccordingToBaseVisibility(baseSymbols)
setOverriddenSymbolsForAccessors(declarationStorage, declaration.isVar, baseSymbols)
}
}
}
}
}
private fun IrProperty.discardAccessorsAccordingToBaseVisibility(baseSymbol: FirPropertySymbol) {
// Do not create fake overrides for accessors if not allowed to do so, e.g., private lateinit var.
if (baseSymbol.fir.getter?.allowsToHaveFakeOverride != true) {
getter = null
}
// or private setter
if (baseSymbol.fir.setter?.allowsToHaveFakeOverride != true) {
setter = null
private fun IrProperty.discardAccessorsAccordingToBaseVisibility(baseSymbols: List<FirPropertySymbol>) {
for (baseSymbol in baseSymbols) {
// Do not create fake overrides for accessors if not allowed to do so, e.g., private lateinit var.
if (baseSymbol.fir.getter?.allowsToHaveFakeOverride != true) {
getter = null
}
// or private setter
if (baseSymbol.fir.setter?.allowsToHaveFakeOverride != true) {
setter = null
}
}
}
private fun IrProperty.setOverriddenSymbolsForAccessors(
declarationStorage: Fir2IrDeclarationStorage,
isVar: Boolean,
firOverriddenSymbol: FirPropertySymbol
firOverriddenSymbols: List<FirPropertySymbol>
): IrProperty {
val irSymbol = declarationStorage.getIrPropertySymbol(firOverriddenSymbol) as? IrPropertySymbol ?: return this
val overriddenProperty = irSymbol.owner
val overriddenIrProperties = firOverriddenSymbols.mapNotNull {
(declarationStorage.getIrPropertySymbol(it) as? IrPropertySymbol)?.owner
}
getter?.apply {
overriddenProperty.getter?.symbol?.let { overriddenSymbols = listOf(it) }
overriddenSymbols = overriddenIrProperties.mapNotNull { it.getter?.symbol }
}
if (isVar) {
setter?.apply {
overriddenProperty.setter?.symbol?.let { overriddenSymbols = listOf(it) }
overriddenSymbols = overriddenIrProperties.mapNotNull { it.setter?.symbol }
}
}
return this
@@ -168,7 +168,7 @@ class Fir2IrLazyClass(
}
}
with(fakeOverrideGenerator) {
val fakeOverrides = getFakeOverrides(fir, processedNames)
val fakeOverrides = getFakeOverrides(fir, fir.declarations)
bindOverriddenSymbols(fakeOverrides)
result += fakeOverrides
}