Support enums according to new design

This commit is contained in:
Leonid Startsev
2018-10-12 17:26:35 +03:00
parent f101e17dfa
commit dba6396e95
2 changed files with 45 additions and 17 deletions
@@ -107,7 +107,6 @@ interface IrBuilderExtension {
} }
fun IrBuilderWithScope.createArrayOfExpression( fun IrBuilderWithScope.createArrayOfExpression(
resultingType: IrType,
arrayElementType: IrType, arrayElementType: IrType,
arrayElements: List<IrExpression> arrayElements: List<IrExpression>
): IrExpression { ): IrExpression {
@@ -116,7 +115,7 @@ interface IrBuilderExtension {
val arg0 = IrVarargImpl(startOffset, endOffset, arrayType, arrayElementType, arrayElements) val arg0 = IrVarargImpl(startOffset, endOffset, arrayType, arrayElementType, arrayElements)
val typeArguments = listOf(arrayElementType) val typeArguments = listOf(arrayElementType)
return irCall(compilerContext.ir.symbols.arrayOf, resultingType, typeArguments = typeArguments).apply { return irCall(compilerContext.ir.symbols.arrayOf, arrayType, typeArguments = typeArguments).apply {
putValueArgument(0, arg0) putValueArgument(0, arg0)
} }
} }
@@ -381,6 +380,21 @@ interface IrBuilderExtension {
) )
} }
fun findEnumValuesMethod(enumClass: ClassDescriptor): IrFunction {
assert(enumClass.kind == ClassKind.ENUM_CLASS)
return compilerContext.externalSymbols.referenceClass(enumClass).owner.functions
.find { it.origin == IrDeclarationOrigin.ENUM_CLASS_SPECIAL_MEMBER && it.name == Name.identifier("values") }
?: throw AssertionError("Enum class does not have .values() function")
}
private fun getEnumMembersNames(enumClass: ClassDescriptor): Sequence<String> {
assert(enumClass.kind == ClassKind.ENUM_CLASS)
return enumClass.unsubstitutedMemberScope.getContributedDescriptors().asSequence()
.filterIsInstance<ClassDescriptor>()
.filter { it.kind == ClassKind.ENUM_ENTRY }
.map { it.name.toString() }
}
// Does not use sti and therefore does not perform encoder calls optimization // Does not use sti and therefore does not perform encoder calls optimization
fun IrBuilderWithScope.serializerTower(generator: SerializerIrGenerator, property: SerializableProperty): IrExpression? { fun IrBuilderWithScope.serializerTower(generator: SerializerIrGenerator, property: SerializableProperty): IrExpression? {
val nullableSerClass = val nullableSerClass =
@@ -400,33 +414,48 @@ interface IrBuilderExtension {
fun IrBuilderWithScope.serializerInstance( fun IrBuilderWithScope.serializerInstance(
enclosingGenerator: SerializerIrGenerator, enclosingGenerator: SerializerIrGenerator,
serializableDescriptor: ClassDescriptor, serializableDescriptor: ClassDescriptor,
serializerClass: ClassDescriptor?, serializerClassOriginal: ClassDescriptor?,
module: ModuleDescriptor, module: ModuleDescriptor,
kType: KotlinType, kType: KotlinType,
genericIndex: Int? = null genericIndex: Int? = null
): IrExpression? { ): IrExpression? {
val nullableSerClass = val nullableSerClass =
compilerContext.externalSymbols.referenceClass(module.getClassFromInternalSerializationPackage(SpecialBuiltins.nullableSerializer)) compilerContext.externalSymbols.referenceClass(module.getClassFromInternalSerializationPackage(SpecialBuiltins.nullableSerializer))
if (serializerClass == null) { if (serializerClassOriginal == null) {
if (genericIndex == null) return null if (genericIndex == null) return null
val thiz = enclosingGenerator.irClass.thisReceiver!! val thiz = enclosingGenerator.irClass.thisReceiver!!
val prop = enclosingGenerator.localSerializersFieldsDescriptors[genericIndex] val prop = enclosingGenerator.localSerializersFieldsDescriptors[genericIndex]
return irGetField(irGet(thiz), compilerContext.localSymbolTable.referenceField(prop).owner) return irGetField(irGet(thiz), compilerContext.localSymbolTable.referenceField(prop).owner)
} }
if (serializerClass.kind == ClassKind.OBJECT) { if (serializerClassOriginal.kind == ClassKind.OBJECT) {
return irGetObject(serializerClass) return irGetObject(serializerClassOriginal)
} else { } else {
// todo: special instantiation of enum serializer according to new design var serializerClass = serializerClassOriginal
var args = if (serializerClass.classId == enumSerializerId || serializerClass.classId == contextSerializerId) var args: List<IrExpression> = when (serializerClassOriginal.classId) {
listOf(classReference(kType)) contextSerializerId -> listOf(classReference(kType))
else kType.arguments.map { enumSerializerId -> {
val argSer = enclosingGenerator.findTypeSerializerOrContext(module, it.type, sourceElement = serializerClass.findPsi()) serializerClass = serializableDescriptor.getClassFromInternalSerializationPackage("CommonEnumSerializer")
val expr = serializerInstance(enclosingGenerator, serializableDescriptor, argSer, module, it.type, it.type.genericIndex) kType.toClassDescriptor!!.let { enumDesc ->
?: return null listOf(
if (it.type.isMarkedNullable) irInvoke(null, nullableSerClass.constructors.toList()[0], expr) else expr irString(enumDesc.name.toString()),
irCall(findEnumValuesMethod(enumDesc)),
createArrayOfExpression(
compilerContext.irBuiltIns.stringType,
getEnumMembersNames(enumDesc).map { irString(it) }.toList()
)
)
}
}
else -> kType.arguments.map {
val argSer = enclosingGenerator.findTypeSerializerOrContext(module, it.type, sourceElement = serializerClassOriginal.findPsi())
val expr = serializerInstance(enclosingGenerator, serializableDescriptor, argSer, module, it.type, it.type.genericIndex)
?: return null
if (it.type.isMarkedNullable) irInvoke(null, nullableSerClass.constructors.toList()[0], expr) else expr
}
} }
if (serializerClass.classId == referenceArraySerializerId) if (serializerClassOriginal.classId == referenceArraySerializerId)
args = listOf(classReference(kType.arguments[0].type)) + args args = listOf(classReference(kType.arguments[0].type)) + args
val serializable = getSerializableClassDescriptorBySerializer(serializerClass) val serializable = getSerializableClassDescriptorBySerializer(serializerClass)
val ctor = if (serializable?.declaredTypeParameters?.isNotEmpty() == true) { val ctor = if (serializable?.declaredTypeParameters?.isNotEmpty() == true) {
requireNotNull( requireNotNull(
@@ -142,10 +142,9 @@ class SerializerIrGenerator(val irClass: IrClass, override val compilerContext:
serializerTower(this@SerializerIrGenerator, it)) { "Property ${it.name} must have a serializer" } serializerTower(this@SerializerIrGenerator, it)) { "Property ${it.name} must have a serializer" }
} }
val arrayOfKSerType = irFun.returnType
val kSer = serializableDescriptor.module.getClassFromSerializationPackage(KSERIALIZER_NAME.identifier) val kSer = serializableDescriptor.module.getClassFromSerializationPackage(KSERIALIZER_NAME.identifier)
val kSerType = compilerContext.externalSymbols.referenceClass(kSer).owner.defaultType val kSerType = compilerContext.externalSymbols.referenceClass(kSer).owner.defaultType
val array = createArrayOfExpression(arrayOfKSerType, kSerType, allSerializers) val array = createArrayOfExpression(kSerType, allSerializers)
+irReturn(array) +irReturn(array)
} }