[FIR] Return symbols instead of FIR from FirPredicateBasedProvider

This commit is contained in:
Dmitriy Novozhilov
2021-10-04 16:28:02 +03:00
committed by TeamCityServer
parent 1cfe4deda9
commit 9a802e7cd7
6 changed files with 19 additions and 12 deletions
@@ -15,6 +15,7 @@ import org.jetbrains.kotlin.fir.declarations.FirAnnotatedDeclaration
import org.jetbrains.kotlin.fir.declarations.FirFile import org.jetbrains.kotlin.fir.declarations.FirFile
import org.jetbrains.kotlin.fir.extensions.predicate.* import org.jetbrains.kotlin.fir.extensions.predicate.*
import org.jetbrains.kotlin.fir.resolve.fqName import org.jetbrains.kotlin.fir.resolve.fqName
import org.jetbrains.kotlin.fir.symbols.FirBasedSymbol
abstract class FirPredicateBasedProvider : FirSessionComponent { abstract class FirPredicateBasedProvider : FirSessionComponent {
companion object { companion object {
@@ -23,8 +24,8 @@ abstract class FirPredicateBasedProvider : FirSessionComponent {
} }
} }
abstract fun getSymbolsByPredicate(predicate: DeclarationPredicate): List<FirAnnotatedDeclaration> abstract fun getSymbolsByPredicate(predicate: DeclarationPredicate): List<FirBasedSymbol<*>>
abstract fun getOwnersOfDeclaration(declaration: FirAnnotatedDeclaration): List<FirAnnotatedDeclaration>? abstract fun getOwnersOfDeclaration(declaration: FirAnnotatedDeclaration): List<FirBasedSymbol<*>>?
abstract fun fileHasPluginAnnotations(file: FirFile): Boolean abstract fun fileHasPluginAnnotations(file: FirFile): Boolean
abstract fun matches(predicate: DeclarationPredicate, declaration: FirAnnotatedDeclaration): Boolean abstract fun matches(predicate: DeclarationPredicate, declaration: FirAnnotatedDeclaration): Boolean
@@ -38,13 +39,13 @@ class FirPredicateBasedProviderImpl(private val session: FirSession) : FirPredic
private val registeredPluginAnnotations = session.registeredPluginAnnotations private val registeredPluginAnnotations = session.registeredPluginAnnotations
private val cache = Cache() private val cache = Cache()
override fun getSymbolsByPredicate(predicate: DeclarationPredicate): List<FirAnnotatedDeclaration> { override fun getSymbolsByPredicate(predicate: DeclarationPredicate): List<FirBasedSymbol<*>> {
val annotations = registeredPluginAnnotations.getAnnotationsForPredicate(predicate) val annotations = registeredPluginAnnotations.getAnnotationsForPredicate(predicate)
if (annotations.isEmpty()) return emptyList() if (annotations.isEmpty()) return emptyList()
val declarations = annotations.flatMapTo(mutableSetOf()) { val declarations = annotations.flatMapTo(mutableSetOf()) {
cache.declarationByAnnotation[it] + cache.declarationsUnderAnnotated[it] cache.declarationByAnnotation[it] + cache.declarationsUnderAnnotated[it]
} }
return declarations.filter { matches(predicate, it) } return declarations.filter { matches(predicate, it) }.map { it.symbol }
} }
override fun fileHasPluginAnnotations(file: FirFile): Boolean { override fun fileHasPluginAnnotations(file: FirFile): Boolean {
@@ -65,8 +66,8 @@ class FirPredicateBasedProviderImpl(private val session: FirSession) : FirPredic
cache.filesWithPluginAnnotations += file cache.filesWithPluginAnnotations += file
} }
override fun getOwnersOfDeclaration(declaration: FirAnnotatedDeclaration): List<FirAnnotatedDeclaration>? { override fun getOwnersOfDeclaration(declaration: FirAnnotatedDeclaration): List<FirBasedSymbol<*>>? {
return cache.ownersForDeclaration[declaration] return cache.ownersForDeclaration[declaration]?.map { it.symbol }
} }
private fun registerOwnersDeclarations(declaration: FirAnnotatedDeclaration, owners: PersistentList<FirAnnotatedDeclaration>) { private fun registerOwnersDeclarations(declaration: FirAnnotatedDeclaration, owners: PersistentList<FirAnnotatedDeclaration>) {
@@ -7,8 +7,10 @@ package org.jetbrains.kotlin.fir.symbols
import org.jetbrains.kotlin.fir.FirModuleData import org.jetbrains.kotlin.fir.FirModuleData
import org.jetbrains.kotlin.fir.FirSourceElement import org.jetbrains.kotlin.fir.FirSourceElement
import org.jetbrains.kotlin.fir.declarations.FirAnnotatedDeclaration
import org.jetbrains.kotlin.fir.declarations.FirDeclaration import org.jetbrains.kotlin.fir.declarations.FirDeclaration
import org.jetbrains.kotlin.fir.declarations.FirDeclarationOrigin import org.jetbrains.kotlin.fir.declarations.FirDeclarationOrigin
import org.jetbrains.kotlin.fir.expressions.FirAnnotation
abstract class FirBasedSymbol<E : FirDeclaration> { abstract class FirBasedSymbol<E : FirDeclaration> {
private var _fir: E? = null private var _fir: E? = null
@@ -30,6 +32,9 @@ abstract class FirBasedSymbol<E : FirDeclaration> {
val moduleData: FirModuleData val moduleData: FirModuleData
get() = fir.moduleData get() = fir.moduleData
val annotations: List<FirAnnotation>
get() = (fir as? FirAnnotatedDeclaration)?.annotations ?: emptyList()
} }
@RequiresOptIn @RequiresOptIn
@@ -22,6 +22,7 @@ import org.jetbrains.kotlin.fir.extensions.predicate.hasOrUnder
import org.jetbrains.kotlin.fir.extensions.predicateBasedProvider import org.jetbrains.kotlin.fir.extensions.predicateBasedProvider
import org.jetbrains.kotlin.fir.extensions.transform import org.jetbrains.kotlin.fir.extensions.transform
import org.jetbrains.kotlin.fir.references.FirNamedReference import org.jetbrains.kotlin.fir.references.FirNamedReference
import org.jetbrains.kotlin.fir.symbols.FirBasedSymbol
import org.jetbrains.kotlin.fir.types.ConeClassLikeType import org.jetbrains.kotlin.fir.types.ConeClassLikeType
import org.jetbrains.kotlin.fir.types.coneTypeSafe import org.jetbrains.kotlin.fir.types.coneTypeSafe
import org.jetbrains.kotlin.name.ClassId import org.jetbrains.kotlin.name.ClassId
@@ -48,15 +49,15 @@ class AllOpenVisibilityTransformer(session: FirSession) : FirStatusTransformerEx
return status.transform(visibility = visibility) return status.transform(visibility = visibility)
} }
private fun findVisibility(declaration: FirDeclaration, owners: List<FirAnnotatedDeclaration>): Visibility? { private fun findVisibility(declaration: FirDeclaration, owners: List<FirBasedSymbol<*>>): Visibility? {
(declaration as? FirAnnotatedDeclaration)?.visibilityFromAnnotation()?.let { return it } declaration.symbol.visibilityFromAnnotation()?.let { return it }
for (owner in owners) { for (owner in owners) {
owner.visibilityFromAnnotation()?.let { return it } owner.visibilityFromAnnotation()?.let { return it }
} }
return null return null
} }
private fun FirAnnotatedDeclaration.visibilityFromAnnotation(): Visibility? { private fun FirBasedSymbol<*>.visibilityFromAnnotation(): Visibility? {
val annotation = annotations.firstOrNull { val annotation = annotations.firstOrNull {
it.annotationTypeRef.coneTypeSafe<ConeClassLikeType>()?.lookupTag?.classId == AllPublicClassId it.annotationTypeRef.coneTypeSafe<ConeClassLikeType>()?.lookupTag?.classId == AllPublicClassId
} as? FirAnnotationCall ?: return null } as? FirAnnotationCall ?: return null
@@ -55,7 +55,7 @@ class AllOpenClassGenerator(session: FirSession) : FirDeclarationGenerationExten
private val predicateBasedProvider = session.predicateBasedProvider private val predicateBasedProvider = session.predicateBasedProvider
private val matchedClasses by lazy { private val matchedClasses by lazy {
predicateBasedProvider.getSymbolsByPredicate(PREDICATE).map { it.symbol }.filterIsInstance<FirRegularClassSymbol>() predicateBasedProvider.getSymbolsByPredicate(PREDICATE).filterIsInstance<FirRegularClassSymbol>()
} }
private val classIdsForMatchedClasses: Map<ClassId, FirRegularClassSymbol> by lazy { private val classIdsForMatchedClasses: Map<ClassId, FirRegularClassSymbol> by lazy {
matchedClasses.associateBy { matchedClasses.associateBy {
@@ -42,7 +42,7 @@ class AllOpenMembersGenerator(session: FirSession) : FirDeclarationGenerationExt
private val predicateBasedProvider = session.predicateBasedProvider private val predicateBasedProvider = session.predicateBasedProvider
private val matchedClasses by lazy { private val matchedClasses by lazy {
predicateBasedProvider.getSymbolsByPredicate(PREDICATE).map { it.symbol }.filterIsInstance<FirRegularClassSymbol>() predicateBasedProvider.getSymbolsByPredicate(PREDICATE).filterIsInstance<FirRegularClassSymbol>()
} }
override fun generateFunctions(callableId: CallableId, owner: FirClassSymbol<*>?): List<FirNamedFunctionSymbol> { override fun generateFunctions(callableId: CallableId, owner: FirClassSymbol<*>?): List<FirNamedFunctionSymbol> {
@@ -38,7 +38,7 @@ class AllOpenTopLevelDeclarationsGenerator(session: FirSession) : FirDeclaration
private val predicateBasedProvider = session.predicateBasedProvider private val predicateBasedProvider = session.predicateBasedProvider
private val matchedClasses by lazy { private val matchedClasses by lazy {
predicateBasedProvider.getSymbolsByPredicate(PREDICATE).map { it.symbol }.filterIsInstance<FirRegularClassSymbol>() predicateBasedProvider.getSymbolsByPredicate(PREDICATE).filterIsInstance<FirRegularClassSymbol>()
} }
override fun generateFunctions(callableId: CallableId, owner: FirClassSymbol<*>?): List<FirNamedFunctionSymbol> { override fun generateFunctions(callableId: CallableId, owner: FirClassSymbol<*>?): List<FirNamedFunctionSymbol> {