[FIR] Register type related extensions in libraries sessions

^KT-57140
This commit is contained in:
Dmitriy Novozhilov
2023-03-21 17:09:28 +02:00
committed by Space Team
parent 01fc84ee3a
commit 8ca7b32577
13 changed files with 61 additions and 22 deletions
@@ -75,12 +75,12 @@ fun <F> prepareJvmSessions(
sessionProvider, sessionProvider,
libraryList.moduleDataProvider, libraryList.moduleDataProvider,
projectEnvironment, projectEnvironment,
extensionRegistrars,
librariesScope, librariesScope,
projectEnvironment.getPackagePartProvider(librariesScope), projectEnvironment.getPackagePartProvider(librariesScope),
configuration.languageVersionSettings, configuration.languageVersionSettings,
registerExtraComponents = {}, registerExtraComponents = {},
) )
} }
) { moduleFiles, moduleData, sessionProvider, sessionConfigurator -> ) { moduleFiles, moduleData, sessionProvider, sessionConfigurator ->
FirJvmSessionFactory.createModuleBasedSession( FirJvmSessionFactory.createModuleBasedSession(
@@ -126,6 +126,7 @@ fun <F> prepareJsSessions(
resolvedLibraries, resolvedLibraries,
sessionProvider, sessionProvider,
libraryList.moduleDataProvider, libraryList.moduleDataProvider,
extensionRegistrars,
configuration.languageVersionSettings, configuration.languageVersionSettings,
registerExtraComponents = {}, registerExtraComponents = {},
) )
@@ -168,6 +169,7 @@ fun <F> prepareNativeSessions(
resolvedLibraries, resolvedLibraries,
sessionProvider, sessionProvider,
libraryList.moduleDataProvider, libraryList.moduleDataProvider,
extensionRegistrars,
configuration.languageVersionSettings, configuration.languageVersionSettings,
registerExtraComponents = {}, registerExtraComponents = {},
) )
@@ -211,6 +213,7 @@ fun <F> prepareCommonSessions(
sessionProvider, sessionProvider,
libraryList.moduleDataProvider, libraryList.moduleDataProvider,
projectEnvironment, projectEnvironment,
extensionRegistrars,
librariesScope, librariesScope,
resolvedLibraries, resolvedLibraries,
projectEnvironment.getPackagePartProvider(librariesScope) as PackageAndMetadataPartProvider, projectEnvironment.getPackagePartProvider(librariesScope) as PackageAndMetadataPartProvider,
@@ -35,6 +35,11 @@ abstract class FirExtensionRegistrar : FirExtensionRegistrarAdapter() {
FirFunctionTypeKindExtension::class, FirFunctionTypeKindExtension::class,
FirDeclarationsForMetadataProviderExtension::class, FirDeclarationsForMetadataProviderExtension::class,
) )
internal val ALLOWED_EXTENSIONS_FOR_LIBRARY_SESSION = listOf(
FirTypeAttributeExtension::class,
FirFunctionTypeKindExtension::class,
)
} }
protected abstract fun ExtensionRegistrarContext.configurePlugin() protected abstract fun ExtensionRegistrarContext.configurePlugin()
@@ -34,6 +34,7 @@ abstract class FirAbstractSessionFactory {
sessionProvider: FirProjectSessionProvider, sessionProvider: FirProjectSessionProvider,
moduleDataProvider: ModuleDataProvider, moduleDataProvider: ModuleDataProvider,
languageVersionSettings: LanguageVersionSettings, languageVersionSettings: LanguageVersionSettings,
extensionRegistrars: List<FirExtensionRegistrar>,
registerExtraComponents: ((FirSession) -> Unit), registerExtraComponents: ((FirSession) -> Unit),
createKotlinScopeProvider: () -> FirKotlinScopeProvider, createKotlinScopeProvider: () -> FirKotlinScopeProvider,
createProviders: (FirSession, FirModuleData, FirKotlinScopeProvider) -> List<FirSymbolProvider> createProviders: (FirSession, FirModuleData, FirKotlinScopeProvider) -> List<FirSymbolProvider>
@@ -59,6 +60,12 @@ abstract class FirAbstractSessionFactory {
) )
builtinsModuleData.bindSession(this) builtinsModuleData.bindSession(this)
FirSessionConfigurator(this).apply {
for (extensionRegistrar in extensionRegistrars) {
registerExtensions(extensionRegistrar.configure())
}
}.configure()
val providers = createProviders(this, builtinsModuleData, kotlinScopeProvider) val providers = createProviders(this, builtinsModuleData, kotlinScopeProvider)
val symbolProvider = FirCachingCompositeSymbolProvider(this, providers) val symbolProvider = FirCachingCompositeSymbolProvider(this, providers)
@@ -37,6 +37,7 @@ object FirCommonSessionFactory : FirAbstractSessionFactory() {
sessionProvider: FirProjectSessionProvider, sessionProvider: FirProjectSessionProvider,
moduleDataProvider: ModuleDataProvider, moduleDataProvider: ModuleDataProvider,
projectEnvironment: AbstractProjectEnvironment, projectEnvironment: AbstractProjectEnvironment,
extensionRegistrars: List<FirExtensionRegistrar>,
librariesScope: AbstractProjectFileSearchScope, librariesScope: AbstractProjectFileSearchScope,
resolvedKLibs: List<KotlinResolvedLibrary>, resolvedKLibs: List<KotlinResolvedLibrary>,
packageAndMetadataPartProvider: PackageAndMetadataPartProvider, packageAndMetadataPartProvider: PackageAndMetadataPartProvider,
@@ -48,6 +49,7 @@ object FirCommonSessionFactory : FirAbstractSessionFactory() {
sessionProvider, sessionProvider,
moduleDataProvider, moduleDataProvider,
languageVersionSettings, languageVersionSettings,
extensionRegistrars,
registerExtraComponents = { registerExtraComponents = {
registerExtraComponents(it) registerExtraComponents(it)
}, },
@@ -66,13 +66,15 @@ object FirJsSessionFactory : FirAbstractSessionFactory() {
resolvedLibraries: List<KotlinLibrary>, resolvedLibraries: List<KotlinLibrary>,
sessionProvider: FirProjectSessionProvider, sessionProvider: FirProjectSessionProvider,
moduleDataProvider: ModuleDataProvider, moduleDataProvider: ModuleDataProvider,
extensionRegistrars: List<FirExtensionRegistrar>,
languageVersionSettings: LanguageVersionSettings = LanguageVersionSettingsImpl.DEFAULT, languageVersionSettings: LanguageVersionSettings = LanguageVersionSettingsImpl.DEFAULT,
registerExtraComponents: ((FirSession) -> Unit), registerExtraComponents: ((FirSession) -> Unit),
) = createLibrarySession( ): FirSession = createLibrarySession(
mainModuleName, mainModuleName,
sessionProvider, sessionProvider,
moduleDataProvider, moduleDataProvider,
languageVersionSettings, languageVersionSettings,
extensionRegistrars,
registerExtraComponents = { registerExtraComponents = {
it.registerJsSpecificResolveComponents() it.registerJsSpecificResolveComponents()
registerExtraComponents(it) registerExtraComponents(it)
@@ -83,7 +85,7 @@ object FirJsSessionFactory : FirAbstractSessionFactory() {
KlibBasedSymbolProvider(session, moduleDataProvider, kotlinScopeProvider, resolvedLibraries), KlibBasedSymbolProvider(session, moduleDataProvider, kotlinScopeProvider, resolvedLibraries),
// (Most) builtins should be taken from the dependencies in JS compilation, therefore builtins provider is the last one // (Most) builtins should be taken from the dependencies in JS compilation, therefore builtins provider is the last one
// TODO: consider using "poisoning" provider for builtins to ensure that proper ones are taken from dependencies // TODO: consider using "poisoning" provider for builtins to ensure that proper ones are taken from dependencies
// NOTE: it requires precise filtering for true nuiltins, like Function* // NOTE: it requires precise filtering for true builtins, like Function*
FirBuiltinSymbolProvider(session, builtinsModuleData, kotlinScopeProvider), FirBuiltinSymbolProvider(session, builtinsModuleData, kotlinScopeProvider),
FirExtensionSyntheticFunctionInterfaceProvider.createIfNeeded(session, builtinsModuleData, kotlinScopeProvider), FirExtensionSyntheticFunctionInterfaceProvider.createIfNeeded(session, builtinsModuleData, kotlinScopeProvider),
) )
@@ -37,6 +37,7 @@ object FirJvmSessionFactory : FirAbstractSessionFactory() {
sessionProvider: FirProjectSessionProvider, sessionProvider: FirProjectSessionProvider,
moduleDataProvider: ModuleDataProvider, moduleDataProvider: ModuleDataProvider,
projectEnvironment: AbstractProjectEnvironment, projectEnvironment: AbstractProjectEnvironment,
extensionRegistrars: List<FirExtensionRegistrar>,
scope: AbstractProjectFileSearchScope, scope: AbstractProjectFileSearchScope,
packagePartProvider: PackagePartProvider, packagePartProvider: PackagePartProvider,
languageVersionSettings: LanguageVersionSettings, languageVersionSettings: LanguageVersionSettings,
@@ -47,6 +48,7 @@ object FirJvmSessionFactory : FirAbstractSessionFactory() {
sessionProvider, sessionProvider,
moduleDataProvider, moduleDataProvider,
languageVersionSettings, languageVersionSettings,
extensionRegistrars,
registerExtraComponents = { registerExtraComponents = {
it.registerCommonJavaComponents(projectEnvironment.getJavaModuleResolver()) it.registerCommonJavaComponents(projectEnvironment.getJavaModuleResolver())
registerExtraComponents(it) registerExtraComponents(it)
@@ -27,6 +27,7 @@ object FirNativeSessionFactory : FirAbstractSessionFactory() {
resolvedLibraries: List<KotlinResolvedLibrary>, resolvedLibraries: List<KotlinResolvedLibrary>,
sessionProvider: FirProjectSessionProvider, sessionProvider: FirProjectSessionProvider,
moduleDataProvider: ModuleDataProvider, moduleDataProvider: ModuleDataProvider,
extensionRegistrars: List<FirExtensionRegistrar>,
languageVersionSettings: LanguageVersionSettings, languageVersionSettings: LanguageVersionSettings,
registerExtraComponents: ((FirSession) -> Unit) = {}, registerExtraComponents: ((FirSession) -> Unit) = {},
): FirSession { ): FirSession {
@@ -35,6 +36,7 @@ object FirNativeSessionFactory : FirAbstractSessionFactory() {
sessionProvider, sessionProvider,
moduleDataProvider, moduleDataProvider,
languageVersionSettings, languageVersionSettings,
extensionRegistrars,
registerExtraComponents, registerExtraComponents,
createKotlinScopeProvider = { FirKotlinScopeProvider { _, declaredMemberScope, _, _, _ -> declaredMemberScope } }, createKotlinScopeProvider = { FirKotlinScopeProvider { _, declaredMemberScope, _, _, _ -> declaredMemberScope } },
createProviders = { session, builtinsModuleData, kotlinScopeProvider -> createProviders = { session, builtinsModuleData, kotlinScopeProvider ->
@@ -12,9 +12,7 @@ import org.jetbrains.kotlin.fir.analysis.checkers.expression.ExpressionCheckers
import org.jetbrains.kotlin.fir.analysis.checkers.type.TypeCheckers import org.jetbrains.kotlin.fir.analysis.checkers.type.TypeCheckers
import org.jetbrains.kotlin.fir.analysis.checkersComponent import org.jetbrains.kotlin.fir.analysis.checkersComponent
import org.jetbrains.kotlin.fir.analysis.extensions.additionalCheckers import org.jetbrains.kotlin.fir.analysis.extensions.additionalCheckers
import org.jetbrains.kotlin.fir.extensions.BunchOfRegisteredExtensions import org.jetbrains.kotlin.fir.extensions.*
import org.jetbrains.kotlin.fir.extensions.extensionService
import org.jetbrains.kotlin.fir.extensions.registerExtensions
class FirSessionConfigurator(private val session: FirSession) { class FirSessionConfigurator(private val session: FirSession) {
private val registeredExtensions: MutableList<BunchOfRegisteredExtensions> = mutableListOf(BunchOfRegisteredExtensions.empty()) private val registeredExtensions: MutableList<BunchOfRegisteredExtensions> = mutableListOf(BunchOfRegisteredExtensions.empty())
@@ -38,9 +36,17 @@ class FirSessionConfigurator(private val session: FirSession) {
session.checkersComponent.register(checkers) session.checkersComponent.register(checkers)
} }
@OptIn(PluginServicesInitialization::class)
@SessionConfiguration @SessionConfiguration
fun configure() { fun configure() {
session.extensionService.registerExtensions(registeredExtensions.reduce(BunchOfRegisteredExtensions::plus)) var extensions = registeredExtensions.reduce(BunchOfRegisteredExtensions::plus)
session.extensionService.additionalCheckers.forEach(session.checkersComponent::register) if (session.kind == FirSession.Kind.Library) {
val filteredExtensions = extensions.extensions.filterKeys { it in FirExtensionRegistrar.ALLOWED_EXTENSIONS_FOR_LIBRARY_SESSION }
extensions = BunchOfRegisteredExtensions(filteredExtensions)
}
session.extensionService.registerExtensions(extensions)
if (session.kind == FirSession.Kind.Source) {
session.extensionService.additionalCheckers.forEach(session.checkersComponent::register)
}
} }
} }
@@ -51,6 +51,7 @@ object FirSessionFactoryHelper {
sessionProvider, sessionProvider,
dependencyList.moduleDataProvider, dependencyList.moduleDataProvider,
projectEnvironment, projectEnvironment,
extensionRegistrars,
librariesScope, librariesScope,
packagePartProvider, packagePartProvider,
languageVersionSettings, languageVersionSettings,
@@ -129,4 +130,4 @@ object FirSessionFactoryHelper {
register(FirOverridesBackwardCompatibilityHelper::class, FirOverridesBackwardCompatibilityHelper.Default()) register(FirOverridesBackwardCompatibilityHelper::class, FirOverridesBackwardCompatibilityHelper.Default())
register(FirEnumEntriesSupport::class, FirEnumEntriesSupport(this)) register(FirEnumEntriesSupport::class, FirEnumEntriesSupport(this))
} }
} }
@@ -88,18 +88,21 @@ open class FirFrontendFacade(
val (moduleDataMap, moduleDataProvider) = initializeModuleData(sortedModules) val (moduleDataMap, moduleDataProvider) = initializeModuleData(sortedModules)
val project = testServices.compilerConfigurationProvider.getProject(module)
val extensionRegistrars = FirExtensionRegistrar.getInstances(project)
val projectEnvironment = createLibrarySession( val projectEnvironment = createLibrarySession(
module, module,
testServices.compilerConfigurationProvider.getProject(module), project,
Name.special("<${module.name}>"), Name.special("<${module.name}>"),
testServices.firModuleInfoProvider.firSessionProvider, testServices.firModuleInfoProvider.firSessionProvider,
moduleDataProvider, moduleDataProvider,
testServices.compilerConfigurationProvider.getCompilerConfiguration(module) testServices.compilerConfigurationProvider.getCompilerConfiguration(module),
extensionRegistrars
) )
val targetPlatform = module.targetPlatform val targetPlatform = module.targetPlatform
val firOutputPartForDependsOnModules = sortedModules.map { val firOutputPartForDependsOnModules = sortedModules.map {
analyze(it, moduleDataMap[it]!!, targetPlatform, projectEnvironment) analyze(it, moduleDataMap[it]!!, targetPlatform, projectEnvironment, extensionRegistrars)
} }
return FirOutputArtifactImpl(firOutputPartForDependsOnModules) return FirOutputArtifactImpl(firOutputPartForDependsOnModules)
@@ -182,7 +185,8 @@ open class FirFrontendFacade(
moduleName: Name, moduleName: Name,
sessionProvider: FirProjectSessionProvider, sessionProvider: FirProjectSessionProvider,
moduleDataProvider: ModuleDataProvider, moduleDataProvider: ModuleDataProvider,
configuration: CompilerConfiguration configuration: CompilerConfiguration,
extensionRegistrars: List<FirExtensionRegistrar>
): AbstractProjectEnvironment? { ): AbstractProjectEnvironment? {
val compilerConfigurationProvider = testServices.compilerConfigurationProvider val compilerConfigurationProvider = testServices.compilerConfigurationProvider
val projectEnvironment: AbstractProjectEnvironment? val projectEnvironment: AbstractProjectEnvironment?
@@ -202,6 +206,7 @@ open class FirFrontendFacade(
sessionProvider, sessionProvider,
moduleDataProvider, moduleDataProvider,
projectEnvironment, projectEnvironment,
extensionRegistrars,
projectFileSearchScope, projectFileSearchScope,
packagePartProvider, packagePartProvider,
languageVersionSettings, languageVersionSettings,
@@ -217,6 +222,7 @@ open class FirFrontendFacade(
module, module,
testServices, testServices,
configuration, configuration,
extensionRegistrars,
languageVersionSettings, languageVersionSettings,
registerExtraComponents = ::registerExtraComponents, registerExtraComponents = ::registerExtraComponents,
) )
@@ -228,6 +234,7 @@ open class FirFrontendFacade(
listOf(), listOf(),
sessionProvider, sessionProvider,
moduleDataProvider, moduleDataProvider,
extensionRegistrars,
languageVersionSettings, languageVersionSettings,
registerExtraComponents = ::registerExtraComponents, registerExtraComponents = ::registerExtraComponents,
) )
@@ -241,7 +248,8 @@ open class FirFrontendFacade(
module: TestModule, module: TestModule,
moduleData: FirModuleData, moduleData: FirModuleData,
targetPlatform: TargetPlatform, targetPlatform: TargetPlatform,
projectEnvironment: AbstractProjectEnvironment? projectEnvironment: AbstractProjectEnvironment?,
extensionRegistrars: List<FirExtensionRegistrar>,
): FirOutputPartForDependsOnModule { ): FirOutputPartForDependsOnModule {
val compilerConfigurationProvider = testServices.compilerConfigurationProvider val compilerConfigurationProvider = testServices.compilerConfigurationProvider
val moduleInfoProvider = testServices.firModuleInfoProvider val moduleInfoProvider = testServices.firModuleInfoProvider
@@ -258,7 +266,6 @@ open class FirFrontendFacade(
FirParser.Psi -> testServices.sourceFileProvider.getKtFilesForSourceFiles(module.files, project).values to emptyList() FirParser.Psi -> testServices.sourceFileProvider.getKtFilesForSourceFiles(module.files, project).values to emptyList()
} }
val extensionRegistrars = FirExtensionRegistrar.getInstances(project)
val sessionConfigurator: FirSessionConfigurator.() -> Unit = { val sessionConfigurator: FirSessionConfigurator.() -> Unit = {
if (FirDiagnosticsDirectives.WITH_EXTENDED_CHECKERS in module.directives) { if (FirDiagnosticsDirectives.WITH_EXTENDED_CHECKERS in module.directives) {
registerExtendedCommonCheckers() registerExtendedCommonCheckers()
@@ -33,6 +33,7 @@ object TestFirJsSessionFactory {
module: TestModule, module: TestModule,
testServices: TestServices, testServices: TestServices,
configuration: CompilerConfiguration, configuration: CompilerConfiguration,
extensionRegistrars: List<FirExtensionRegistrar>,
languageVersionSettings: LanguageVersionSettings, languageVersionSettings: LanguageVersionSettings,
registerExtraComponents: ((FirSession) -> Unit), registerExtraComponents: ((FirSession) -> Unit),
): FirSession { ): FirSession {
@@ -45,6 +46,7 @@ object TestFirJsSessionFactory {
resolvedLibraries.map { it.library }, resolvedLibraries.map { it.library },
sessionProvider, sessionProvider,
moduleDataProvider, moduleDataProvider,
extensionRegistrars,
languageVersionSettings, languageVersionSettings,
registerExtraComponents, registerExtraComponents,
) )
@@ -11,13 +11,13 @@ FILE: dependencyWithoutAttributePlugin.kt
R|org/jetbrains/kotlin/fir/plugin/consumePositiveInt|(R|<local>/someInt|) R|org/jetbrains/kotlin/fir/plugin/consumePositiveInt|(R|<local>/someInt|)
} }
public final fun test_2(): R|kotlin/Unit| { public final fun test_2(): R|kotlin/Unit| {
lval x: R|@R|org/jetbrains/kotlin/fir/plugin/Positive|() kotlin/Int| = R|org/jetbrains/kotlin/fir/plugin/producePositiveInt|() lval x: R|@Positive kotlin/Int| = R|org/jetbrains/kotlin/fir/plugin/producePositiveInt|()
R|/takePositive|(R|<local>/x|) R|/takePositive|(R|<local>/x|)
R|/takeNegative|(R|<local>/x|) R|/takeNegative|(R|<local>/x|)
R|/takeAny|(R|<local>/x|) R|/takeAny|(R|<local>/x|)
} }
public final fun test_3(): R|kotlin/Unit| { public final fun test_3(): R|kotlin/Unit| {
lval x: R|@R|org/jetbrains/kotlin/fir/plugin/Positive|() kotlin/Int| = R|org/jetbrains/kotlin/fir/plugin/produceBoxedPositiveInt|().R|SubstitutionOverride<org/jetbrains/kotlin/fir/plugin/Box.value: R|@R|org/jetbrains/kotlin/fir/plugin/Positive|() kotlin/Int|>| lval x: R|@Positive kotlin/Int| = R|org/jetbrains/kotlin/fir/plugin/produceBoxedPositiveInt|().R|SubstitutionOverride<org/jetbrains/kotlin/fir/plugin/Box.value: R|@Positive kotlin/Int|>|
R|/takePositive|(R|<local>/x|) R|/takePositive|(R|<local>/x|)
R|/takeNegative|(R|<local>/x|) R|/takeNegative|(R|<local>/x|)
R|/takeAny|(R|<local>/x|) R|/takeAny|(R|<local>/x|)
@@ -10,20 +10,20 @@ fun test_1(
someInt: Int someInt: Int
) { ) {
consumePositiveInt(positiveInt) consumePositiveInt(positiveInt)
consumePositiveInt(negativeInt) // should be error consumePositiveInt(<!ILLEGAL_NUMBER_SIGN!>negativeInt<!>) // should be error
consumePositiveInt(someInt) // should be error consumePositiveInt(<!ILLEGAL_NUMBER_SIGN!>someInt<!>) // should be error
} }
fun test_2() { fun test_2() {
val x = producePositiveInt() val x = producePositiveInt()
takePositive(<!ILLEGAL_NUMBER_SIGN!>x<!>) takePositive(x)
takeNegative(<!ILLEGAL_NUMBER_SIGN!>x<!>) // should be error takeNegative(<!ILLEGAL_NUMBER_SIGN!>x<!>) // should be error
takeAny(x) takeAny(x)
} }
fun test_3() { fun test_3() {
val x = produceBoxedPositiveInt().value val x = produceBoxedPositiveInt().value
takePositive(<!ILLEGAL_NUMBER_SIGN!>x<!>) takePositive(x)
takeNegative(<!ILLEGAL_NUMBER_SIGN!>x<!>) // should be error takeNegative(<!ILLEGAL_NUMBER_SIGN!>x<!>) // should be error
takeAny(x) takeAny(x)
} }