[LL FIR] KT-58325 Implement JvmStubBasedFirDeserializedSymbolProvider as LLFirKotlinSymbolProvider

- This enables merging `JvmStubBasedFirDeserializedSymbolProvider`
  into `LLFirCombinedKotlinSymbolProvider` automatically.
- Crucially, we have to also allow callable caches to take already found
  callables as a context.
- We can also simplify the name cache boilerplate using
  `LLFirKotlinSymbolProviderWithNameCache`.
This commit is contained in:
Marco Pennekamp
2023-05-08 16:16:44 +02:00
committed by Space Team
parent 3322ec9dd9
commit b82a589e56
3 changed files with 85 additions and 56 deletions
@@ -30,18 +30,19 @@ internal class LLFirModuleWithDependenciesSymbolProvider(
val dependencyProvider: LLFirDependenciesSymbolProvider, val dependencyProvider: LLFirDependenciesSymbolProvider,
) : FirSymbolProvider(session) { ) : FirSymbolProvider(session) {
override fun getClassLikeSymbolByClassId(classId: ClassId): FirClassLikeSymbol<*>? = override fun getClassLikeSymbolByClassId(classId: ClassId): FirClassLikeSymbol<*>? =
getClassLikeSymbolByFqNameWithoutDependencies(classId) getClassLikeSymbolByClassIdWithoutDependencies(classId)
?: dependencyProvider.getClassLikeSymbolByClassId(classId) ?: dependencyProvider.getClassLikeSymbolByClassId(classId)
fun getClassLikeSymbolByFqNameWithoutDependencies(classId: ClassId): FirClassLikeSymbol<*>? = fun getClassLikeSymbolByClassIdWithoutDependencies(classId: ClassId): FirClassLikeSymbol<*>? =
providers.firstNotNullOfOrNull { it.getClassLikeSymbolByClassId(classId) } providers.firstNotNullOfOrNull { it.getClassLikeSymbolByClassId(classId) }
fun getClassLikeSymbolByFqNameWithoutDependencies( @OptIn(FirSymbolProviderInternals::class)
fun getDeserializedClassLikeSymbolByClassIdWithoutDependencies(
classId: ClassId,
classLikeDeclaration: KtClassLikeDeclaration, classLikeDeclaration: KtClassLikeDeclaration,
classId: ClassId ): FirClassLikeSymbol<*>? = providers.firstNotNullOfOrNull { provider ->
): FirClassLikeSymbol<*>? { if (provider !is JvmStubBasedFirDeserializedSymbolProvider) return@firstNotNullOfOrNull null
return providers.filterIsInstance(JvmStubBasedFirDeserializedSymbolProvider::class.java) provider.getClassLikeSymbolByClassId(classId, classLikeDeclaration)
.firstNotNullOfOrNull { it.getClassLikeSymbolByClassId(classLikeDeclaration, classId) }
} }
@FirSymbolProviderInternals @FirSymbolProviderInternals
@@ -53,12 +54,14 @@ internal class LLFirModuleWithDependenciesSymbolProvider(
@FirSymbolProviderInternals @FirSymbolProviderInternals
fun getTopLevelDeserializedCallableSymbolsToWithoutDependencies( fun getTopLevelDeserializedCallableSymbolsToWithoutDependencies(
destination: MutableList<FirCallableSymbol<*>>, destination: MutableList<FirCallableSymbol<*>>,
callableDeclaration: KtCallableDeclaration,
packageFqName: FqName, packageFqName: FqName,
shortName: Name shortName: Name,
callableDeclaration: KtCallableDeclaration,
) { ) {
providers.filterIsInstance(JvmStubBasedFirDeserializedSymbolProvider::class.java) providers.forEach { provider ->
.forEach { destination.addIfNotNull(it.getTopLevelCallableSymbol(callableDeclaration, packageFqName, shortName)) } if (provider !is JvmStubBasedFirDeserializedSymbolProvider) return@forEach
destination.addIfNotNull(provider.getTopLevelCallableSymbol(packageFqName, shortName, callableDeclaration))
}
} }
@FirSymbolProviderInternals @FirSymbolProviderInternals
@@ -7,19 +7,19 @@ package org.jetbrains.kotlin.analysis.low.level.api.fir.stubBased.deserializatio
import com.intellij.openapi.project.Project import com.intellij.openapi.project.Project
import com.intellij.psi.search.GlobalSearchScope import com.intellij.psi.search.GlobalSearchScope
import org.jetbrains.kotlin.analysis.low.level.api.fir.providers.LLFirKotlinSymbolProviderWithNameCache
import org.jetbrains.kotlin.analysis.low.level.api.fir.util.LLFirKotlinSymbolProviderNameCache import org.jetbrains.kotlin.analysis.low.level.api.fir.util.LLFirKotlinSymbolProviderNameCache
import org.jetbrains.kotlin.analysis.providers.KotlinDeclarationProvider import org.jetbrains.kotlin.analysis.providers.KotlinDeclarationProvider
import org.jetbrains.kotlin.analysis.providers.createDeclarationProvider import org.jetbrains.kotlin.analysis.providers.createDeclarationProvider
import org.jetbrains.kotlin.fir.FirSession import org.jetbrains.kotlin.fir.FirSession
import org.jetbrains.kotlin.fir.caches.FirCache import org.jetbrains.kotlin.fir.caches.FirCache
import org.jetbrains.kotlin.fir.caches.createCache
import org.jetbrains.kotlin.fir.caches.firCachesFactory import org.jetbrains.kotlin.fir.caches.firCachesFactory
import org.jetbrains.kotlin.fir.caches.getValue import org.jetbrains.kotlin.fir.caches.getValue
import org.jetbrains.kotlin.fir.declarations.FirDeclarationOrigin import org.jetbrains.kotlin.fir.declarations.FirDeclarationOrigin
import org.jetbrains.kotlin.fir.deserialization.SingleModuleDataProvider import org.jetbrains.kotlin.fir.deserialization.SingleModuleDataProvider
import org.jetbrains.kotlin.fir.java.deserialization.JvmClassFileBasedSymbolProvider
import org.jetbrains.kotlin.fir.java.deserialization.KotlinBuiltins import org.jetbrains.kotlin.fir.java.deserialization.KotlinBuiltins
import org.jetbrains.kotlin.fir.realPsi import org.jetbrains.kotlin.fir.realPsi
import org.jetbrains.kotlin.fir.resolve.providers.FirSymbolProvider
import org.jetbrains.kotlin.fir.resolve.providers.FirSymbolProviderInternals import org.jetbrains.kotlin.fir.resolve.providers.FirSymbolProviderInternals
import org.jetbrains.kotlin.fir.scopes.FirKotlinScopeProvider import org.jetbrains.kotlin.fir.scopes.FirKotlinScopeProvider
import org.jetbrains.kotlin.fir.symbols.impl.* import org.jetbrains.kotlin.fir.symbols.impl.*
@@ -27,6 +27,7 @@ import org.jetbrains.kotlin.name.*
import org.jetbrains.kotlin.psi.* import org.jetbrains.kotlin.psi.*
import org.jetbrains.kotlin.resolve.jvm.JvmClassName import org.jetbrains.kotlin.resolve.jvm.JvmClassName
import org.jetbrains.kotlin.serialization.deserialization.MetadataPackageFragment import org.jetbrains.kotlin.serialization.deserialization.MetadataPackageFragment
import org.jetbrains.kotlin.utils.addToStdlib.ifNotEmpty
typealias DeserializedTypeAliasPostProcessor = (FirTypeAliasSymbol) -> Unit typealias DeserializedTypeAliasPostProcessor = (FirTypeAliasSymbol) -> Unit
@@ -49,16 +50,11 @@ internal open class JvmStubBasedFirDeserializedSymbolProvider(
project: Project, project: Project,
scope: GlobalSearchScope, scope: GlobalSearchScope,
private val initialOrigin: FirDeclarationOrigin private val initialOrigin: FirDeclarationOrigin
) : FirSymbolProvider(session) { ) : LLFirKotlinSymbolProviderWithNameCache(session) {
private val declarationProvider by lazy(LazyThreadSafetyMode.PUBLICATION) { project.createDeclarationProvider(scope, module = null) } private val declarationProvider by lazy(LazyThreadSafetyMode.PUBLICATION) { project.createDeclarationProvider(scope, module = null) }
private val moduleData = moduleDataProvider.getModuleData(null) private val moduleData = moduleDataProvider.getModuleData(null)
private val namesByPackageCache by lazy(LazyThreadSafetyMode.PUBLICATION) { override val symbolNameCache: LLFirKotlinSymbolProviderNameCache = LLFirKotlinSymbolProviderNameCache(session, declarationProvider)
LLFirKotlinSymbolProviderNameCache(
session,
declarationProvider
)
}
private val typeAliasCache: FirCache<ClassId, FirTypeAliasSymbol?, StubBasedFirDeserializationContext?> = private val typeAliasCache: FirCache<ClassId, FirTypeAliasSymbol?, StubBasedFirDeserializationContext?> =
session.firCachesFactory.createCacheWithPostCompute( session.firCachesFactory.createCacheWithPostCompute(
@@ -69,6 +65,7 @@ internal open class JvmStubBasedFirDeserializedSymbolProvider(
} }
} }
) )
private val classCache: FirCache<ClassId, FirRegularClassSymbol?, StubBasedFirDeserializationContext?> = private val classCache: FirCache<ClassId, FirRegularClassSymbol?, StubBasedFirDeserializationContext?> =
session.firCachesFactory.createCache( session.firCachesFactory.createCache(
createValue = { classId, context -> findAndDeserializeClass(classId, context) } createValue = { classId, context -> findAndDeserializeClass(classId, context) }
@@ -77,23 +74,12 @@ internal open class JvmStubBasedFirDeserializedSymbolProvider(
private val functionCache = session.firCachesFactory.createCache(::loadFunctionsByCallableId) private val functionCache = session.firCachesFactory.createCache(::loadFunctionsByCallableId)
private val propertyCache = session.firCachesFactory.createCache(::loadPropertiesByCallableId) private val propertyCache = session.firCachesFactory.createCache(::loadPropertiesByCallableId)
override fun computePackageSetWithTopLevelCallables(): Set<String>? {
return namesByPackageCache.getPackageNamesWithTopLevelCallables()
}
override fun computeCallableNamesInPackage(packageFqName: FqName): Set<Name>? =
namesByPackageCache.getTopLevelCallableNamesInPackage(packageFqName)
override fun knownTopLevelClassifiersInPackage(packageFqName: FqName): Set<String>? {
return namesByPackageCache.getTopLevelClassifierNamesInPackage(packageFqName)
}
private fun findAndDeserializeTypeAlias( private fun findAndDeserializeTypeAlias(
classId: ClassId, classId: ClassId,
context: StubBasedFirDeserializationContext? context: StubBasedFirDeserializationContext?,
): Pair<FirTypeAliasSymbol?, DeserializedTypeAliasPostProcessor?> { ): Pair<FirTypeAliasSymbol?, DeserializedTypeAliasPostProcessor?> {
val classLikeDeclaration = val classLikeDeclaration =
context?.classLikeDeclaration ?: declarationProvider.getClassLikeDeclarationByClassId(classId)?.originalElement (context?.classLikeDeclaration ?: declarationProvider.getClassLikeDeclarationByClassId(classId))?.originalElement
if (classLikeDeclaration is KtTypeAlias) { if (classLikeDeclaration is KtTypeAlias) {
val symbol = FirTypeAliasSymbol(classId) val symbol = FirTypeAliasSymbol(classId)
val postProcessor: DeserializedTypeAliasPostProcessor = { val postProcessor: DeserializedTypeAliasPostProcessor = {
@@ -114,14 +100,15 @@ internal open class JvmStubBasedFirDeserializedSymbolProvider(
private fun findAndDeserializeClass( private fun findAndDeserializeClass(
classId: ClassId, classId: ClassId,
parentContext: StubBasedFirDeserializationContext? = null parentContext: StubBasedFirDeserializationContext?,
): FirRegularClassSymbol? { ): FirRegularClassSymbol? {
val (classLikeDeclaration, context) = val (classLikeDeclaration, context) =
if (parentContext?.classLikeDeclaration != null) { if (parentContext?.classLikeDeclaration != null) {
parentContext.classLikeDeclaration to null parentContext.classLikeDeclaration.originalElement to null
} else { } else {
(declarationProvider.getClassLikeDeclarationByClassId(classId)?.originalElement ?: return null) to parentContext (declarationProvider.getClassLikeDeclarationByClassId(classId)?.originalElement ?: return null) to parentContext
} }
val symbol = FirRegularClassSymbol(classId) val symbol = FirRegularClassSymbol(classId)
if (classLikeDeclaration is KtClassOrObject) { if (classLikeDeclaration is KtClassOrObject) {
deserializeClassToSymbol( deserializeClassToSymbol(
@@ -146,8 +133,11 @@ internal open class JvmStubBasedFirDeserializedSymbolProvider(
return null return null
} }
private fun loadFunctionsByCallableId(callableId: CallableId): List<FirNamedFunctionSymbol> { private fun loadFunctionsByCallableId(
val topLevelFunctions = declarationProvider.getTopLevelFunctions(callableId) callableId: CallableId,
foundFunctions: Collection<KtNamedFunction>?,
): List<FirNamedFunctionSymbol> {
val topLevelFunctions = foundFunctions ?: declarationProvider.getTopLevelFunctions(callableId)
val origins = if (topLevelFunctions.size > 1) mutableSetOf<KtNamedFunction>() else null val origins = if (topLevelFunctions.size > 1) mutableSetOf<KtNamedFunction>() else null
return topLevelFunctions return topLevelFunctions
.mapNotNull { function -> .mapNotNull { function ->
@@ -166,8 +156,8 @@ internal open class JvmStubBasedFirDeserializedSymbolProvider(
} }
} }
private fun loadPropertiesByCallableId(callableId: CallableId): List<FirPropertySymbol> { private fun loadPropertiesByCallableId(callableId: CallableId, foundProperties: Collection<KtProperty>?): List<FirPropertySymbol> {
val topLevelProperties = declarationProvider.getTopLevelProperties(callableId) val topLevelProperties = foundProperties ?: declarationProvider.getTopLevelProperties(callableId)
val origins = if (topLevelProperties.size > 1) mutableSetOf<KtProperty>() else null val origins = if (topLevelProperties.size > 1) mutableSetOf<KtProperty>() else null
return topLevelProperties return topLevelProperties
.mapNotNull { property -> .mapNotNull { property ->
@@ -187,45 +177,81 @@ internal open class JvmStubBasedFirDeserializedSymbolProvider(
return classCache.getValue(classId, parentContext) return classCache.getValue(classId, parentContext)
} }
private fun getTypeAlias(classId: ClassId): FirTypeAliasSymbol? { private fun getTypeAlias(classId: ClassId, context: StubBasedFirDeserializationContext? = null): FirTypeAliasSymbol? {
if (!classId.relativeClassName.isOneSegmentFQN()) return null if (!classId.relativeClassName.isOneSegmentFQN()) return null
return typeAliasCache.getValue(classId) return typeAliasCache.getValue(classId, context)
} }
@FirSymbolProviderInternals @FirSymbolProviderInternals
override fun getTopLevelCallableSymbolsTo(destination: MutableList<FirCallableSymbol<*>>, packageFqName: FqName, name: Name) { override fun getTopLevelCallableSymbolsTo(destination: MutableList<FirCallableSymbol<*>>, packageFqName: FqName, name: Name) {
val callableId = CallableId(packageFqName, name) val callableId = CallableId(packageFqName, name)
destination += functionCache.getCallables(callableId) destination += functionCache.getCallablesWithoutContext(callableId)
destination += propertyCache.getCallables(callableId) destination += propertyCache.getCallablesWithoutContext(callableId)
} }
private fun <C : FirCallableSymbol<*>> FirCache<CallableId, List<C>, Nothing?>.getCallables(id: CallableId): List<C> { private fun <C : FirCallableSymbol<*>, CONTEXT> FirCache<CallableId, List<C>, CONTEXT?>.getCallablesWithoutContext(
if (!namesByPackageCache.mayHaveTopLevelCallable(id.packageName, id.callableName)) return emptyList() id: CallableId,
return getValue(id) ): List<C> {
if (!symbolNameCache.mayHaveTopLevelCallable(id.packageName, id.callableName)) return emptyList()
return getValue(id, null)
}
@FirSymbolProviderInternals
override fun getTopLevelCallableSymbolsTo(
destination: MutableList<FirCallableSymbol<*>>,
callableId: CallableId,
callables: Collection<KtCallableDeclaration>,
) {
callables.filterIsInstance<KtNamedFunction>().ifNotEmpty {
destination += functionCache.getValue(callableId, this)
}
callables.filterIsInstance<KtProperty>().ifNotEmpty {
destination += propertyCache.getValue(callableId, this)
}
} }
@FirSymbolProviderInternals @FirSymbolProviderInternals
override fun getTopLevelFunctionSymbolsTo(destination: MutableList<FirNamedFunctionSymbol>, packageFqName: FqName, name: Name) { override fun getTopLevelFunctionSymbolsTo(destination: MutableList<FirNamedFunctionSymbol>, packageFqName: FqName, name: Name) {
destination += functionCache.getCallables(CallableId(packageFqName, name)) destination += functionCache.getCallablesWithoutContext(CallableId(packageFqName, name))
}
@FirSymbolProviderInternals
override fun getTopLevelFunctionSymbolsTo(
destination: MutableList<FirNamedFunctionSymbol>,
callableId: CallableId,
functions: Collection<KtNamedFunction>,
) {
destination += functionCache.getValue(callableId, functions)
} }
@FirSymbolProviderInternals @FirSymbolProviderInternals
override fun getTopLevelPropertySymbolsTo(destination: MutableList<FirPropertySymbol>, packageFqName: FqName, name: Name) { override fun getTopLevelPropertySymbolsTo(destination: MutableList<FirPropertySymbol>, packageFqName: FqName, name: Name) {
destination += propertyCache.getCallables(CallableId(packageFqName, name)) destination += propertyCache.getCallablesWithoutContext(CallableId(packageFqName, name))
}
@FirSymbolProviderInternals
override fun getTopLevelPropertySymbolsTo(
destination: MutableList<FirPropertySymbol>,
callableId: CallableId,
properties: Collection<KtProperty>,
) {
destination += propertyCache.getValue(callableId, properties)
} }
override fun getPackage(fqName: FqName): FqName? = override fun getPackage(fqName: FqName): FqName? =
fqName.takeIf { fqName.takeIf {
namesByPackageCache.getTopLevelClassifierNamesInPackage(fqName)?.isNotEmpty() == true || symbolNameCache.getTopLevelClassifierNamesInPackage(fqName)?.isNotEmpty() == true ||
namesByPackageCache.getPackageNamesWithTopLevelCallables()?.contains(fqName.asString()) == true symbolNameCache.getPackageNamesWithTopLevelCallables()?.contains(fqName.asString()) == true
} }
override fun getClassLikeSymbolByClassId(classId: ClassId): FirClassLikeSymbol<*>? { override fun getClassLikeSymbolByClassId(classId: ClassId): FirClassLikeSymbol<*>? {
if (!namesByPackageCache.mayHaveTopLevelClassifier(classId, mayHaveFunctionClass = false)) return null if (!symbolNameCache.mayHaveTopLevelClassifier(classId, mayHaveFunctionClass = false)) return null
return getClass(classId) ?: getTypeAlias(classId) return getClass(classId) ?: getTypeAlias(classId)
} }
fun getClassLikeSymbolByClassId(classLikeDeclaration: KtClassLikeDeclaration, classId: ClassId): FirClassLikeSymbol<*>? { @FirSymbolProviderInternals
override fun getClassLikeSymbolByClassId(classId: ClassId, classLikeDeclaration: KtClassLikeDeclaration): FirClassLikeSymbol<*>? {
val annotationDeserializer = StubBasedAnnotationDeserializer(session) val annotationDeserializer = StubBasedAnnotationDeserializer(session)
val deserializationContext = StubBasedFirDeserializationContext( val deserializationContext = StubBasedFirDeserializationContext(
moduleData, moduleData,
@@ -256,9 +282,9 @@ internal open class JvmStubBasedFirDeserializedSymbolProvider(
} }
fun getTopLevelCallableSymbol( fun getTopLevelCallableSymbol(
callableDeclaration: KtCallableDeclaration,
packageFqName: FqName, packageFqName: FqName,
shortName: Name shortName: Name,
callableDeclaration: KtCallableDeclaration,
): FirCallableSymbol<*>? { ): FirCallableSymbol<*>? {
//possible overloads spoils here //possible overloads spoils here
//we can't use only this callable instead of index access to fill the cache //we can't use only this callable instead of index access to fill the cache
@@ -75,7 +75,7 @@ internal class FirDeclarationForCompiledElementSearcher(private val symbolProvid
val classCandidate = when (symbolProvider) { val classCandidate = when (symbolProvider) {
is LLFirModuleWithDependenciesSymbolProvider -> { is LLFirModuleWithDependenciesSymbolProvider -> {
symbolProvider.getClassLikeSymbolByFqNameWithoutDependencies(declaration, classId) symbolProvider.getDeserializedClassLikeSymbolByClassIdWithoutDependencies(classId, declaration)
?: symbolProvider.friendBuiltinsProvider?.getClassLikeSymbolByClassId(classId) ?: symbolProvider.friendBuiltinsProvider?.getClassLikeSymbolByClassId(classId)
} }
else -> { else -> {
@@ -170,7 +170,7 @@ private fun FirSymbolProvider.findCallableCandidates(
@OptIn(FirSymbolProviderInternals::class) @OptIn(FirSymbolProviderInternals::class)
return when (this) { return when (this) {
is LLFirModuleWithDependenciesSymbolProvider -> buildList { is LLFirModuleWithDependenciesSymbolProvider -> buildList {
getTopLevelDeserializedCallableSymbolsToWithoutDependencies(this, declaration, packageFqName, shortName) getTopLevelDeserializedCallableSymbolsToWithoutDependencies(this, packageFqName, shortName, declaration)
friendBuiltinsProvider?.getTopLevelCallableSymbolsTo(this, packageFqName, shortName) friendBuiltinsProvider?.getTopLevelCallableSymbolsTo(this, packageFqName, shortName)
} }
else -> getTopLevelCallableSymbols(packageFqName, shortName) else -> getTopLevelCallableSymbols(packageFqName, shortName)