[LL API] Move modification tracker right inside a 'LLFirSession'

This commit is contained in:
Yan Zhulanow
2023-03-06 18:34:39 +09:00
committed by Space Team
parent be71d75f9e
commit 88636c8dbf
3 changed files with 91 additions and 101 deletions
@@ -6,13 +6,18 @@
package org.jetbrains.kotlin.analysis.low.level.api.fir.sessions package org.jetbrains.kotlin.analysis.low.level.api.fir.sessions
import com.intellij.openapi.project.Project import com.intellij.openapi.project.Project
import org.jetbrains.kotlin.analysis.project.structure.KtModule import com.intellij.openapi.util.ModificationTracker
import org.jetbrains.kotlin.analysis.low.level.api.fir.util.BooleanModificationTracker
import org.jetbrains.kotlin.analysis.project.structure.*
import org.jetbrains.kotlin.analysis.providers.KotlinModificationTrackerFactory
import org.jetbrains.kotlin.analysis.utils.trackers.CompositeModificationTracker
import org.jetbrains.kotlin.fir.BuiltinTypes import org.jetbrains.kotlin.fir.BuiltinTypes
import org.jetbrains.kotlin.fir.FirElementWithResolvePhase import org.jetbrains.kotlin.fir.FirElementWithResolvePhase
import org.jetbrains.kotlin.fir.FirSession import org.jetbrains.kotlin.fir.FirSession
import org.jetbrains.kotlin.fir.PrivateSessionConstructor import org.jetbrains.kotlin.fir.PrivateSessionConstructor
import org.jetbrains.kotlin.fir.resolve.ScopeSession import org.jetbrains.kotlin.fir.resolve.ScopeSession
import org.jetbrains.kotlin.fir.symbols.FirBasedSymbol import org.jetbrains.kotlin.fir.symbols.FirBasedSymbol
import java.util.concurrent.atomic.AtomicBoolean
@OptIn(PrivateSessionConstructor::class) @OptIn(PrivateSessionConstructor::class)
abstract class LLFirSession( abstract class LLFirSession(
@@ -20,10 +25,46 @@ abstract class LLFirSession(
override val builtinTypes: BuiltinTypes, override val builtinTypes: BuiltinTypes,
kind: Kind kind: Kind
) : FirSession(sessionProvider = null, kind) { ) : FirSession(sessionProvider = null, kind) {
abstract fun getScopeSession(): ScopeSession
val modificationTracker: ModificationTracker
private val initialModificationCount: Long
private val isExplicitlyInvalidated = AtomicBoolean(false)
val project: Project val project: Project
get() = ktModule.project get() = ktModule.project
abstract fun getScopeSession(): ScopeSession init {
val trackerFactory = KotlinModificationTrackerFactory.getService(ktModule.project)
val validityTracker = trackerFactory.createModuleStateTracker(ktModule)
val outOfBlockTracker = when (ktModule) {
is KtSourceModule -> trackerFactory.createModuleWithoutDependenciesOutOfBlockModificationTracker(ktModule)
is KtNotUnderContentRootModule -> ModificationTracker { ktModule.file?.modificationStamp ?: 0 }
is KtScriptModule -> ModificationTracker { ktModule.file.modificationStamp }
is KtScriptDependencyModule -> ModificationTracker { ktModule.file?.modificationStamp ?: 0 }
else -> ModificationTracker.NEVER_CHANGED
}
modificationTracker = CompositeModificationTracker.create(
listOf(
outOfBlockTracker,
ModificationTracker { validityTracker.rootModificationCount },
BooleanModificationTracker { validityTracker.isValid },
BooleanModificationTracker { !isExplicitlyInvalidated.get() }
)
)
initialModificationCount = modificationTracker.modificationCount
}
fun invalidate() {
isExplicitlyInvalidated.set(true)
}
val isValid: Boolean
get() = modificationTracker.modificationCount == initialModificationCount
} }
abstract class LLFirModuleSession( abstract class LLFirModuleSession(
@@ -6,16 +6,12 @@
package org.jetbrains.kotlin.analysis.low.level.api.fir.sessions package org.jetbrains.kotlin.analysis.low.level.api.fir.sessions
import com.intellij.openapi.project.Project import com.intellij.openapi.project.Project
import com.intellij.openapi.util.ModificationTracker
import org.jetbrains.kotlin.analysis.low.level.api.fir.LLFirGlobalResolveComponents import org.jetbrains.kotlin.analysis.low.level.api.fir.LLFirGlobalResolveComponents
import org.jetbrains.kotlin.analysis.low.level.api.fir.project.structure.LLFirBuiltinsSessionFactory import org.jetbrains.kotlin.analysis.low.level.api.fir.project.structure.LLFirBuiltinsSessionFactory
import org.jetbrains.kotlin.analysis.low.level.api.fir.project.structure.LLFirLibrarySessionFactory import org.jetbrains.kotlin.analysis.low.level.api.fir.project.structure.LLFirLibrarySessionFactory
import org.jetbrains.kotlin.analysis.low.level.api.fir.project.structure.llFirModuleData import org.jetbrains.kotlin.analysis.low.level.api.fir.project.structure.llFirModuleData
import org.jetbrains.kotlin.analysis.low.level.api.fir.util.addValueFor import org.jetbrains.kotlin.analysis.low.level.api.fir.util.addValueFor
import org.jetbrains.kotlin.analysis.project.structure.* import org.jetbrains.kotlin.analysis.project.structure.*
import org.jetbrains.kotlin.analysis.providers.KotlinModificationTrackerFactory
import org.jetbrains.kotlin.analysis.providers.KtModuleStateTracker
import org.jetbrains.kotlin.analysis.utils.trackers.CompositeModificationTracker
import java.lang.ref.SoftReference import java.lang.ref.SoftReference
class LLFirSessionProviderStorage(val project: Project) { class LLFirSessionProviderStorage(val project: Project) {
@@ -51,12 +47,11 @@ class LLFirSessionProviderStorage(val project: Project) {
else -> error("Unexpected ${useSiteKtModule::class.simpleName}") else -> error("Unexpected ${useSiteKtModule::class.simpleName}")
} }
private fun createSessionProviderForSourceSession( private fun createSessionProviderForSourceSession(
useSiteKtModule: KtSourceModule, useSiteKtModule: KtSourceModule,
configureSession: (LLFirSession.() -> Unit)? configureSession: (LLFirSession.() -> Unit)?
): LLFirSessionProvider { ): LLFirSessionProvider {
val (sessions, session) = sourceAsUseSiteSessionCache.withMappings(project) { mappings -> val (sessions, session) = sourceAsUseSiteSessionCache.withMappings { mappings ->
val sessions = mutableMapOf<KtModule, LLFirSession>().apply { putAll(mappings) } val sessions = mutableMapOf<KtModule, LLFirSession>().apply { putAll(mappings) }
val session = LLFirSessionFactory.createSourcesSession( val session = LLFirSessionFactory.createSourcesSession(
project, project,
@@ -77,7 +72,7 @@ class LLFirSessionProviderStorage(val project: Project) {
useSiteKtModule: KtModule, useSiteKtModule: KtModule,
configureSession: (LLFirSession.() -> Unit)? configureSession: (LLFirSession.() -> Unit)?
): LLFirSessionProvider { ): LLFirSessionProvider {
val (sessions, session) = libraryAsUseSiteSessionCache.withMappings(project) { mappings -> val (sessions, session) = libraryAsUseSiteSessionCache.withMappings { mappings ->
val sessions = mutableMapOf<KtModule, LLFirSession>().apply { putAll(mappings) } val sessions = mutableMapOf<KtModule, LLFirSession>().apply { putAll(mappings) }
val session = LLFirSessionFactory.createLibraryOrLibrarySourceResolvableSession( val session = LLFirSessionFactory.createLibraryOrLibrarySourceResolvableSession(
project, project,
@@ -97,7 +92,7 @@ class LLFirSessionProviderStorage(val project: Project) {
useSiteKtModule: KtScriptModule, useSiteKtModule: KtScriptModule,
configureSession: (LLFirSession.() -> Unit)? configureSession: (LLFirSession.() -> Unit)?
): LLFirSessionProvider { ): LLFirSessionProvider {
val (sessions, session) = sourceAsUseSiteSessionCache.withMappings(project) { mappings -> val (sessions, session) = sourceAsUseSiteSessionCache.withMappings { mappings ->
val sessions = mutableMapOf<KtModule, LLFirSession>().apply { putAll(mappings) } val sessions = mutableMapOf<KtModule, LLFirSession>().apply { putAll(mappings) }
val session = LLFirSessionFactory.createScriptSession( val session = LLFirSessionFactory.createScriptSession(
project, project,
@@ -116,7 +111,7 @@ class LLFirSessionProviderStorage(val project: Project) {
useSiteKtModule: KtNotUnderContentRootModule, useSiteKtModule: KtNotUnderContentRootModule,
configureSession: (LLFirSession.() -> Unit)? configureSession: (LLFirSession.() -> Unit)?
): LLFirSessionProvider { ): LLFirSessionProvider {
val (sessions, session) = notUnderContentRootSessionCache.withMappings(project) { mappings -> val (sessions, session) = notUnderContentRootSessionCache.withMappings { mappings ->
val sessions = mutableMapOf<KtModule, LLFirSession>().apply { putAll(mappings) } val sessions = mutableMapOf<KtModule, LLFirSession>().apply { putAll(mappings) }
val session = LLFirSessionFactory.createNotUnderContentRootResolvableSession( val session = LLFirSessionFactory.createNotUnderContentRootResolvableSession(
project, project,
@@ -133,111 +128,60 @@ class LLFirSessionProviderStorage(val project: Project) {
private class LLFirSessionsCache { private class LLFirSessionsCache {
@Volatile @Volatile
private var mappings: Map<KtModule, FirSessionWithModificationTracker> = emptyMap() private var mappings: Map<KtModule, SoftReference<LLFirSession>> = emptyMap()
val sessionInvalidator: LLFirSessionInvalidator = LLFirSessionInvalidator { session -> val sessionInvalidator: LLFirSessionInvalidator = LLFirSessionInvalidator { session ->
mappings[session.llFirModuleData.ktModule]?.invalidate() mappings[session.llFirModuleData.ktModule]?.get()?.invalidate()
} }
inline fun <R> withMappings( inline fun <R> withMappings(
project: Project,
action: (Map<KtModule, LLFirSession>) -> Pair<Map<KtModule, LLFirSession>, R> action: (Map<KtModule, LLFirSession>) -> Pair<Map<KtModule, LLFirSession>, R>
): Pair<Map<KtModule, LLFirSession>, R> { ): Pair<Map<KtModule, LLFirSession>, R> {
val (newMappings, result) = action(getSessions().mapValues { it.value }) val (newMappings, result) = action(getSessions().mapValues { it.value })
mappings = newMappings.mapValues { FirSessionWithModificationTracker(project, it.value) } mappings = newMappings.mapValues { SoftReference(it.value) }
return newMappings to result return newMappings to result
} }
private fun getSessions(): Map<KtModule, LLFirSession> = buildMap { private fun getSessions(): Map<KtModule, LLFirSession> = buildMap {
val sessions = mappings.values // Initially, all sessions are considered to be valid ('true').
val wasSessionInvalidated = sessions.associateWithTo(hashMapOf()) { false } val sessions = LinkedHashMap<LLFirSession, Boolean>().apply {
for (sessionRef in mappings.values) {
val reversedDependencies = sessions.reversedDependencies { session -> val session = sessionRef.get() ?: continue
if (session.validityTracker.isValid) { put(session, true)
session.ktModule.directRegularDependencies.mapNotNull { mappings[it] }
} else emptyList()
}
fun markAsInvalidWithDfs(session: FirSessionWithModificationTracker) {
if (wasSessionInvalidated.getValue(session)) {
// we already was in that branch
return
}
wasSessionInvalidated[session] = true
reversedDependencies[session]?.forEach { dependsOn ->
markAsInvalidWithDfs(dependsOn)
} }
} }
for (session in sessions) { val reversedDependencies = buildMap {
if (!session.isValid) { for (session in sessions.keys) {
markAsInvalidWithDfs(session) if (session.isValid) {
} val module = session.ktModule
} for (dependency in module.directRegularDependencies) {
return wasSessionInvalidated.entries addValueFor(dependency, module)
.mapNotNull { (sessionWithTracker, wasInvalidated) -> }
if (wasInvalidated) return@mapNotNull null
val firSession = sessionWithTracker.firSessionSoftReference.get() ?: return@mapNotNull null
sessionWithTracker.ktModule to firSession
}.toMap()
}
private fun <T> Collection<T>.reversedDependencies(getDependencies: (T) -> List<T>): Map<T, List<T>> {
val result = hashMapOf<T, MutableList<T>>()
forEach { from ->
getDependencies(from).forEach { to ->
result.addValueFor(to, from)
}
}
return result
}
}
private class FirSessionWithModificationTracker(
project: Project,
firSession: LLFirSession,
) {
val firSessionSoftReference: SoftReference<LLFirSession> = SoftReference(firSession)
val ktModule = firSession.llFirModuleData.ktModule
val validityTracker: KtModuleStateTracker
private val modificationTracker: ModificationTracker
init {
val trackerFactory = KotlinModificationTrackerFactory.getService(project)
validityTracker = trackerFactory.createModuleStateTracker(ktModule)
val outOfBlockTracker = when (ktModule) {
is KtSourceModule -> trackerFactory.createModuleWithoutDependenciesOutOfBlockModificationTracker(ktModule)
is KtNotUnderContentRootModule -> ModificationTracker { ktModule.file?.modificationStamp ?: 0 }
is KtScriptModule -> ModificationTracker { ktModule.file.modificationStamp }
is KtScriptDependencyModule -> ModificationTracker { ktModule.file?.modificationStamp ?: 0 }
else -> null
}
modificationTracker = CompositeModificationTracker.create(
listOfNotNull(
outOfBlockTracker,
object : ModificationTracker {
override fun getModificationCount() = validityTracker.rootModificationCount
} }
) }
) }
fun invalidateRecursively(session: LLFirSession) {
// Invalidate all dependent sessions only if we didn't that before
if (sessions.put(session, false) == true) {
for (dependentModule in reversedDependencies[session.ktModule].orEmpty()) {
val dependentSession = mappings[dependentModule]?.get() ?: continue
invalidateRecursively(dependentSession)
}
}
}
for (session in sessions.keys) {
if (session.isValid) continue
invalidateRecursively(session)
}
return buildMap {
for ((session, isValid) in sessions) {
if (!isValid) continue
put(session.ktModule, session)
}
}
} }
private val timeStamp = modificationTracker.modificationCount
@Volatile
private var isInvalidated = false
fun invalidate() {
isInvalidated = true
}
val isValid: Boolean
get() = validityTracker.isValid
&& !isInvalidated
&& modificationTracker.modificationCount == timeStamp
} }
@@ -6,6 +6,7 @@
package org.jetbrains.kotlin.analysis.low.level.api.fir.util package org.jetbrains.kotlin.analysis.low.level.api.fir.util
import com.intellij.openapi.progress.ProgressManager import com.intellij.openapi.progress.ProgressManager
import com.intellij.openapi.util.ModificationTracker
import org.jetbrains.kotlin.fir.FirElement import org.jetbrains.kotlin.fir.FirElement
import org.jetbrains.kotlin.fir.declarations.FirDeclaration import org.jetbrains.kotlin.fir.declarations.FirDeclaration
import org.jetbrains.kotlin.fir.diagnostics.FirDiagnosticHolder import org.jetbrains.kotlin.fir.diagnostics.FirDiagnosticHolder
@@ -65,3 +66,7 @@ internal fun KtDeclaration.isNonAnonymousClassOrObject() =
this is KtClassOrObject this is KtClassOrObject
&& !this.isObjectLiteral() && !this.isObjectLiteral()
internal fun BooleanModificationTracker(provider: () -> Boolean): ModificationTracker {
return ModificationTracker { if (provider()) 0 else 1 }
}