FIR: use serialization extension protocol correctly

otherwise the deserialization breaks on KLibs
This commit is contained in:
Ilya Chernikov
2022-07-19 12:13:20 +02:00
parent 112f91ba3b
commit 8f18ab19f7
12 changed files with 43 additions and 28 deletions
@@ -36,18 +36,15 @@ import org.jetbrains.kotlin.name.Name
import org.jetbrains.kotlin.name.StandardClassIds
import org.jetbrains.kotlin.protobuf.MessageLite
import org.jetbrains.kotlin.serialization.SerializerExtensionProtocol
import org.jetbrains.kotlin.serialization.deserialization.builtins.BuiltInSerializerProtocol
import org.jetbrains.kotlin.serialization.deserialization.descriptors.DeserializedContainerSource
import org.jetbrains.kotlin.serialization.deserialization.getClassId
import org.jetbrains.kotlin.serialization.deserialization.getName
import org.jetbrains.kotlin.types.ConstantValueKind
abstract class AbstractAnnotationDeserializer(
private val session: FirSession
private val session: FirSession,
protected val protocol: SerializerExtensionProtocol
) {
protected open val protocol: SerializerExtensionProtocol
get() = BuiltInSerializerProtocol
open fun inheritAnnotationInfo(parent: AbstractAnnotationDeserializer) {
}
@@ -20,6 +20,7 @@ import org.jetbrains.kotlin.fir.symbols.impl.*
import org.jetbrains.kotlin.metadata.ProtoBuf
import org.jetbrains.kotlin.metadata.deserialization.NameResolver
import org.jetbrains.kotlin.name.*
import org.jetbrains.kotlin.serialization.SerializerExtensionProtocol
import org.jetbrains.kotlin.serialization.deserialization.descriptors.DeserializedContainerSource
import org.jetbrains.kotlin.serialization.deserialization.getName
import java.nio.file.Path
@@ -68,6 +69,7 @@ abstract class AbstractFirDeserializedSymbolProvider(
val moduleDataProvider: ModuleDataProvider,
val kotlinScopeProvider: FirKotlinScopeProvider,
val defaultDeserializationOrigin: FirDeclarationOrigin,
private val serializerExtensionProtocol: SerializerExtensionProtocol
) : FirSymbolProvider(session) {
// ------------------------ Caches ------------------------
@@ -146,6 +148,7 @@ abstract class AbstractFirDeserializedSymbolProvider(
moduleData,
annotationDeserializer,
kotlinScopeProvider,
serializerExtensionProtocol,
parentContext,
sourceElement,
origin = defaultDeserializationOrigin,
@@ -30,7 +30,9 @@ import org.jetbrains.kotlin.metadata.ProtoBuf
import org.jetbrains.kotlin.metadata.deserialization.*
import org.jetbrains.kotlin.metadata.jvm.JvmProtoBuf
import org.jetbrains.kotlin.name.*
import org.jetbrains.kotlin.serialization.SerializerExtensionProtocol
import org.jetbrains.kotlin.serialization.deserialization.ProtoEnumFlags
import org.jetbrains.kotlin.serialization.deserialization.builtins.BuiltInSerializerProtocol
import org.jetbrains.kotlin.serialization.deserialization.descriptors.DeserializedContainerSource
import org.jetbrains.kotlin.serialization.deserialization.getName
import org.jetbrains.kotlin.serialization.deserialization.loadValueClassRepresentation
@@ -44,6 +46,7 @@ fun deserializeClassToSymbol(
moduleData: FirModuleData,
defaultAnnotationDeserializer: AbstractAnnotationDeserializer?,
scopeProvider: FirScopeProvider,
serializerExtensionProtocol: SerializerExtensionProtocol,
parentContext: FirDeserializationContext? = null,
containerSource: DeserializedContainerSource? = null,
origin: FirDeclarationOrigin = FirDeclarationOrigin.Library,
@@ -71,9 +74,9 @@ fun deserializeClassToSymbol(
val annotationDeserializer = defaultAnnotationDeserializer ?: FirBuiltinAnnotationDeserializer(session)
val jvmBinaryClass = (containerSource as? KotlinJvmBinarySourceElement)?.binaryClass
val constDeserializer = if (jvmBinaryClass != null) {
FirJvmConstDeserializer(session, jvmBinaryClass)
FirJvmConstDeserializer(session, jvmBinaryClass, serializerExtensionProtocol)
} else {
FirConstDeserializer(session)
FirConstDeserializer(session, serializerExtensionProtocol)
}
val context =
parentContext?.childContext(
@@ -87,8 +90,9 @@ fun deserializeClassToSymbol(
if (status.isCompanion) {
parentContext.constDeserializer
} else {
((containerSource as? KotlinJvmBinarySourceElement)?.binaryClass)?.let { FirJvmConstDeserializer(session, it) }
?: parentContext.constDeserializer
((containerSource as? KotlinJvmBinarySourceElement)?.binaryClass)?.let {
FirJvmConstDeserializer(session, it, serializerExtensionProtocol)
} ?: parentContext.constDeserializer
},
status.isInner
) ?: FirDeserializationContext.createForClass(
@@ -10,10 +10,11 @@ import org.jetbrains.kotlin.fir.expressions.FirAnnotation
import org.jetbrains.kotlin.metadata.ProtoBuf
import org.jetbrains.kotlin.metadata.deserialization.Flags
import org.jetbrains.kotlin.metadata.deserialization.NameResolver
import org.jetbrains.kotlin.serialization.deserialization.builtins.BuiltInSerializerProtocol
class FirBuiltinAnnotationDeserializer(
session: FirSession
) : AbstractAnnotationDeserializer(session) {
) : AbstractAnnotationDeserializer(session, BuiltInSerializerProtocol) {
override fun loadTypeAnnotations(typeProto: ProtoBuf.Type, nameResolver: NameResolver): List<FirAnnotation> {
if (!Flags.HAS_ANNOTATIONS.get(typeProto.flags)) return emptyList()
@@ -8,24 +8,25 @@ package org.jetbrains.kotlin.fir.deserialization
import org.jetbrains.kotlin.fir.FirSession
import org.jetbrains.kotlin.fir.expressions.FirExpression
import org.jetbrains.kotlin.fir.expressions.builder.buildConstExpression
import org.jetbrains.kotlin.name.CallableId
import org.jetbrains.kotlin.metadata.ProtoBuf
import org.jetbrains.kotlin.metadata.deserialization.Flags
import org.jetbrains.kotlin.metadata.deserialization.NameResolver
import org.jetbrains.kotlin.metadata.deserialization.getExtensionOrNull
import org.jetbrains.kotlin.name.CallableId
import org.jetbrains.kotlin.name.Name
import org.jetbrains.kotlin.serialization.deserialization.builtins.BuiltInSerializerProtocol
import org.jetbrains.kotlin.serialization.SerializerExtensionProtocol
import org.jetbrains.kotlin.types.ConstantValueKind
open class FirConstDeserializer(
val session: FirSession
val session: FirSession,
private val protocol: SerializerExtensionProtocol
) {
protected val constantCache = mutableMapOf<CallableId, FirExpression>()
open fun loadConstant(propertyProto: ProtoBuf.Property, callableId: CallableId, nameResolver: NameResolver): FirExpression? {
if (!Flags.HAS_CONSTANT.get(propertyProto.flags)) return null
constantCache[callableId]?.let { return it }
val value = propertyProto.getExtensionOrNull(BuiltInSerializerProtocol.compileTimeValue) ?: return null
val value = propertyProto.getExtensionOrNull(protocol.compileTimeValue) ?: return null
return buildFirConstant(value, null, value.type.name, nameResolver)?.also { constantCache[callableId] = it }
}
}
@@ -13,11 +13,13 @@ import org.jetbrains.kotlin.metadata.ProtoBuf
import org.jetbrains.kotlin.metadata.deserialization.Flags
import org.jetbrains.kotlin.metadata.deserialization.NameResolver
import org.jetbrains.kotlin.name.Name
import org.jetbrains.kotlin.serialization.SerializerExtensionProtocol
class FirJvmConstDeserializer(
session: FirSession,
private val binaryClass: KotlinJvmBinaryClass,
) : FirConstDeserializer(session) {
protocol: SerializerExtensionProtocol,
) : FirConstDeserializer(session, protocol) {
override fun loadConstant(propertyProto: ProtoBuf.Property, callableId: CallableId, nameResolver: NameResolver): FirExpression? {
if (!Flags.HAS_CONSTANT.get(propertyProto.flags)) return null
constantCache[callableId]?.let { return it }
@@ -118,7 +118,7 @@ open class FirBuiltinSymbolProvider(
FirDeserializationContext.createForPackage(
fqName, packageProto.`package`, nameResolver, moduleData,
FirBuiltinAnnotationDeserializer(moduleData.session),
FirConstDeserializer(moduleData.session),
FirConstDeserializer(moduleData.session, BuiltInSerializerProtocol),
containerSource = null
).memberDeserializer
}
@@ -131,7 +131,7 @@ open class FirBuiltinSymbolProvider(
deserializeClassToSymbol(
classId, classProto, symbol, nameResolver, moduleData.session, moduleData,
null, kotlinScopeProvider, parentContext,
null, kotlinScopeProvider, BuiltInSerializerProtocol, parentContext,
null,
origin = FirDeclarationOrigin.BuiltIns,
this::findAndDeserializeClass,