[FIR] Extract extension functional type from an annotated functional type during deserialization

^KT-57140 Fixed
This commit is contained in:
Dmitriy Novozhilov
2023-03-22 11:27:31 +02:00
committed by Space Team
parent 8ca7b32577
commit 336b6ba9f0
6 changed files with 67 additions and 23 deletions
@@ -47,7 +47,6 @@ abstract class FirAbstractSessionFactory {
registerCliCompilerOnlyComponents() registerCliCompilerOnlyComponents()
registerCommonComponents(languageVersionSettings) registerCommonComponents(languageVersionSettings)
registerCommonComponentsAfterExtensionsAreConfigured()
registerExtraComponents(this) registerExtraComponents(this)
val kotlinScopeProvider = createKotlinScopeProvider.invoke() val kotlinScopeProvider = createKotlinScopeProvider.invoke()
@@ -65,6 +64,7 @@ abstract class FirAbstractSessionFactory {
registerExtensions(extensionRegistrar.configure()) registerExtensions(extensionRegistrar.configure())
} }
}.configure() }.configure()
registerCommonComponentsAfterExtensionsAreConfigured()
val providers = createProviders(this, builtinsModuleData, kotlinScopeProvider) val providers = createProviders(this, builtinsModuleData, kotlinScopeProvider)
@@ -20,6 +20,7 @@ import org.jetbrains.kotlin.fir.symbols.ConeClassLikeLookupTag
import org.jetbrains.kotlin.fir.symbols.ConeClassifierLookupTag import org.jetbrains.kotlin.fir.symbols.ConeClassifierLookupTag
import org.jetbrains.kotlin.fir.symbols.ConeTypeParameterLookupTag import org.jetbrains.kotlin.fir.symbols.ConeTypeParameterLookupTag
import org.jetbrains.kotlin.fir.symbols.FirBasedSymbol import org.jetbrains.kotlin.fir.symbols.FirBasedSymbol
import org.jetbrains.kotlin.fir.symbols.impl.ConeClassLikeLookupTagImpl
import org.jetbrains.kotlin.fir.symbols.impl.FirClassLikeSymbol import org.jetbrains.kotlin.fir.symbols.impl.FirClassLikeSymbol
import org.jetbrains.kotlin.fir.symbols.impl.FirTypeParameterSymbol import org.jetbrains.kotlin.fir.symbols.impl.FirTypeParameterSymbol
import org.jetbrains.kotlin.fir.types.* import org.jetbrains.kotlin.fir.types.*
@@ -176,10 +177,28 @@ class FirTypeDeserializer(
argumentList + outerType(typeTable)?.collectAllArguments().orEmpty() argumentList + outerType(typeTable)?.collectAllArguments().orEmpty()
val arguments = proto.collectAllArguments().map(this::typeArgument).toTypedArray() val arguments = proto.collectAllArguments().map(this::typeArgument).toTypedArray()
val simpleType = if (Flags.SUSPEND_TYPE.get(proto.flags)) {
createSuspendFunctionType(constructor, arguments, isNullable = proto.nullable, attributes) val extensionFunctionalKind = moduleData.session.functionTypeService.extractSingleExtensionKindForDeserializedConeType(
} else { constructor.classId, attributes.customAnnotations
ConeClassLikeTypeImpl(constructor, arguments, isNullable = proto.nullable, attributes) )
val simpleType = when {
extensionFunctionalKind != null -> {
val newConstructor = if (arguments.isNotEmpty()) {
ConeClassLikeLookupTagImpl(extensionFunctionalKind.numberedClassId(arguments.size - 1))
} else {
return ConeErrorType(
ConeSimpleDiagnostic("Illegal number of arguments for extension functional type $extensionFunctionalKind"),
typeArguments = arguments,
attributes = attributes
)
}
ConeClassLikeTypeImpl(newConstructor, arguments, isNullable = proto.nullable, attributes)
}
Flags.SUSPEND_TYPE.get(proto.flags) -> {
createSuspendFunctionType(constructor, arguments, isNullable = proto.nullable, attributes)
}
else -> ConeClassLikeTypeImpl(constructor, arguments, isNullable = proto.nullable, attributes)
} }
val abbreviatedTypeProto = proto.abbreviatedType(typeTable) ?: return simpleType val abbreviatedTypeProto = proto.abbreviatedType(typeTable) ?: return simpleType
return simpleType(abbreviatedTypeProto, attributes) return simpleType(abbreviatedTypeProto, attributes)
@@ -10,6 +10,7 @@ import org.jetbrains.kotlin.builtins.functions.FunctionTypeKindExtractor
import org.jetbrains.kotlin.fir.FirSession import org.jetbrains.kotlin.fir.FirSession
import org.jetbrains.kotlin.fir.declarations.toAnnotationClassId import org.jetbrains.kotlin.fir.declarations.toAnnotationClassId
import org.jetbrains.kotlin.fir.declarations.utils.isSuspend import org.jetbrains.kotlin.fir.declarations.utils.isSuspend
import org.jetbrains.kotlin.fir.expressions.FirAnnotation
import org.jetbrains.kotlin.fir.extensions.FirFunctionTypeKindExtension import org.jetbrains.kotlin.fir.extensions.FirFunctionTypeKindExtension
import org.jetbrains.kotlin.fir.extensions.extensionService import org.jetbrains.kotlin.fir.extensions.extensionService
import org.jetbrains.kotlin.fir.extensions.functionTypeKindExtensions import org.jetbrains.kotlin.fir.extensions.functionTypeKindExtensions
@@ -75,6 +76,23 @@ class FirFunctionTypeKindServiceImpl(private val session: FirSession) : FirFunct
return extractSpecialKindsImpl(typeRef, { isSuspend }, { annotations.mapNotNull { it.toAnnotationClassId(session) } }) return extractSpecialKindsImpl(typeRef, { isSuspend }, { annotations.mapNotNull { it.toAnnotationClassId(session) } })
} }
override fun extractSingleExtensionKindForDeserializedConeType(
classId: ClassId,
annotations: List<FirAnnotation>
): FunctionTypeKind? {
if (nonReflectKindsFromExtensions.isEmpty() || annotations.isEmpty()) return null
val baseKind = extractor.getFunctionalClassKind(classId.packageFqName, classId.shortClassName.asString()) ?: return null
if (baseKind.nonReflectKind() != FunctionTypeKind.Function) return null
val matchingExtensionKinds = buildList {
extractKindsFromAnnotations(annotations.mapNotNull { it.toAnnotationClassId(session) })
}
val matchingKind = matchingExtensionKinds.singleOrNull() ?: return null
return when (baseKind.isReflectType) {
false -> matchingKind
true -> matchingKind.reflectKind()
}
}
private inline fun <T> extractSpecialKindsImpl( private inline fun <T> extractSpecialKindsImpl(
source: T, source: T,
isSuspend: T.() -> Boolean, isSuspend: T.() -> Boolean,
@@ -85,12 +103,16 @@ class FirFunctionTypeKindServiceImpl(private val session: FirSession) : FirFunct
add(FunctionTypeKind.SuspendFunction) add(FunctionTypeKind.SuspendFunction)
} }
if (nonReflectKindsFromExtensions.isNotEmpty()) { if (nonReflectKindsFromExtensions.isNotEmpty()) {
for (annotationClassId in source.annotations()) { extractKindsFromAnnotations(source.annotations())
for (kind in nonReflectKindsFromExtensions) { }
if (kind.annotationOnInvokeClassId == annotationClassId) { }
add(kind) }
}
} private fun MutableList<FunctionTypeKind>.extractKindsFromAnnotations(annotations: List<ClassId>) {
for (annotationClassId in annotations) {
for (kind in nonReflectKindsFromExtensions) {
if (kind.annotationOnInvokeClassId == annotationClassId) {
add(kind)
} }
} }
} }
@@ -9,7 +9,9 @@ import org.jetbrains.kotlin.builtins.functions.FunctionTypeKind
import org.jetbrains.kotlin.builtins.functions.FunctionTypeKindExtractor import org.jetbrains.kotlin.builtins.functions.FunctionTypeKindExtractor
import org.jetbrains.kotlin.fir.FirSession import org.jetbrains.kotlin.fir.FirSession
import org.jetbrains.kotlin.fir.FirSessionComponent import org.jetbrains.kotlin.fir.FirSessionComponent
import org.jetbrains.kotlin.fir.expressions.FirAnnotation
import org.jetbrains.kotlin.fir.symbols.impl.FirFunctionSymbol import org.jetbrains.kotlin.fir.symbols.impl.FirFunctionSymbol
import org.jetbrains.kotlin.name.ClassId
import org.jetbrains.kotlin.name.FqName import org.jetbrains.kotlin.name.FqName
abstract class FirFunctionTypeKindService : FirSessionComponent { abstract class FirFunctionTypeKindService : FirSessionComponent {
@@ -36,6 +38,7 @@ abstract class FirFunctionTypeKindService : FirSessionComponent {
abstract fun extractSingleSpecialKindForFunction(functionSymbol: FirFunctionSymbol<*>): FunctionTypeKind? abstract fun extractSingleSpecialKindForFunction(functionSymbol: FirFunctionSymbol<*>): FunctionTypeKind?
abstract fun extractAllSpecialKindsForFunction(functionSymbol: FirFunctionSymbol<*>): List<FunctionTypeKind> abstract fun extractAllSpecialKindsForFunction(functionSymbol: FirFunctionSymbol<*>): List<FunctionTypeKind>
abstract fun extractAllSpecialKindsForFunctionTypeRef(typeRef: FirFunctionTypeRef): List<FunctionTypeKind> abstract fun extractAllSpecialKindsForFunctionTypeRef(typeRef: FirFunctionTypeRef): List<FunctionTypeKind>
abstract fun extractSingleExtensionKindForDeserializedConeType(classId: ClassId, annotations: List<FirAnnotation>): FunctionTypeKind?
} }
val FirSession.functionTypeService: FirFunctionTypeKindService by FirSession.sessionComponentAccessor() val FirSession.functionTypeService: FirFunctionTypeKindService by FirSession.sessionComponentAccessor()
@@ -7,20 +7,20 @@ FILE: dependencyWithoutFunctionalKindPlugin.kt
} }
public final fun test_1(block: R|() -> kotlin/Unit|, composableBlock: R|@R|org/jetbrains/kotlin/fir/plugin/MyComposable|() some/MyComposableFunction0<kotlin/Unit>|, suspendBlock: R|suspend () -> kotlin/Unit|): R|kotlin/Unit| { public final fun test_1(block: R|() -> kotlin/Unit|, composableBlock: R|@R|org/jetbrains/kotlin/fir/plugin/MyComposable|() some/MyComposableFunction0<kotlin/Unit>|, suspendBlock: R|suspend () -> kotlin/Unit|): R|kotlin/Unit| {
R|org/jetbrains/kotlin/fir/plugin/consumeComposableFunction|(R|<local>/block|) R|org/jetbrains/kotlin/fir/plugin/consumeComposableFunction|(R|<local>/block|)
R|org/jetbrains/kotlin/fir/plugin/consumeComposableFunction<Inapplicable(INAPPLICABLE): org/jetbrains/kotlin/fir/plugin/consumeComposableFunction>#|(R|<local>/composableBlock|) R|org/jetbrains/kotlin/fir/plugin/consumeComposableFunction|(R|<local>/composableBlock|)
R|org/jetbrains/kotlin/fir/plugin/consumeComposableFunction<Inapplicable(INAPPLICABLE): org/jetbrains/kotlin/fir/plugin/consumeComposableFunction>#|(R|<local>/suspendBlock|) R|org/jetbrains/kotlin/fir/plugin/consumeComposableFunction<Inapplicable(INAPPLICABLE): org/jetbrains/kotlin/fir/plugin/consumeComposableFunction>#|(R|<local>/suspendBlock|)
} }
public final fun test_2(): R|kotlin/Unit| { public final fun test_2(): R|kotlin/Unit| {
lval block: R|() -> kotlin/Unit| = R|org/jetbrains/kotlin/fir/plugin/produceComposableFunction|() lval block: R|@R|org/jetbrains/kotlin/fir/plugin/MyComposable|() some/MyComposableFunction0<kotlin/Unit>| = R|org/jetbrains/kotlin/fir/plugin/produceComposableFunction|()
R|/consumeRegularFunction|(R|<local>/block|) R|/consumeRegularFunction<Inapplicable(INAPPLICABLE): /consumeRegularFunction>#|(R|<local>/block|)
R|/consumeSuspendFunction|(R|<local>/block|) R|/consumeSuspendFunction<Inapplicable(INAPPLICABLE): /consumeSuspendFunction>#|(R|<local>/block|)
R|/consumeOurComposableFunction|(R|<local>/block|) R|/consumeOurComposableFunction|(R|<local>/block|)
R|org/jetbrains/kotlin/fir/plugin/consumeComposableFunction|(R|<local>/block|) R|org/jetbrains/kotlin/fir/plugin/consumeComposableFunction|(R|<local>/block|)
} }
public final fun test_3(): R|kotlin/Unit| { public final fun test_3(): R|kotlin/Unit| {
lval block: R|() -> kotlin/Unit| = R|org/jetbrains/kotlin/fir/plugin/produceBoxedComposableFunction|().R|SubstitutionOverride<org/jetbrains/kotlin/fir/plugin/Box.value: R|() -> kotlin/Unit|>| lval block: R|@R|org/jetbrains/kotlin/fir/plugin/MyComposable|() some/MyComposableFunction0<kotlin/Unit>| = R|org/jetbrains/kotlin/fir/plugin/produceBoxedComposableFunction|().R|SubstitutionOverride<org/jetbrains/kotlin/fir/plugin/Box.value: R|@R|org/jetbrains/kotlin/fir/plugin/MyComposable|() some/MyComposableFunction0<kotlin/Unit>|>|
R|/consumeRegularFunction|(R|<local>/block|) R|/consumeRegularFunction<Inapplicable(INAPPLICABLE): /consumeRegularFunction>#|(R|<local>/block|)
R|/consumeSuspendFunction|(R|<local>/block|) R|/consumeSuspendFunction<Inapplicable(INAPPLICABLE): /consumeSuspendFunction>#|(R|<local>/block|)
R|/consumeOurComposableFunction|(R|<local>/block|) R|/consumeOurComposableFunction|(R|<local>/block|)
R|org/jetbrains/kotlin/fir/plugin/consumeComposableFunction|(R|<local>/block|) R|org/jetbrains/kotlin/fir/plugin/consumeComposableFunction|(R|<local>/block|)
} }
@@ -10,22 +10,22 @@ fun test_1(
suspendBlock: suspend () -> Unit, suspendBlock: suspend () -> Unit,
) { ) {
consumeComposableFunction(block) consumeComposableFunction(block)
consumeComposableFunction(<!ARGUMENT_TYPE_MISMATCH!>composableBlock<!>) consumeComposableFunction(composableBlock)
consumeComposableFunction(<!ARGUMENT_TYPE_MISMATCH!>suspendBlock<!>) // should be error consumeComposableFunction(<!ARGUMENT_TYPE_MISMATCH!>suspendBlock<!>) // should be error
} }
fun test_2() { fun test_2() {
val block = produceComposableFunction() val block = produceComposableFunction()
consumeRegularFunction(block) // should be error consumeRegularFunction(<!ARGUMENT_TYPE_MISMATCH!>block<!>) // should be error
consumeSuspendFunction(block) // should be error consumeSuspendFunction(<!ARGUMENT_TYPE_MISMATCH!>block<!>) // should be error
consumeOurComposableFunction(block) consumeOurComposableFunction(block)
consumeComposableFunction(block) consumeComposableFunction(block)
} }
fun test_3() { fun test_3() {
val block = produceBoxedComposableFunction().value val block = produceBoxedComposableFunction().value
consumeRegularFunction(block) // should be error consumeRegularFunction(<!ARGUMENT_TYPE_MISMATCH!>block<!>) // should be error
consumeSuspendFunction(block) // should be error consumeSuspendFunction(<!ARGUMENT_TYPE_MISMATCH!>block<!>) // should be error
consumeOurComposableFunction(block) consumeOurComposableFunction(block)
consumeComposableFunction(block) consumeComposableFunction(block)
} }