FIR: support > 1 most specific members in type intersection scopes

This commit is contained in:
pyos
2022-10-12 10:15:11 +02:00
committed by teamcity
parent 35450e6e04
commit 2879e7a74c
2 changed files with 49 additions and 70 deletions
@@ -130,7 +130,7 @@ abstract class AbstractFirUseSiteMemberScope(
* *
* TODO: is it enough to check only one function? * TODO: is it enough to check only one function?
*/ */
mostSpecific keySymbol
} else { } else {
chosenSymbol chosenSymbol
} }
@@ -19,7 +19,6 @@ import org.jetbrains.kotlin.fir.declarations.utils.visibility
import org.jetbrains.kotlin.fir.resolve.substitution.ConeSubstitutor import org.jetbrains.kotlin.fir.resolve.substitution.ConeSubstitutor
import org.jetbrains.kotlin.fir.resolve.transformers.ReturnTypeCalculatorForFullBodyResolve import org.jetbrains.kotlin.fir.resolve.transformers.ReturnTypeCalculatorForFullBodyResolve
import org.jetbrains.kotlin.fir.scopes.* import org.jetbrains.kotlin.fir.scopes.*
import org.jetbrains.kotlin.fir.scopes.impl.FirIntersectionOverrideStorage.ContextForIntersectionOverrideConstruction
import org.jetbrains.kotlin.fir.scopes.impl.FirTypeIntersectionScopeContext.ResultOfIntersection import org.jetbrains.kotlin.fir.scopes.impl.FirTypeIntersectionScopeContext.ResultOfIntersection
import org.jetbrains.kotlin.fir.symbols.impl.* import org.jetbrains.kotlin.fir.symbols.impl.*
import org.jetbrains.kotlin.fir.types.ConeKotlinType import org.jetbrains.kotlin.fir.types.ConeKotlinType
@@ -40,8 +39,8 @@ class FirTypeIntersectionScopeContext(
) { ) {
private val overrideService = session.overrideService private val overrideService = session.overrideService
val intersectionOverrides: FirCache<FirCallableSymbol<*>, MemberWithBaseScope<FirCallableSymbol<*>>, ContextForIntersectionOverrideConstruction<*>> = val intersectionOverrides: FirCache<FirCallableSymbol<*>, MemberWithBaseScope<FirCallableSymbol<*>>, ResultOfIntersection.NonTrivial<*>> =
session.intersectionOverrideStorage.cacheByScope.getValue(dispatchReceiverType).intersectionOverrides session.intersectionOverrideStorage.cacheByScope.getValue(dispatchReceiverType)
sealed class ResultOfIntersection<D : FirCallableSymbol<*>>( sealed class ResultOfIntersection<D : FirCallableSymbol<*>>(
val overriddenMembers: List<MemberWithBaseScope<D>>, val overriddenMembers: List<MemberWithBaseScope<D>>,
@@ -61,21 +60,18 @@ class FirTypeIntersectionScopeContext(
} }
class NonTrivial<D : FirCallableSymbol<*>>( class NonTrivial<D : FirCallableSymbol<*>>(
private val intersectionOverridesCache: FirCache<FirCallableSymbol<*>, MemberWithBaseScope<FirCallableSymbol<*>>, ContextForIntersectionOverrideConstruction<*>>, val context: FirTypeIntersectionScopeContext,
private val context: ContextForIntersectionOverrideConstruction<D>, val mostSpecific: List<MemberWithBaseScope<D>>,
overriddenMembers: List<MemberWithBaseScope<D>>, overriddenMembers: List<MemberWithBaseScope<D>>,
containingScope: FirTypeScope? containingScope: FirTypeScope?
) : ResultOfIntersection<D>(overriddenMembers, containingScope) { ) : ResultOfIntersection<D>(overriddenMembers, containingScope) {
override val chosenSymbol: D by lazy { override val chosenSymbol: D by lazy {
@Suppress("UNCHECKED_CAST") @Suppress("UNCHECKED_CAST")
intersectionOverridesCache.getValue( context.intersectionOverrides.getValue(keySymbol, this).member as D
context.mostSpecific,
context
).member as D
} }
val mostSpecific: D val keySymbol: D
get() = context.mostSpecific get() = mostSpecific.first().member
} }
} }
@@ -168,19 +164,9 @@ class FirTypeIntersectionScopeContext(
}.takeIf { it.isNotEmpty() } ?: extractBothWaysWithPrivate }.takeIf { it.isNotEmpty() } ?: extractBothWaysWithPrivate
val baseMembersForIntersection = extractedOverrides.calcBaseMembersForIntersectionOverride() val baseMembersForIntersection = extractedOverrides.calcBaseMembersForIntersectionOverride()
if (baseMembersForIntersection.size > 1) { if (baseMembersForIntersection.size > 1) {
val (mostSpecific, scopeForMostSpecific) = overrideService.selectMostSpecificMember(
baseMembersForIntersection,
ReturnTypeCalculatorForFullBodyResolve
)
val intersectionOverrideContext = ContextForIntersectionOverrideConstruction(
mostSpecific,
this,
extractedOverrides,
scopeForMostSpecific
)
result += ResultOfIntersection.NonTrivial( result += ResultOfIntersection.NonTrivial(
intersectionOverrides, this,
intersectionOverrideContext, overrideService.selectMostSpecificMembers(baseMembersForIntersection, ReturnTypeCalculatorForFullBodyResolve),
extractedOverrides, extractedOverrides,
containingScope = null containingScope = null
) )
@@ -199,18 +185,23 @@ class FirTypeIntersectionScopeContext(
} }
fun <D : FirCallableSymbol<*>> createIntersectionOverride( fun <D : FirCallableSymbol<*>> createIntersectionOverride(
mostSpecific: List<MemberWithBaseScope<D>>,
extractedOverrides: List<MemberWithBaseScope<D>>, extractedOverrides: List<MemberWithBaseScope<D>>,
mostSpecific: D,
scopeForMostSpecific: FirTypeScope
): MemberWithBaseScope<FirCallableSymbol<*>> { ): MemberWithBaseScope<FirCallableSymbol<*>> {
val newModality = chooseIntersectionOverrideModality(extractedOverrides) val newModality = chooseIntersectionOverrideModality(extractedOverrides)
val newVisibility = chooseIntersectionVisibility(extractedOverrides) val newVisibility = chooseIntersectionVisibility(extractedOverrides)
val mostSpecificSymbols = mostSpecific.map { it.member }
val extractedOverridesSymbols = extractedOverrides.map { it.member } val extractedOverridesSymbols = extractedOverrides.map { it.member }
return when (mostSpecific) { val key = mostSpecific.first()
is FirNamedFunctionSymbol -> createIntersectionOverride(mostSpecific, extractedOverridesSymbols, newModality, newVisibility) return when (key.member) {
is FirPropertySymbol -> createIntersectionOverride(mostSpecific, extractedOverridesSymbols, newModality, newVisibility) is FirNamedFunctionSymbol ->
createIntersectionOverrideFunction(mostSpecificSymbols, extractedOverridesSymbols, newModality, newVisibility)
is FirPropertySymbol ->
createIntersectionOverrideProperty(mostSpecificSymbols, extractedOverridesSymbols, newModality, newVisibility)
else -> throw IllegalStateException("Should not be here") else -> throw IllegalStateException("Should not be here")
}.withScope(scopeForMostSpecific) }.withScope(key.baseScope)
} }
private fun <S : FirCallableSymbol<*>> List<MemberWithBaseScope<S>>.calcBaseMembersForIntersectionOverride(): List<MemberWithBaseScope<S>> { private fun <S : FirCallableSymbol<*>> List<MemberWithBaseScope<S>>.calcBaseMembersForIntersectionOverride(): List<MemberWithBaseScope<S>> {
@@ -372,54 +363,50 @@ class FirTypeIntersectionScopeContext(
return maxVisibility return maxVisibility
} }
private fun createIntersectionOverride( private fun createIntersectionOverrideFunction(
mostSpecific: FirNamedFunctionSymbol, mostSpecific: Collection<FirCallableSymbol<*>>,
overrides: Collection<FirCallableSymbol<*>>, overrides: Collection<FirCallableSymbol<*>>,
newModality: Modality?, newModality: Modality?,
newVisibility: Visibility, newVisibility: Visibility,
): FirNamedFunctionSymbol { ): FirNamedFunctionSymbol {
val key = mostSpecific.first() as FirNamedFunctionSymbol
val newSymbol = val keyFir = key.fir
FirIntersectionOverrideFunctionSymbol( val callableId = CallableId(
CallableId( dispatchReceiverType.classId ?: keyFir.dispatchReceiverClassOrNull()?.classId!!,
dispatchReceiverType.classId ?: mostSpecific.dispatchReceiverClassOrNull()?.classId!!, keyFir.name
mostSpecific.fir.name )
), val newSymbol = FirIntersectionOverrideFunctionSymbol(callableId, overrides)
overrides
)
val mostSpecificFunction = mostSpecific.fir
FirFakeOverrideGenerator.createCopyForFirFunction( FirFakeOverrideGenerator.createCopyForFirFunction(
newSymbol, newSymbol, keyFir, session, FirDeclarationOrigin.IntersectionOverride, keyFir.isExpect,
mostSpecificFunction, session, FirDeclarationOrigin.IntersectionOverride,
mostSpecificFunction.isExpect,
newDispatchReceiverType = dispatchReceiverType,
newModality = newModality, newModality = newModality,
newVisibility = newVisibility, newVisibility = newVisibility,
newDispatchReceiverType = dispatchReceiverType,
).apply { ).apply {
originalForIntersectionOverrideAttr = mostSpecific.fir originalForIntersectionOverrideAttr = keyFir
} }
return newSymbol return newSymbol
} }
private fun createIntersectionOverride( private fun createIntersectionOverrideProperty(
mostSpecific: FirPropertySymbol, mostSpecific: Collection<FirCallableSymbol<*>>,
overrides: Collection<FirCallableSymbol<*>>, overrides: Collection<FirCallableSymbol<*>>,
newModality: Modality?, newModality: Modality?,
newVisibility: Visibility, newVisibility: Visibility,
): FirPropertySymbol { ): FirPropertySymbol {
val key = mostSpecific.first() as FirPropertySymbol
val keyFir = key.fir
val callableId = CallableId( val callableId = CallableId(
dispatchReceiverType.classId ?: mostSpecific.dispatchReceiverClassOrNull()?.classId!!, dispatchReceiverType.classId ?: keyFir.dispatchReceiverClassOrNull()?.classId!!,
mostSpecific.fir.name keyFir.name
) )
val newSymbol = FirIntersectionOverridePropertySymbol(callableId, overrides) val newSymbol = FirIntersectionOverridePropertySymbol(callableId, overrides)
val mostSpecificProperty = mostSpecific.fir
FirFakeOverrideGenerator.createCopyForFirProperty( FirFakeOverrideGenerator.createCopyForFirProperty(
newSymbol, mostSpecificProperty, session, FirDeclarationOrigin.IntersectionOverride, newSymbol, keyFir, session, FirDeclarationOrigin.IntersectionOverride,
newModality = newModality, newModality = newModality,
newVisibility = newVisibility, newVisibility = newVisibility,
newDispatchReceiverType = dispatchReceiverType, newDispatchReceiverType = dispatchReceiverType,
).apply { ).apply {
originalForIntersectionOverrideAttr = mostSpecific.fir originalForIntersectionOverrideAttr = keyFir
} }
return newSymbol return newSymbol
} }
@@ -427,26 +414,18 @@ class FirTypeIntersectionScopeContext(
private fun <D : FirCallableSymbol<*>> D.withScope(baseScope: FirTypeScope) = MemberWithBaseScope(this, baseScope) private fun <D : FirCallableSymbol<*>> D.withScope(baseScope: FirTypeScope) = MemberWithBaseScope(this, baseScope)
typealias FirIntersectionOverrideCache =
FirCache<FirCallableSymbol<*>, MemberWithBaseScope<FirCallableSymbol<*>>, ResultOfIntersection.NonTrivial<*>>
class FirIntersectionOverrideStorage(val session: FirSession) : FirSessionComponent { class FirIntersectionOverrideStorage(val session: FirSession) : FirSessionComponent {
private val cachesFactory = session.firCachesFactory private val cachesFactory = session.firCachesFactory
class CacheForScope(cachesFactory: FirCachesFactory) { val cacheByScope: FirCache<ConeKotlinType, FirIntersectionOverrideCache, Nothing?> =
val intersectionOverrides: FirCache<FirCallableSymbol<*>, MemberWithBaseScope<FirCallableSymbol<*>>, ContextForIntersectionOverrideConstruction<*>> = cachesFactory.createCache { _ ->
cachesFactory.createCache { mostSpecific, context -> cachesFactory.createCache { _, result ->
val (_, intersectionScope, extractedOverrides, scopeForMostSpecific) = context result.context.createIntersectionOverride(result.mostSpecific, result.overriddenMembers)
intersectionScope.createIntersectionOverride(extractedOverrides, mostSpecific, scopeForMostSpecific)
} }
} }
data class ContextForIntersectionOverrideConstruction<D : FirCallableSymbol<*>>(
val mostSpecific: D,
val intersectionContext: FirTypeIntersectionScopeContext,
val extractedOverrides: List<MemberWithBaseScope<D>>,
val scopeForMostSpecific: FirTypeScope
)
val cacheByScope: FirCache<ConeKotlinType, CacheForScope, Nothing?> =
cachesFactory.createCache { _ -> CacheForScope(cachesFactory) }
} }
private val FirSession.intersectionOverrideStorage: FirIntersectionOverrideStorage by FirSession.sessionComponentAccessor() private val FirSession.intersectionOverrideStorage: FirIntersectionOverrideStorage by FirSession.sessionComponentAccessor()