[K2] Pass FirFile to metadata serialization

This is needed to extract correct const value from `ConstValueProvider`.

#KT-57812
This commit is contained in:
Ivan Kylchik
2023-04-21 12:47:42 +02:00
committed by Space Team
parent d26e3871ba
commit 951e30b683
8 changed files with 53 additions and 39 deletions
@@ -78,7 +78,7 @@ class FirElementSerializer private constructor(
fun packagePartProto( fun packagePartProto(
packageFqName: FqName, packageFqName: FqName,
files: List<FirFile>, file: FirFile,
actualizedExpectDeclarations: Set<FirDeclaration>? actualizedExpectDeclarations: Set<FirDeclaration>?
): ProtoBuf.Package.Builder { ): ProtoBuf.Package.Builder {
val builder = ProtoBuf.Package.newBuilder() val builder = ProtoBuf.Package.newBuilder()
@@ -88,7 +88,7 @@ class FirElementSerializer private constructor(
if (!declaration.shouldBeSerialized(actualizedExpectDeclarations)) return if (!declaration.shouldBeSerialized(actualizedExpectDeclarations)) return
when (declaration) { when (declaration) {
is FirProperty -> propertyProto(declaration)?.let { builder.addProperty(it) } is FirProperty -> propertyProto(declaration)?.let { builder.addProperty(it) }
is FirSimpleFunction -> functionProto(declaration)?.let { builder.addFunction(it) } is FirSimpleFunction -> privateFunctionProto(declaration)?.let { builder.addFunction(it) }
is FirTypeAlias -> typeAliasProto(declaration)?.let { builder.addTypeAlias(it) } is FirTypeAlias -> typeAliasProto(declaration)?.let { builder.addTypeAlias(it) }
else -> onUnsupportedDeclaration(declaration) else -> onUnsupportedDeclaration(declaration)
} }
@@ -97,18 +97,16 @@ class FirElementSerializer private constructor(
} }
} }
for (file in files) { processFile(file) {
extension.processFile(file) { for (declaration in file.declarations) {
for (declaration in file.declarations) { addDeclaration(declaration) {}
addDeclaration(declaration) {}
}
} }
} extension.serializePackage(packageFqName, builder)
extension.serializePackage(packageFqName, builder) for (declaration in providedDeclarationsService.getProvidedTopLevelDeclarations(packageFqName, scopeSession)) {
for (declaration in providedDeclarationsService.getProvidedTopLevelDeclarations(packageFqName, scopeSession)) { addDeclaration(declaration) {
addDeclaration(declaration) { error("Unsupported top-level declaration type: ${it.render()}")
error("Unsupported top-level declaration type: ${it.render()}") }
} }
} }
@@ -118,7 +116,20 @@ class FirElementSerializer private constructor(
return builder return builder
} }
fun classProto(klass: FirClass): ProtoBuf.Class.Builder = whileAnalysing(session, klass) { private inline fun <T> processFile(firFile: FirFile, crossinline action: () -> T): T {
return extension.processFile(firFile) {
action()
}
}
// Note: we could try to extract FirFile from `session.firProvider.getFirClassifierContainerFile` but it doesn't work for anonymous objects
fun classProto(klass: FirClass, firFile: FirFile): ProtoBuf.Class.Builder {
return processFile(firFile) {
privateClassProto(klass)
}
}
private fun privateClassProto(klass: FirClass): ProtoBuf.Class.Builder = whileAnalysing(session, klass) {
val builder = ProtoBuf.Class.newBuilder() val builder = ProtoBuf.Class.newBuilder()
val regularClass = klass as? FirRegularClass val regularClass = klass as? FirRegularClass
@@ -198,7 +209,7 @@ class FirElementSerializer private constructor(
if (declaration !is FirEnumEntry && declaration.isStatic) continue // ??? Miss values() & valueOf() if (declaration !is FirEnumEntry && declaration.isStatic) continue // ??? Miss values() & valueOf()
when (declaration) { when (declaration) {
is FirProperty -> propertyProto(declaration)?.let { builder.addProperty(it) } is FirProperty -> propertyProto(declaration)?.let { builder.addProperty(it) }
is FirSimpleFunction -> functionProto(declaration)?.let { builder.addFunction(it) } is FirSimpleFunction -> privateFunctionProto(declaration)?.let { builder.addFunction(it) }
is FirEnumEntry -> enumEntryProto(declaration).let { builder.addEnumEntry(it) } is FirEnumEntry -> enumEntryProto(declaration).let { builder.addEnumEntry(it) }
else -> {} else -> {}
} }
@@ -489,7 +500,13 @@ class FirElementSerializer private constructor(
return builder return builder
} }
fun functionProto(function: FirFunction): ProtoBuf.Function.Builder? = whileAnalysing(session, function) { fun functionProto(function: FirFunction, firFile: FirFile): ProtoBuf.Function.Builder? {
return processFile(firFile) {
privateFunctionProto(function)
}
}
fun privateFunctionProto(function: FirFunction): ProtoBuf.Function.Builder? = whileAnalysing(session, function) {
val builder = ProtoBuf.Function.newBuilder() val builder = ProtoBuf.Function.newBuilder()
val simpleFunction = function as? FirSimpleFunction val simpleFunction = function as? FirSimpleFunction
@@ -8,8 +8,8 @@ package org.jetbrains.kotlin.fir.serialization
import org.jetbrains.kotlin.fir.FirSession import org.jetbrains.kotlin.fir.FirSession
import org.jetbrains.kotlin.fir.declarations.* import org.jetbrains.kotlin.fir.declarations.*
import org.jetbrains.kotlin.fir.expressions.FirAnnotation import org.jetbrains.kotlin.fir.expressions.FirAnnotation
import org.jetbrains.kotlin.fir.serialization.constant.ConstValueProviderInternals
import org.jetbrains.kotlin.fir.serialization.constant.ConstValueProvider import org.jetbrains.kotlin.fir.serialization.constant.ConstValueProvider
import org.jetbrains.kotlin.fir.serialization.constant.ConstValueProviderInternals
import org.jetbrains.kotlin.fir.types.ConeErrorType import org.jetbrains.kotlin.fir.types.ConeErrorType
import org.jetbrains.kotlin.fir.types.ConeFlexibleType import org.jetbrains.kotlin.fir.types.ConeFlexibleType
import org.jetbrains.kotlin.metadata.ProtoBuf import org.jetbrains.kotlin.metadata.ProtoBuf
@@ -29,7 +29,7 @@ abstract class FirSerializerExtension {
protected abstract val constValueProvider: ConstValueProvider? protected abstract val constValueProvider: ConstValueProvider?
@OptIn(ConstValueProviderInternals::class) @OptIn(ConstValueProviderInternals::class)
internal inline fun <T> processFile(firFile: FirFile, action: () -> T): T { internal inline fun <T> processFile(firFile: FirFile, crossinline action: () -> T): T {
val previousFile = constValueProvider?.processingFirFile val previousFile = constValueProvider?.processingFirFile
constValueProvider?.processingFirFile = firFile constValueProvider?.processingFirFile = firFile
return try { return try {
@@ -35,7 +35,7 @@ fun serializeSingleFirFile(
// TODO: split package fragment (see klib serializer) // TODO: split package fragment (see klib serializer)
// TODO: handle incremental/monolothic (see klib serializer) - maybe externally // TODO: handle incremental/monolothic (see klib serializer) - maybe externally
val packageProto = packageSerializer.packagePartProto(file.packageFqName, listOf(file), actualizedExpectDeclarations).build() val packageProto = packageSerializer.packagePartProto(file.packageFqName, file, actualizedExpectDeclarations).build()
val classesProto = mutableListOf<Pair<ProtoBuf.Class, Int>>() val classesProto = mutableListOf<Pair<ProtoBuf.Class, Int>>()
@@ -51,14 +51,12 @@ fun serializeSingleFirFile(
) )
val index = classSerializer.stringTable.getFqNameIndex(klass) val index = classSerializer.stringTable.getFqNameIndex(klass)
classesProto += classSerializer.classProto(klass).build() to index classesProto += classSerializer.classProto(klass, file).build() to index
classSerializer.computeNestedClassifiersForClass(symbol).filterIsInstance<FirClassSymbol<*>>().makeClassesProtoWithNested() classSerializer.computeNestedClassifiersForClass(symbol).filterIsInstance<FirClassSymbol<*>>().makeClassesProtoWithNested()
} }
} }
serializerExtension.processFile(file) { file.declarations.mapNotNull { it.symbol as? FirClassSymbol<*> }.makeClassesProtoWithNested()
file.declarations.mapNotNull { it.symbol as? FirClassSymbol<*> }.makeClassesProtoWithNested()
}
val hasTopLevelDeclarations = file.declarations.any { val hasTopLevelDeclarations = file.declarations.any {
it is FirMemberDeclaration && it.shouldBeSerialized(actualizedExpectDeclarations) && it is FirMemberDeclaration && it.shouldBeSerialized(actualizedExpectDeclarations) &&
@@ -134,9 +134,6 @@ class FirJvmSerializerExtension(
if (moduleName != JvmProtoBufUtil.DEFAULT_MODULE_NAME) { if (moduleName != JvmProtoBufUtil.DEFAULT_MODULE_NAME) {
proto.setExtension(JvmProtoBuf.packageModuleName, stringTable.getStringIndex(moduleName)) proto.setExtension(JvmProtoBuf.packageModuleName, stringTable.getStringIndex(moduleName))
} }
}
fun serializeJvmPackage(proto: ProtoBuf.Package.Builder) {
writeLocalProperties(proto, JvmProtoBuf.packageLocalVariable) writeLocalProperties(proto, JvmProtoBuf.packageLocalVariable)
} }
@@ -26,13 +26,15 @@ import org.jetbrains.kotlin.fir.resolve.toFirRegularClass
import org.jetbrains.kotlin.fir.serialization.FirElementAwareStringTable import org.jetbrains.kotlin.fir.serialization.FirElementAwareStringTable
import org.jetbrains.kotlin.fir.serialization.FirElementSerializer import org.jetbrains.kotlin.fir.serialization.FirElementSerializer
import org.jetbrains.kotlin.fir.serialization.TypeApproximatorForMetadataSerializer import org.jetbrains.kotlin.fir.serialization.TypeApproximatorForMetadataSerializer
import org.jetbrains.kotlin.fir.serialization.constant.* import org.jetbrains.kotlin.fir.symbols.impl.FirAnonymousFunctionSymbol
import org.jetbrains.kotlin.fir.symbols.impl.* import org.jetbrains.kotlin.fir.symbols.impl.FirDelegateFieldSymbol
import org.jetbrains.kotlin.fir.symbols.impl.FirPropertyAccessorSymbol
import org.jetbrains.kotlin.fir.symbols.impl.FirPropertySymbol
import org.jetbrains.kotlin.fir.types.* import org.jetbrains.kotlin.fir.types.*
import org.jetbrains.kotlin.ir.declarations.IrClass import org.jetbrains.kotlin.ir.declarations.IrClass
import org.jetbrains.kotlin.ir.declarations.IrDeclarationOrigin import org.jetbrains.kotlin.ir.declarations.IrDeclarationOrigin
import org.jetbrains.kotlin.ir.declarations.MetadataSource import org.jetbrains.kotlin.ir.declarations.MetadataSource
import org.jetbrains.kotlin.ir.types.* import org.jetbrains.kotlin.ir.util.file
import org.jetbrains.kotlin.metadata.jvm.serialization.JvmStringTable import org.jetbrains.kotlin.metadata.jvm.serialization.JvmStringTable
import org.jetbrains.kotlin.modules.TargetId import org.jetbrains.kotlin.modules.TargetId
import org.jetbrains.kotlin.name.ClassId import org.jetbrains.kotlin.name.ClassId
@@ -61,9 +63,9 @@ fun makeFirMetadataSerializerForIrClass(
approximator, context.defaultTypeMapper, components approximator, context.defaultTypeMapper, components
) )
return FirMetadataSerializer( return FirMetadataSerializer(
(irClass.file.metadata as? FirMetadataSource.File)?.file,
context.state.globalSerializationBindings, context.state.globalSerializationBindings,
serializationBindings, serializationBindings,
firSerializerExtension,
approximator, approximator,
makeElementSerializer( makeElementSerializer(
irClass.metadata, components.session, components.scopeSession, firSerializerExtension, approximator, parent, irClass.metadata, components.session, components.scopeSession, firSerializerExtension, approximator, parent,
@@ -73,8 +75,8 @@ fun makeFirMetadataSerializerForIrClass(
) )
} }
@OptIn(LookupTagInternals::class)
fun makeLocalFirMetadataSerializerForMetadataSource( fun makeLocalFirMetadataSerializerForMetadataSource(
firFile: FirFile,
metadata: MetadataSource?, metadata: MetadataSource?,
session: FirSession, session: FirSession,
scopeSession: ScopeSession, scopeSession: ScopeSession,
@@ -108,9 +110,9 @@ fun makeLocalFirMetadataSerializerForMetadataSource(
constValueProvider = null constValueProvider = null
) )
return FirMetadataSerializer( return FirMetadataSerializer(
firFile,
globalSerializationBindings, globalSerializationBindings,
serializationBindings, serializationBindings,
firSerializerExtension,
approximator, approximator,
makeElementSerializer( makeElementSerializer(
metadata, session, scopeSession, firSerializerExtension, approximator, parent, metadata, session, scopeSession, firSerializerExtension, approximator, parent,
@@ -121,9 +123,9 @@ fun makeLocalFirMetadataSerializerForMetadataSource(
} }
class FirMetadataSerializer( class FirMetadataSerializer(
private val firFile: FirFile?,
private val globalSerializationBindings: JvmSerializationBindings, private val globalSerializationBindings: JvmSerializationBindings,
private val serializationBindings: JvmSerializationBindings, private val serializationBindings: JvmSerializationBindings,
private val serializerExtension: FirJvmSerializerExtension,
private val approximator: AbstractTypeApproximator, private val approximator: AbstractTypeApproximator,
internal val serializer: FirElementSerializer?, internal val serializer: FirElementSerializer?,
irActualizedResult: IrActualizedResult? irActualizedResult: IrActualizedResult?
@@ -132,16 +134,15 @@ class FirMetadataSerializer(
override fun serialize(metadata: MetadataSource): Pair<MessageLite, JvmStringTable>? { override fun serialize(metadata: MetadataSource): Pair<MessageLite, JvmStringTable>? {
val message = when (metadata) { val message = when (metadata) {
is FirMetadataSource.Class -> serializer!!.classProto(metadata.fir).build() is FirMetadataSource.Class -> serializer!!.classProto(metadata.fir, firFile!!).build()
is FirMetadataSource.File -> is FirMetadataSource.File ->
serializer!!.packagePartProto(metadata.files.first().packageFqName, metadata.files, actualizedExpectDeclarations) serializer!!.packagePartProto(metadata.file.packageFqName, metadata.file, actualizedExpectDeclarations).build()
.apply { serializerExtension.serializeJvmPackage(this) }.build()
is FirMetadataSource.Function -> { is FirMetadataSource.Function -> {
val withTypeParameters = metadata.fir.copyToFreeAnonymousFunction(approximator) val withTypeParameters = metadata.fir.copyToFreeAnonymousFunction(approximator)
serializationBindings.get(FirJvmSerializerExtension.METHOD_FOR_FIR_FUNCTION, metadata.fir)?.let { serializationBindings.get(FirJvmSerializerExtension.METHOD_FOR_FIR_FUNCTION, metadata.fir)?.let {
serializationBindings.put(FirJvmSerializerExtension.METHOD_FOR_FIR_FUNCTION, withTypeParameters, it) serializationBindings.put(FirJvmSerializerExtension.METHOD_FOR_FIR_FUNCTION, withTypeParameters, it)
} }
serializer!!.functionProto(withTypeParameters)?.build() serializer!!.functionProto(withTypeParameters, firFile!!)?.build()
} }
else -> null else -> null
} ?: return null } ?: return null
@@ -111,7 +111,7 @@ class Fir2IrVisitor(
it.toIrDeclaration() it.toIrDeclaration()
} }
annotationGenerator.generate(this, file) annotationGenerator.generate(this, file)
metadata = FirMetadataSource.File(listOf(file)) metadata = FirMetadataSource.File(file)
} }
} }
@@ -23,7 +23,7 @@ sealed class FirMetadataSource : MetadataSource {
else -> null else -> null
} }
class File(val files: List<FirFile>) : FirMetadataSource(), MetadataSource.File { class File(val file: FirFile) : FirMetadataSource(), MetadataSource.File {
override var serializedIr: ByteArray? = null override var serializedIr: ByteArray? = null
override val fir: FirDeclaration? override val fir: FirDeclaration?
@@ -46,6 +46,7 @@ internal fun collectNewDirtySources(
body: (MetadataSerializer) -> Unit body: (MetadataSerializer) -> Unit
) { ) {
val serializer = makeLocalFirMetadataSerializerForMetadataSource( val serializer = makeLocalFirMetadataSerializerForMetadataSource(
it,
metadata, metadata,
analyzedOutput.session, analyzedOutput.session,
analyzedOutput.scopeSession, analyzedOutput.scopeSession,
@@ -73,7 +74,7 @@ internal fun collectNewDirtySources(
} }
override fun visitFile(file: FirFile, data: MutableList<MetadataSerializer>) { override fun visitFile(file: FirFile, data: MutableList<MetadataSerializer>) {
val metadata = FirMetadataSource.File(listOf(file)) val metadata = FirMetadataSource.File(file)
withMetadataSerializer(metadata, data) { withMetadataSerializer(metadata, data) {
file.acceptChildren(this, data) file.acceptChildren(this, data)
// TODO: compare package fragments? // TODO: compare package fragments?