[FIR] Replace single supertype scope with list of scopes of supertypes in use site scopes

This big refactoring is needed to cleanup building of overrides
  mappings and prevent creating redundant intersection overrides in
  cases when there is no need in them:

```kotlin
interface A {
    fun foo()
}

interface B {
    fun foo()
}

interface C : A, B {
    override fun foo()
}
```

Before this refactoring there was next override tree:
C.foo
  intersection override (A.foo, B.foo)
    A.foo
    B.foo

Also this commit fixes special mapping of overrides in jvm scopes
  for declarations which have kotlin builtins in supertypes with
  special java mapping rules (collections, for example)
This commit is contained in:
Dmitriy Novozhilov
2022-01-13 18:22:42 +03:00
parent 17916d4a63
commit c80cfb0fdb
82 changed files with 2564 additions and 636 deletions
@@ -55,12 +55,9 @@ class FirKotlinScopeProvider(
useSiteSuperType.scopeForSupertype(useSiteSession, scopeSession, klass)
}
FirClassUseSiteMemberScope(
klass.classId,
klass,
useSiteSession,
FirTypeIntersectionScope.prepareIntersectionScope(
useSiteSession, FirStandardOverrideChecker(useSiteSession), scopes,
klass.defaultType(),
),
scopes,
decoratedDeclaredMemberScope,
)
}
@@ -8,6 +8,7 @@ package org.jetbrains.kotlin.fir.scopes
import org.jetbrains.kotlin.fir.declarations.FirCallableDeclaration
import org.jetbrains.kotlin.fir.declarations.FirProperty
import org.jetbrains.kotlin.fir.declarations.FirSimpleFunction
import org.jetbrains.kotlin.fir.symbols.impl.FirNamedFunctionSymbol
interface FirOverrideChecker {
fun isOverriddenFunction(
@@ -20,3 +21,8 @@ interface FirOverrideChecker {
baseDeclaration: FirProperty
): Boolean
}
fun FirOverrideChecker.isOverriddenFunction(
overrideCandidate: FirNamedFunctionSymbol,
baseDeclaration: FirNamedFunctionSymbol
): Boolean = isOverriddenFunction(overrideCandidate.fir, baseDeclaration.fir)
@@ -56,3 +56,10 @@ internal fun FirOverrideChecker.similarFunctionsOrBothProperties(
else -> error("Unknown fir callable type: $overrideCandidate, $baseDeclaration")
}
}
fun FirOverrideChecker.similarFunctionsOrBothProperties(
overrideCandidate: FirCallableSymbol<*>,
baseDeclaration: FirCallableSymbol<*>
): Boolean {
return similarFunctionsOrBothProperties(overrideCandidate.fir, baseDeclaration.fir)
}
@@ -7,11 +7,10 @@ package org.jetbrains.kotlin.fir.scopes.impl
import org.jetbrains.kotlin.fir.FirSession
import org.jetbrains.kotlin.fir.resolve.substitution.ConeSubstitutor
import org.jetbrains.kotlin.fir.scopes.FirContainingNamesAwareScope
import org.jetbrains.kotlin.fir.scopes.FirOverrideChecker
import org.jetbrains.kotlin.fir.scopes.FirTypeScope
import org.jetbrains.kotlin.fir.scopes.ProcessorAction
import org.jetbrains.kotlin.fir.scopes.*
import org.jetbrains.kotlin.fir.scopes.impl.FirTypeIntersectionScopeContext.ResultOfIntersection
import org.jetbrains.kotlin.fir.symbols.impl.*
import org.jetbrains.kotlin.fir.types.ConeSimpleKotlinType
import org.jetbrains.kotlin.name.ClassId
import org.jetbrains.kotlin.name.Name
@@ -19,100 +18,207 @@ abstract class AbstractFirUseSiteMemberScope(
val classId: ClassId,
session: FirSession,
overrideChecker: FirOverrideChecker,
val superTypesScope: FirTypeScope,
protected val superTypeScopes: List<FirTypeScope>,
dispatchReceiverType: ConeSimpleKotlinType,
protected val declaredMemberScope: FirContainingNamesAwareScope
) : AbstractFirOverrideScope(session, overrideChecker) {
protected val supertypeScopeContext = FirTypeIntersectionScopeContext(session, overrideChecker, superTypeScopes, dispatchReceiverType)
private val functions = hashMapOf<Name, Collection<FirNamedFunctionSymbol>>()
private val properties = hashMapOf<Name, Collection<FirVariableSymbol<*>>>()
val directOverriddenFunctions = hashMapOf<FirNamedFunctionSymbol, Collection<FirNamedFunctionSymbol>>()
protected val directOverriddenProperties = hashMapOf<FirPropertySymbol, MutableList<FirPropertySymbol>>()
private val functions: MutableMap<Name, Collection<FirNamedFunctionSymbol>> = hashMapOf()
private val properties: MutableMap<Name, Collection<FirVariableSymbol<*>>> = hashMapOf()
protected val directOverriddenFunctions: MutableMap<FirNamedFunctionSymbol, List<ResultOfIntersection<FirNamedFunctionSymbol>>> =
hashMapOf()
protected val directOverriddenProperties: MutableMap<FirPropertySymbol, List<ResultOfIntersection<FirPropertySymbol>>> = hashMapOf()
protected val functionsFromSupertypes: MutableMap<Name, List<ResultOfIntersection<FirNamedFunctionSymbol>>> = mutableMapOf()
protected val propertiesFromSupertypes: MutableMap<Name, List<ResultOfIntersection<FirPropertySymbol>>> = mutableMapOf()
protected val fieldsFromSupertypes: MutableMap<Name, List<FirFieldSymbol>> = mutableMapOf()
private val absentClassifiersFromSupertypes = mutableSetOf<Name>()
private val callableNamesCached by lazy(LazyThreadSafetyMode.PUBLICATION) {
declaredMemberScope.getCallableNames() + superTypesScope.getCallableNames()
buildSet {
addAll(declaredMemberScope.getCallableNames())
superTypeScopes.flatMapTo(this) { it.getCallableNames() }
}
}
override fun processFunctionsByName(name: Name, processor: (FirNamedFunctionSymbol) -> Unit) {
private val classifierNamesCached by lazy(LazyThreadSafetyMode.PUBLICATION) {
buildSet {
addAll(declaredMemberScope.getClassifierNames())
superTypeScopes.flatMapTo(this) { it.getClassifierNames() }
}
}
final override fun processFunctionsByName(name: Name, processor: (FirNamedFunctionSymbol) -> Unit) {
functions.getOrPut(name) {
doProcessFunctions(name)
collectFunctions(name)
}.forEach {
processor(it)
}
}
private fun doProcessFunctions(
protected open fun collectFunctions(
name: Name
): Collection<FirNamedFunctionSymbol> = mutableListOf<FirNamedFunctionSymbol>().apply {
val overrideCandidates = mutableSetOf<FirFunctionSymbol<*>>()
collectDeclaredFunctions(name, this)
val explicitlyDeclaredFunctions = this.toSet()
collectFunctionsFromSupertypes(name, this, explicitlyDeclaredFunctions)
}
protected fun collectDeclaredFunctions(name: Name, destination: MutableList<FirNamedFunctionSymbol>) {
declaredMemberScope.processFunctionsByName(name) { symbol ->
if (symbol.isStatic) return@processFunctionsByName
val directOverridden = computeDirectOverridden(symbol)
this@AbstractFirUseSiteMemberScope.directOverriddenFunctions[symbol] = directOverridden
overrideCandidates += symbol
add(symbol)
if (!symbol.isVisibleInCurrentClass()) return@processFunctionsByName
val directOverridden = computeDirectOverriddenForDeclaredFunction(symbol)
directOverriddenFunctions[symbol] = directOverridden
destination += symbol
}
}
superTypesScope.processFunctionsByName(name) {
val overriddenBy = it.getOverridden(overrideCandidates)
protected abstract fun FirNamedFunctionSymbol.isVisibleInCurrentClass(): Boolean
protected fun collectFunctionsFromSupertypes(
name: Name,
destination: MutableList<FirNamedFunctionSymbol>,
explicitlyDeclaredFunctions: Set<FirNamedFunctionSymbol>
) {
for (chosenSymbolFromSupertype in getFunctionsFromSupertypesByName(name)) {
val superSymbol = chosenSymbolFromSupertype.extractSomeSymbolFromSuperType()
if (!superSymbol.isVisibleInCurrentClass()) continue
val overriddenBy = superSymbol.getOverridden(explicitlyDeclaredFunctions)
if (overriddenBy == null) {
add(it)
destination += chosenSymbolFromSupertype.chosenSymbol
}
}
}
private fun getFunctionsFromSupertypesByName(name: Name): List<ResultOfIntersection<FirNamedFunctionSymbol>> {
return functionsFromSupertypes.getOrPut(name) {
supertypeScopeContext.collectCallables(name, FirScope::processFunctionsByName)
}
}
final override fun processPropertiesByName(name: Name, processor: (FirVariableSymbol<*>) -> Unit) {
properties.getOrPut(name) {
doProcessProperties(name)
collectProperties(name)
}.forEach {
processor(it)
}
}
protected abstract fun doProcessProperties(name: Name): Collection<FirVariableSymbol<*>>
protected abstract fun collectProperties(name: Name): Collection<FirVariableSymbol<*>>
private fun computeDirectOverridden(symbol: FirNamedFunctionSymbol): Collection<FirNamedFunctionSymbol> {
val result = mutableListOf<FirNamedFunctionSymbol>()
val firSimpleFunction = symbol.fir
superTypesScope.processFunctionsByName(symbol.callableId.callableName) { superSymbol ->
if (overrideChecker.isOverriddenFunction(firSimpleFunction, superSymbol.fir)) {
result.add(superSymbol)
private fun computeDirectOverriddenForDeclaredFunction(declaredFunctionSymbol: FirNamedFunctionSymbol): List<ResultOfIntersection<FirNamedFunctionSymbol>> {
val result = mutableListOf<ResultOfIntersection<FirNamedFunctionSymbol>>()
val declaredFunction = declaredFunctionSymbol.fir
for (resultOfIntersection in getFunctionsFromSupertypesByName(declaredFunctionSymbol.name)) {
val symbolFromSupertype = resultOfIntersection.extractSomeSymbolFromSuperType()
if (overrideChecker.isOverriddenFunction(declaredFunction, symbolFromSupertype.fir)) {
result.add(resultOfIntersection)
}
}
return result
}
protected fun <D : FirCallableSymbol<*>> ResultOfIntersection<D>.extractSomeSymbolFromSuperType(): D {
return if (this.isIntersectionOverride()) {
/*
* we don't want to create intersection override if some declared function actually overrides some functions
* from supertypes, so instead of intersection override symbol we check actual symbol from supertype
*
* TODO: is it enough to check only one function?
*/
firstMember
} else {
chosenSymbol
}
}
override fun processDirectOverriddenFunctionsWithBaseScope(
functionSymbol: FirNamedFunctionSymbol,
processor: (FirNamedFunctionSymbol, FirTypeScope) -> ProcessorAction
): ProcessorAction =
//directOverriddenFunctions might be not filled for functionSymbol if it is not from processFunctionsByName call
doProcessDirectOverriddenCallables(
functionSymbol, processor, directOverriddenFunctions, superTypesScope,
): ProcessorAction {
return processDirectOverriddenMembersWithBaseScopeImpl(
directOverriddenFunctions,
functionsFromSupertypes,
functionSymbol,
processor,
FirTypeScope::processDirectOverriddenFunctionsWithBaseScope
)
}
override fun processDirectOverriddenPropertiesWithBaseScope(
propertySymbol: FirPropertySymbol,
processor: (FirPropertySymbol, FirTypeScope) -> ProcessorAction
): ProcessorAction =
doProcessDirectOverriddenCallables(
propertySymbol, processor, directOverriddenProperties, superTypesScope,
): ProcessorAction {
return processDirectOverriddenMembersWithBaseScopeImpl(
directOverriddenProperties,
propertiesFromSupertypes,
propertySymbol,
processor,
FirTypeScope::processDirectOverriddenPropertiesWithBaseScope
)
}
private fun <D : FirCallableSymbol<*>> processDirectOverriddenMembersWithBaseScopeImpl(
directOverriddenMap: Map<D, List<ResultOfIntersection<D>>>,
callablesFromSupertypes: Map<Name, List<ResultOfIntersection<D>>>,
callableSymbol: D,
processor: (D, FirTypeScope) -> ProcessorAction,
processDirectOverriddenCallables: FirTypeScope.(D, (D, FirTypeScope) -> ProcessorAction) -> ProcessorAction
): ProcessorAction {
when (val directOverridden = directOverriddenMap[callableSymbol]) {
null -> {
val resultOfIntersection = callablesFromSupertypes[callableSymbol.name]
?.firstOrNull { it.chosenSymbol == callableSymbol }
?: return ProcessorAction.NONE
if (resultOfIntersection.isIntersectionOverride()) {
for ((overridden, baseScope) in resultOfIntersection.overriddenMembers) {
if (!processor(overridden, baseScope)) return ProcessorAction.STOP
}
return ProcessorAction.NONE
} else {
return resultOfIntersection.containingScope
?.processDirectOverriddenCallables(callableSymbol, processor)
?: ProcessorAction.NONE
}
}
else -> {
for (resultOfIntersection in directOverridden) {
for ((overridden, baseScope) in resultOfIntersection.overriddenMembers) {
if (!processor(overridden, baseScope)) return ProcessorAction.STOP
}
}
return ProcessorAction.NONE
}
}
}
override fun processClassifiersByNameWithSubstitution(name: Name, processor: (FirClassifierSymbol<*>, ConeSubstitutor) -> Unit) {
declaredMemberScope.processClassifiersByNameWithSubstitution(name, processor)
superTypesScope.processClassifiersByNameWithSubstitution(name, processor)
if (name in absentClassifiersFromSupertypes) return
val classifiers = supertypeScopeContext.collectClassifiers(name)
if (classifiers.isEmpty()) {
absentClassifiersFromSupertypes += name
return
}
for ((symbol, substitution) in classifiers) {
processor(symbol, substitution)
}
}
override fun processDeclaredConstructors(processor: (FirConstructorSymbol) -> Unit) {
declaredMemberScope.processDeclaredConstructors(processor)
}
override fun getCallableNames(): Set<Name> = callableNamesCached
override fun getCallableNames(): Set<Name> {
return callableNamesCached
}
override fun getClassifierNames(): Set<Name> {
return declaredMemberScope.getClassifierNames() + superTypesScope.getClassifierNames()
return classifierNamesCached
}
}
@@ -7,53 +7,86 @@ package org.jetbrains.kotlin.fir.scopes.impl
import org.jetbrains.kotlin.fir.FirSession
import org.jetbrains.kotlin.fir.declarations.FirClass
import org.jetbrains.kotlin.fir.declarations.FirProperty
import org.jetbrains.kotlin.fir.declarations.utils.classId
import org.jetbrains.kotlin.fir.resolve.defaultType
import org.jetbrains.kotlin.fir.scopes.FirContainingNamesAwareScope
import org.jetbrains.kotlin.fir.scopes.FirTypeScope
import org.jetbrains.kotlin.fir.symbols.impl.FirPropertySymbol
import org.jetbrains.kotlin.fir.symbols.impl.FirVariableSymbol
import org.jetbrains.kotlin.fir.symbols.impl.isStatic
import org.jetbrains.kotlin.name.ClassId
import org.jetbrains.kotlin.fir.symbols.impl.*
import org.jetbrains.kotlin.name.Name
class FirClassUseSiteMemberScope(
classId: ClassId,
klass: FirClass,
session: FirSession,
superTypesScope: FirTypeScope,
superTypeScopes: List<FirTypeScope>,
declaredMemberScope: FirContainingNamesAwareScope
) : AbstractFirUseSiteMemberScope(classId, session, FirStandardOverrideChecker(session), superTypesScope, declaredMemberScope) {
override fun doProcessProperties(name: Name): Collection<FirVariableSymbol<*>> {
val seen = mutableSetOf<FirVariableSymbol<*>>()
val result = mutableSetOf<FirVariableSymbol<*>>()
declaredMemberScope.processPropertiesByName(name) l@{
if (it.isStatic) return@l
if (it is FirPropertySymbol) {
val directOverridden = computeDirectOverridden(it.fir)
this@FirClassUseSiteMemberScope.directOverriddenProperties[it] = directOverridden
) : AbstractFirUseSiteMemberScope(
klass.classId,
session,
FirStandardOverrideChecker(session),
superTypeScopes,
klass.defaultType(),
declaredMemberScope
) {
override fun collectProperties(name: Name): Collection<FirVariableSymbol<*>> {
return buildList {
val explicitlyDeclaredProperties = mutableSetOf<FirVariableSymbol<*>>()
declaredMemberScope.processPropertiesByName(name) { symbol ->
if (symbol.isStatic) return@processPropertiesByName
if (symbol is FirPropertySymbol) {
val directOverridden = computeDirectOverriddenForDeclaredProperty(symbol)
directOverriddenProperties[symbol] = directOverridden
}
explicitlyDeclaredProperties += symbol
add(symbol)
}
seen += it
result += it
}
superTypesScope.processPropertiesByName(name) {
val overriddenBy = it.getOverridden(seen)
if (overriddenBy == null) {
result += it
val (properties, fields) = getPropertiesAndFieldsFromSupertypesByName(name)
for (propertyFromSupertype in properties) {
val superSymbol = propertyFromSupertype.extractSomeSymbolFromSuperType()
val overriddenBy = superSymbol.getOverridden(explicitlyDeclaredProperties)
if (overriddenBy == null) {
add(propertyFromSupertype.chosenSymbol)
}
}
addAll(fields)
}
}
private fun computeDirectOverriddenForDeclaredProperty(declaredPropertySymbol: FirPropertySymbol): List<FirTypeIntersectionScopeContext.ResultOfIntersection<FirPropertySymbol>> {
val result = mutableListOf<FirTypeIntersectionScopeContext.ResultOfIntersection<FirPropertySymbol>>()
val declaredProperty = declaredPropertySymbol.fir
for (resultOfIntersection in getPropertiesAndFieldsFromSupertypesByName(declaredPropertySymbol.name).first) {
val symbolFromSupertype = resultOfIntersection.extractSomeSymbolFromSuperType()
if (overrideChecker.isOverriddenProperty(declaredProperty, symbolFromSupertype.fir)) {
result.add(resultOfIntersection)
}
}
return result
}
private fun computeDirectOverridden(property: FirProperty): MutableList<FirPropertySymbol> {
val result = mutableListOf<FirPropertySymbol>()
superTypesScope.processPropertiesByName(property.name) l@{ superSymbol ->
if (superSymbol !is FirPropertySymbol) return@l
if (overrideChecker.isOverriddenProperty(property, superSymbol.fir)) {
result.add(superSymbol)
private fun getPropertiesAndFieldsFromSupertypesByName(name: Name): Pair<List<FirTypeIntersectionScopeContext.ResultOfIntersection<FirPropertySymbol>>, List<FirFieldSymbol>> {
propertiesFromSupertypes[name]?.let {
return it to fieldsFromSupertypes.getValue(name)
}
val fields = mutableListOf<FirFieldSymbol>()
val properties = supertypeScopeContext.collectCallables<FirPropertySymbol>(name) { propertyName, processor ->
processPropertiesByName(propertyName) {
when (it) {
is FirPropertySymbol -> processor(it)
is FirFieldSymbol -> fields += it
else -> {}
}
}
}
return result
propertiesFromSupertypes[name] = properties
fieldsFromSupertypes[name] = fields
return properties to fields
}
override fun FirNamedFunctionSymbol.isVisibleInCurrentClass(): Boolean {
return true
}
override fun toString(): String {
@@ -0,0 +1,63 @@
/*
* Copyright 2010-2021 JetBrains s.r.o. and Kotlin Programming Language contributors.
* Use of this source code is governed by the Apache 2.0 license that can be found in the license/LICENSE.txt file.
*/
package org.jetbrains.kotlin.fir.scopes.impl
import org.jetbrains.kotlin.fir.PrivateForInline
import org.jetbrains.kotlin.fir.scopes.FirTypeScope
import org.jetbrains.kotlin.fir.scopes.MemberWithBaseScope
import org.jetbrains.kotlin.fir.scopes.ProcessOverriddenWithBaseScope
import org.jetbrains.kotlin.fir.scopes.ProcessorAction
import org.jetbrains.kotlin.fir.symbols.impl.FirCallableSymbol
import org.jetbrains.kotlin.fir.symbols.impl.FirNamedFunctionSymbol
import org.jetbrains.kotlin.fir.symbols.impl.FirPropertySymbol
fun filterOutOverriddenFunctions(extractedOverridden: Collection<MemberWithBaseScope<FirNamedFunctionSymbol>>): Collection<MemberWithBaseScope<FirNamedFunctionSymbol>> {
return filterOutOverridden(extractedOverridden, FirTypeScope::processDirectOverriddenFunctionsWithBaseScope)
}
fun filterOutOverriddenProperties(extractedOverridden: Collection<MemberWithBaseScope<FirPropertySymbol>>): Collection<MemberWithBaseScope<FirPropertySymbol>> {
return filterOutOverridden(extractedOverridden, FirTypeScope::processDirectOverriddenPropertiesWithBaseScope)
}
@OptIn(PrivateForInline::class)
inline fun <D : FirCallableSymbol<*>> filterOutOverridden(
extractedOverridden: Collection<MemberWithBaseScope<D>>,
processAllOverridden: ProcessOverriddenWithBaseScope<D>,
): Collection<MemberWithBaseScope<D>> {
return extractedOverridden.filter { overridden1 ->
extractedOverridden.none { overridden2 ->
overridden1 !== overridden2 && overrides(
overridden2,
overridden1,
processAllOverridden
)
}
}
}
// Whether f overrides g
@PrivateForInline
inline fun <D : FirCallableSymbol<*>> overrides(
f: MemberWithBaseScope<D>,
g: MemberWithBaseScope<D>,
processAllOverridden: ProcessOverriddenWithBaseScope<D>,
): Boolean {
val (fMember, fScope) = f
val (gMember) = g
var result = false
fScope.processAllOverridden(fMember) { overridden, _ ->
if (overridden == gMember) {
result = true
ProcessorAction.STOP
} else {
ProcessorAction.NEXT
}
}
return result
}
@@ -7,10 +7,7 @@ package org.jetbrains.kotlin.fir.scopes.impl
import org.jetbrains.kotlin.fir.FirSession
import org.jetbrains.kotlin.fir.resolve.substitution.ConeSubstitutor
import org.jetbrains.kotlin.fir.scopes.FirOverrideChecker
import org.jetbrains.kotlin.fir.scopes.FirScope
import org.jetbrains.kotlin.fir.scopes.FirTypeScope
import org.jetbrains.kotlin.fir.scopes.ProcessorAction
import org.jetbrains.kotlin.fir.scopes.*
import org.jetbrains.kotlin.fir.symbols.impl.*
import org.jetbrains.kotlin.fir.types.ConeSimpleKotlinType
import org.jetbrains.kotlin.name.Name
@@ -58,9 +55,10 @@ class FirTypeIntersectionScope private constructor(
return
}
for ((chosenSymbol, overriddenMembers) in callablesWithOverridden) {
overriddenSymbols[chosenSymbol] = overriddenMembers
processor(chosenSymbol)
for (resultOfIntersection in callablesWithOverridden) {
val symbol = resultOfIntersection.chosenSymbol
overriddenSymbols[symbol] = resultOfIntersection.overriddenMembers
processor(symbol)
}
}
@@ -17,16 +17,21 @@ import org.jetbrains.kotlin.fir.declarations.utils.visibility
import org.jetbrains.kotlin.fir.resolve.substitution.ConeSubstitutor
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.symbols.impl.*
import org.jetbrains.kotlin.fir.types.*
import org.jetbrains.kotlin.name.CallableId
import org.jetbrains.kotlin.name.Name
import org.jetbrains.kotlin.types.AbstractTypeChecker
import kotlin.contracts.ExperimentalContracts
import kotlin.contracts.contract
typealias MembersByScope<D> = List<Pair<FirTypeScope, List<D>>>
class FirTypeIntersectionScopeContext(
val session: FirSession,
private val overrideChecker: FirOverrideChecker,
@property:PrivateForInline val scopes: List<FirTypeScope>,
val scopes: List<FirTypeScope>,
private val dispatchReceiverType: ConeSimpleKotlinType,
) {
private val typeCheckerState = session.typeContext.newTypeCheckerState(
@@ -37,14 +42,40 @@ class FirTypeIntersectionScopeContext(
val intersectionOverrides: FirCache<FirCallableSymbol<*>, MemberWithBaseScope<FirCallableSymbol<*>>, ContextForIntersectionOverrideConstruction<*>> =
session.intersectionOverrideStorage.cacheByScope.getValue(dispatchReceiverType).intersectionOverrides
data class ResultOfIntersection<D : FirCallableSymbol<*>>(
val chosenSymbol: D,
val overriddenMembers: List<MemberWithBaseScope<D>>
sealed class ResultOfIntersection<D : FirCallableSymbol<*>>(
val overriddenMembers: List<MemberWithBaseScope<D>>,
val containingScope: FirTypeScope?
) {
constructor(
chosenSymbol: D,
overriddenMember: MemberWithBaseScope<D>
) : this(chosenSymbol, listOf(overriddenMember))
abstract val chosenSymbol: D
class SingleMember<D : FirCallableSymbol<*>>(
override val chosenSymbol: D,
overriddenMembers: List<MemberWithBaseScope<D>>,
containingScope: FirTypeScope?
) : ResultOfIntersection<D>(overriddenMembers, containingScope) {
constructor(
chosenSymbol: D,
overriddenMember: MemberWithBaseScope<D>
) : this(chosenSymbol, listOf(overriddenMember), overriddenMember.baseScope)
}
class NonTrivial<D : FirCallableSymbol<*>>(
private val intersectionOverridesCache: FirCache<FirCallableSymbol<*>, MemberWithBaseScope<FirCallableSymbol<*>>, ContextForIntersectionOverrideConstruction<*>>,
private val context: ContextForIntersectionOverrideConstruction<D>,
overriddenMembers: List<MemberWithBaseScope<D>>,
containingScope: FirTypeScope?
) : ResultOfIntersection<D>(overriddenMembers, containingScope) {
override val chosenSymbol: D by lazy {
@Suppress("UNCHECKED_CAST")
intersectionOverridesCache.getValue(
context.mostSpecific,
context
).member as D
}
val firstMember: D
get() = context.extractedOverrides.first().member
}
}
@OptIn(PrivateForInline::class)
@@ -65,12 +96,16 @@ class FirTypeIntersectionScopeContext(
return result
}
fun collectFunctions(name: Name): List<ResultOfIntersection<FirNamedFunctionSymbol>> {
return collectCallables(name, FirScope::processFunctionsByName)
}
@OptIn(PrivateForInline::class)
inline fun <D : FirCallableSymbol<*>> collectCallables(
inline fun <D : FirCallableSymbol<*>> collectMembersByScope(
name: Name,
processCallables: FirScope.(Name, (D) -> Unit) -> Unit
): List<ResultOfIntersection<D>> {
val membersByScope = scopes.mapNotNull { scope ->
): MembersByScope<D> {
return scopes.mapNotNull { scope ->
val resultForScope = mutableListOf<D>()
scope.processCallables(name) {
if (it !is FirConstructorSymbol) {
@@ -82,19 +117,25 @@ class FirTypeIntersectionScopeContext(
scope to it
}
}
return collectCallablesImpl(membersByScope)
}
@PrivateForInline
@OptIn(PrivateForInline::class)
inline fun <D : FirCallableSymbol<*>> collectCallables(
name: Name,
processCallables: FirScope.(Name, (D) -> Unit) -> Unit
): List<ResultOfIntersection<D>> {
return collectCallablesImpl(collectMembersByScope(name, processCallables))
}
fun <D : FirCallableSymbol<*>> collectCallablesImpl(
membersByScope: List<Pair<FirTypeScope, MutableList<D>>>
membersByScope: List<Pair<FirTypeScope, List<D>>>
): List<ResultOfIntersection<D>> {
if (membersByScope.isEmpty()) {
return emptyList()
}
membersByScope.singleOrNull()?.let { (scope, members) ->
return members.map { ResultOfIntersection(it, MemberWithBaseScope(it, scope)) }
return members.map { ResultOfIntersection.SingleMember(it, MemberWithBaseScope(it, scope)) }
}
val allMembersWithScope = membersByScope.flatMapTo(linkedSetOf()) { (scope, members) ->
@@ -112,31 +153,32 @@ class FirTypeIntersectionScopeContext(
val baseMembersForIntersection = extractedOverrides.calcBaseMembersForIntersectionOverride()
if (baseMembersForIntersection.size > 1) {
val (mostSpecific, scopeForMostSpecific) = selectMostSpecificMember(baseMembersForIntersection)
val intersectionOverride = intersectionOverrides.getValue(
val intersectionOverrideContext = ContextForIntersectionOverrideConstruction(
mostSpecific,
ContextForIntersectionOverrideConstruction(
this,
extractedOverrides,
scopeForMostSpecific
)
this,
extractedOverrides,
scopeForMostSpecific
)
result += ResultOfIntersection.NonTrivial(
intersectionOverrides,
intersectionOverrideContext,
extractedOverrides,
containingScope = null
)
@Suppress("UNCHECKED_CAST")
result += ResultOfIntersection(intersectionOverride.member as D, extractedOverrides)
} else {
val mostSpecific = baseMembersForIntersection.single().member
result += ResultOfIntersection(mostSpecific, extractedOverrides)
val (mostSpecific, containingScope) = baseMembersForIntersection.single()
result += ResultOfIntersection.SingleMember(mostSpecific, extractedOverrides, containingScope)
}
}
if (allMembersWithScope.isNotEmpty()) {
val single = allMembersWithScope.single().member
result += ResultOfIntersection(single, allMembersWithScope.toList())
val (single, containingScope) = allMembersWithScope.single()
result += ResultOfIntersection.SingleMember(single, allMembersWithScope.toList(), containingScope)
}
return result
}
fun <D : FirCallableSymbol<*>> createIntersectionOverride(
extractedOverrides: List<MemberWithBaseScope<D>>,
mostSpecific: D,
@@ -384,45 +426,6 @@ class FirTypeIntersectionScopeContext(
}
}
private fun <D : FirCallableSymbol<*>> filterOutOverridden(
extractedOverridden: Collection<MemberWithBaseScope<D>>,
processAllOverridden: ProcessOverriddenWithBaseScope<D>,
): Collection<MemberWithBaseScope<D>> {
return extractedOverridden.filter { overridden1 ->
extractedOverridden.none { overridden2 ->
overridden1 !== overridden2 && overrides(
overridden2,
overridden1,
processAllOverridden
)
}
}
}
// Whether f overrides g
private fun <D : FirCallableSymbol<*>> overrides(
f: MemberWithBaseScope<D>,
g: MemberWithBaseScope<D>,
processAllOverridden: ProcessOverriddenWithBaseScope<D>,
): Boolean {
val (fMember, fScope) = f
val (gMember) = g
var result = false
fScope.processAllOverridden(fMember) { overridden, _ ->
if (overridden == gMember) {
result = true
ProcessorAction.STOP
} else {
ProcessorAction.NEXT
}
}
return result
}
private fun <D : FirCallableSymbol<*>> chooseIntersectionVisibility(
extractedOverrides: Collection<MemberWithBaseScope<D>>
): Visibility {
@@ -492,19 +495,6 @@ class FirTypeIntersectionScopeContext(
}
}
class MemberWithBaseScope<out D : FirCallableSymbol<*>>(val member: D, val baseScope: FirTypeScope) {
operator fun component1() = member
operator fun component2() = baseScope
override fun equals(other: Any?): Boolean {
return other is MemberWithBaseScope<*> && member == other.member
}
override fun hashCode(): Int {
return member.hashCode()
}
}
private fun <D : FirCallableSymbol<*>> D.withScope(baseScope: FirTypeScope) = MemberWithBaseScope(this, baseScope)
class FirIntersectionOverrideStorage(val session: FirSession) : FirSessionComponent {
@@ -513,12 +503,13 @@ class FirIntersectionOverrideStorage(val session: FirSession) : FirSessionCompon
class CacheForScope(cachesFactory: FirCachesFactory) {
val intersectionOverrides: FirCache<FirCallableSymbol<*>, MemberWithBaseScope<FirCallableSymbol<*>>, ContextForIntersectionOverrideConstruction<*>> =
cachesFactory.createCache { mostSpecific, context ->
val (intersectionScope, extractedOverrides, scopeForMostSpecific) = context
val (_, intersectionScope, extractedOverrides, scopeForMostSpecific) = context
intersectionScope.createIntersectionOverride(extractedOverrides, mostSpecific, scopeForMostSpecific)
}
}
data class ContextForIntersectionOverrideConstruction<D : FirCallableSymbol<*>>(
val mostSpecific: D,
val intersectionContext: FirTypeIntersectionScopeContext,
val extractedOverrides: List<MemberWithBaseScope<D>>,
val scopeForMostSpecific: FirTypeScope
@@ -529,3 +520,11 @@ class FirIntersectionOverrideStorage(val session: FirSession) : FirSessionCompon
}
private val FirSession.intersectionOverrideStorage: FirIntersectionOverrideStorage by FirSession.sessionComponentAccessor()
@OptIn(ExperimentalContracts::class)
fun <D : FirCallableSymbol<*>> ResultOfIntersection<D>.isIntersectionOverride(): Boolean {
contract {
returns(true) implies (this@isIntersectionOverride is ResultOfIntersection.NonTrivial<D>)
}
return this is ResultOfIntersection.NonTrivial
}