GenerateProtoBufCompare: optimization and generate difference and hashCode methods

This commit is contained in:
Michael Nedzelsky
2015-08-06 09:21:06 +03:00
parent a8e1c1f7d3
commit 4c9ec56bc8
2 changed files with 675 additions and 147 deletions
@@ -50,11 +50,19 @@ class GenerateProtoBufCompare {
)
private val RESULT_NAME = "result"
private val STRING_INDEXES_NANE = "StringIndexes"
private val FQ_NAME_INDEXES_NANE = "FqNameIndexes"
private val OLD_PREFIX = "old"
private val NEW_PREFIX = "new"
private val CHECK_EQAULS_NAME = "checkEquals"
private val HASH_CODE_NAME = "hashCode"
val extentionsMap = DebugJvmProtoBuf.getDescriptor().extensions.groupBy { it.containingType }
val doneMessages: MutableSet<Descriptors.Descriptor> = hashSetOf()
val messages: MutableList<Descriptors.Descriptor> = arrayListOf()
val repeatedFields: MutableList<Descriptors.FieldDescriptor> = arrayListOf()
val allMessages: MutableSet<Descriptors.Descriptor> = linkedSetOf()
val messagesToProcess: Queue<Descriptors.Descriptor> = linkedListOf()
val repeatedFields: MutableSet<Descriptors.FieldDescriptor> = linkedSetOf()
fun generate(): String {
val sb = StringBuilder()
@@ -63,77 +71,115 @@ class GenerateProtoBufCompare {
p.println("package org.jetbrains.kotlin.jps.incremental")
p.println()
p.println("import org.jetbrains.kotlin.name.FqName")
p.println("import org.jetbrains.kotlin.serialization.Interner")
p.println("import org.jetbrains.kotlin.serialization.ProtoBuf")
p.println("import org.jetbrains.kotlin.serialization.deserialization.NameResolver")
p.println("import org.jetbrains.kotlin.serialization.jvm.JvmProtoBuf")
p.println("import java.util.EnumSet")
p.println()
p.println("/** This file is generated by org.jetbrains.kotlin.generators.protobuf.GenerateProtoBufCompare. DO NOT MODIFY MANUALLY */")
p.println()
p.println("open class ProtoCompareGenerated(private val oldNameResolver: NameResolver, private val newNameResolver: NameResolver) {")
p.println("open class ProtoCompareGenerated(public val oldNameResolver: NameResolver, public val newNameResolver: NameResolver) {")
p.pushIndent()
p.println("private val stringIdMap: MutableMap<Int, Int> = hashMapOf()")
p.println("private val fqNameIdMap: MutableMap<Int, Int> = hashMapOf()")
p.println("private val strings = Interner<String>()")
p.println("public val $OLD_PREFIX$STRING_INDEXES_NANE: IntArray = oldNameResolver.stringTable.stringList.map { strings.intern(it) }.toIntArray()")
p.println("public val $NEW_PREFIX$STRING_INDEXES_NANE: IntArray = newNameResolver.stringTable.stringList.map { strings.intern(it) }.toIntArray()")
p.println()
p.println("private val fqNames = Interner<FqName>()")
p.println("public val $OLD_PREFIX$FQ_NAME_INDEXES_NANE: IntArray = oldNameResolver.qualifiedNameTable.qualifiedNameList.indices.map { fqNames.intern(oldNameResolver.getFqName(it)) }.toIntArray()")
p.println("public val $NEW_PREFIX$FQ_NAME_INDEXES_NANE: IntArray = newNameResolver.qualifiedNameTable.qualifiedNameList.indices.map { fqNames.intern(newNameResolver.getFqName(it)) }.toIntArray()")
p.println()
val fileDescriptor = DebugProtoBuf.getDescriptor()
messages.add(fileDescriptor.findMessageTypeByName("Package"))
messages.add(fileDescriptor.findMessageTypeByName("Class"))
addMessageToProcessIfNeeded(fileDescriptor.findMessageTypeByName("Package"))
addMessageToProcessIfNeeded(fileDescriptor.findMessageTypeByName("Class"))
val generateDifference = allMessages.toSet()
while (!messages.isEmpty()) {
val messageDescriptor = messages.remove(0)
doneMessages.add(messageDescriptor)
while (messagesToProcess.isNotEmpty()) {
p.println()
generateForMessage(messageDescriptor, p)
val message = messagesToProcess.poll()
generateForMessage(message, p)
if (message in generateDifference) {
generateDiffForMessage(message, p)
}
}
p.println()
repeatedFields.forEach { generateHelperMethodForRepeatedField(it, p) }
generatePredefined(p)
p.popIndent()
p.println("}")
allMessages.forEach { generateHashCodeFun(it, p) }
return sb.toString()
}
fun generatePredefined(p: Printer) {
p.println("fun checkStringIdEquals(old: Int, new: Int): Boolean {")
p.println(" stringIdMap.get(old)?.let { return it == new }")
p.println()
p.println(" val oldValue = oldNameResolver.stringTable.getString(old)")
p.println(" val newValue = newNameResolver.stringTable.getString(new)")
p.println()
p.println(" return if (oldValue == newValue) { stringIdMap[old] = new; true } else false")
p.println("}")
fun generateHashCodeFun(descriptor: Descriptors.Descriptor, p: Printer) {
val typeName = descriptor.typeName
val fields = descriptor.fields.filter { !it.isSkip }
val extFields = extentionsMap[descriptor]?.filter { !it.isSkip } ?: emptyList()
p.println()
p.println("fun checkNameIdEquals(old: Int, new: Int): Boolean = checkStringIdEquals(old, new)")
p.println("public fun $typeName.$HASH_CODE_NAME(stringIndexes: IntArray, fqNameIndexes: IntArray): Int {")
p.pushIndent()
p.println("var $HASH_CODE_NAME = 1")
fields.forEach { field -> generateHashCodeForField(field, p, false) }
extFields.forEach { field -> generateHashCodeForField(field, p, true) }
p.println()
p.println("fun checkFqNameIdEquals(old: Int, new: Int): Boolean {")
p.println(" fqNameIdMap.get(old)?.let { return it == new }")
p.println()
p.println(" val oldValue = oldNameResolver.getFqName(old).asString()")
p.println(" val newValue = newNameResolver.getFqName(new).asString()")
p.println()
p.println(" return if (oldValue == newValue) { fqNameIdMap[old] = new; true } else false")
p.println("return $HASH_CODE_NAME")
p.popIndent()
p.println("}")
}
fun generateForMessage(descriptor: Descriptors.Descriptor, p: Printer) {
val typeName = descriptor.typeName()
fun generateHashCodeForField(field: Descriptors.FieldDescriptor, p: Printer, isExtensionField: Boolean) {
val fieldName = field.name.javaName
val capFieldName = fieldName.capitalize()
val outerClassName = field.file.options.javaOuterClassname.removePrefix("Debug")
val fullFieldName = "$outerClassName.$fieldName"
p.println("open fun checkEquals(old: $typeName, new: $typeName): Boolean {")
val upperBound = if (isExtensionField) "getExtensionCount($fullFieldName)" else "${fieldName}Count"
val hasMethod = if (isExtensionField) "hasExtension($fullFieldName)" else "has$capFieldName()"
val fieldValue = if (isExtensionField) "getExtension($fullFieldName)" else fieldName
val repeatedFieldValue = if (isExtensionField) "getExtension($fullFieldName, i)" else "get$capFieldName(i)"
p.println()
if (field.isRepeated) {
p.println("for(i in 0..$upperBound - 1) {")
p.println(" $HASH_CODE_NAME = 31 * $HASH_CODE_NAME + ${fieldToHashCode(field, repeatedFieldValue)}")
p.println("}")
}
else if (field.isRequired) {
p.println("$HASH_CODE_NAME = 31 * $HASH_CODE_NAME + ${fieldToHashCode(field, fieldValue)}")
}
else if (field.isOptional) {
p.println("if ($hasMethod) {")
p.println(" $HASH_CODE_NAME = 31 * $HASH_CODE_NAME + ${fieldToHashCode(field, fieldValue)}")
p.println("}")
}
}
fun generateForMessage(descriptor: Descriptors.Descriptor, p: Printer) {
val typeName = descriptor.typeName
val fields = descriptor.fields.filter { !it.isSkip }
val extFields = extentionsMap[descriptor]?.filter { !it.isSkip } ?: emptyList()
p.println("open fun $CHECK_EQAULS_NAME(old: $typeName, new: $typeName): Boolean {")
p.pushIndent()
descriptor.fields.forEach { generateForField(it, p) }
extentionsMap[descriptor]?.let { it.forEach { field -> generateForExtensionField(field, p) }}
fields.forEach { field -> FieldGeneratorImpl(field, p).generate() }
extFields.forEach { field -> ExtFieldGeneratorImpl(field, p).generate() }
p.println("return true")
@@ -141,127 +187,216 @@ class GenerateProtoBufCompare {
p.println("}")
}
fun generateForField(field: Descriptors.FieldDescriptor, p: Printer) {
val fieldName = field.name.toJavaName()
val capFieldName = fieldName.capitalize()
fun generateDiffForMessage(descriptor: Descriptors.Descriptor, p: Printer) {
val typeName = descriptor.typeName
val className = typeName.replace(".", "")
if (field.options.getExtension(DebugExtOptionsProtoBuf.skipInComparison)) return
val fields = descriptor.fields.filter { !it.isSkip }
val extFields = extentionsMap[descriptor]?.filter { !it.isSkip } ?: emptyList()
val allFields = fields + extFields
if (field.isRepeated) {
repeatedFields.add(field)
p.println("if (!${field.helperMethodName()}(old, new)) return false")
}
else if (field.isRequired) {
p.printlnIfWithComparison(field, fieldName)
}
else if (field.isOptional) {
p.println("if (old.has$capFieldName() != new.has$capFieldName()) return false")
p.println("if (old.has$capFieldName()) {")
p.printlnIfWithComparison(field, fieldName, withIndent = true)
p.println("}")
}
p.println("public enum class ${className}Kind {")
p.println(allFields.map { " " + it.enumName }.join(",\n "))
p.println("}")
p.println()
p.println("public fun difference(old: $typeName, new: $typeName): EnumSet<${className}Kind> {")
p.pushIndent()
addMessageTypeToProcessIfNeeded(field)
p.println("val $RESULT_NAME = EnumSet.noneOf(javaClass<${className}Kind>())")
p.println()
fields.forEach { field -> FieldGeneratorForDiff(field, p).generate() }
extFields.forEach { field -> ExtFieldGeneratorForDiff(field, p).generate() }
p.println("return $RESULT_NAME")
p.popIndent()
p.println("}")
}
fun generateHelperMethodForRepeatedField(field: Descriptors.FieldDescriptor, p: Printer) {
assert(field.isRepeated, "expected repeated field: ${field.name}")
assert(field.isRepeated) { "expected repeated field: ${field.name}" }
val typeName = field.containingType.typeName()
val fieldName = field.name.toJavaName()
val typeName = field.containingType.typeName
val fieldName = field.name.javaName
val capFieldName = fieldName.capitalize()
val methodName = field.helperMethodName()
p.println()
p.println("open fun $methodName(old: $typeName, new: $typeName): Boolean {")
p.pushIndent()
p.println("if (old.${fieldName}Count != new.${fieldName}Count) return false")
p.println()
p.println("for(i in 0..old.${fieldName}Count - 1) {")
p.printlnIfWithComparison(field, "get$capFieldName(i)", withIndent = true)
p.printlnIfWithComparisonIndent(field, "get$capFieldName(i)")
p.println("}")
p.println()
p.println("return true")
p.popIndent()
p.println("}")
p.println()
}
fun generateForExtensionField(field: Descriptors.FieldDescriptor, p: Printer) {
val outerClassName = field.file.options.javaOuterClassname.removePrefix("Debug")
val fieldName = field.name.toJavaName()
if (field.options.getExtension(DebugExtOptionsProtoBuf.skipInComparison)) return
abstract inner class FieldGenerator(val field: Descriptors.FieldDescriptor, val p: Printer) {
val statement = field.getStatement()
abstract fun Descriptors.FieldDescriptor.getStatement(): String
fun generate() {
if (field.isRepeated) {
printRepeatedField()
}
else if (field.isRequired) {
printRequiredField()
}
else if (field.isOptional) {
printOptionalField()
}
p.println()
addMessageTypeToProcessIfNeeded(field)
}
abstract fun printRepeatedField()
abstract fun printRequiredField()
abstract fun printOptionalField()
}
open inner class FieldGeneratorImpl(field: Descriptors.FieldDescriptor, p: Printer) : FieldGenerator(field, p) {
val fieldName = field.name.javaName
val capFieldName = fieldName.capitalize()
override fun printRepeatedField() {
repeatedFields.add(field)
p.println("if (!${field.helperMethodName()}(old, new)) $statement")
}
override fun printRequiredField() {
p.printlnIfWithComparison(field, fieldName, statement)
}
override fun printOptionalField() {
p.println("if (old.has$capFieldName() != new.has$capFieldName()) $statement")
p.println("if (old.has$capFieldName()) {")
p.printlnIfWithComparisonIndent(field, fieldName, statement)
p.println("}")
}
override fun Descriptors.FieldDescriptor.getStatement(): String = "return false"
}
open inner class ExtFieldGeneratorImpl(field: Descriptors.FieldDescriptor, p: Printer) : FieldGenerator(field, p) {
val outerClassName = field.file.options.javaOuterClassname.removePrefix("Debug")
val fieldName = field.name.javaName
val fullFieldName = "$outerClassName.$fieldName"
if (field.isRepeated) {
p.println("if (old.getExtensionCount($fullFieldName) != new.getExtensionCount($fullFieldName)) return false")
override fun printRepeatedField() {
p.println("if (old.getExtensionCount($fullFieldName) != new.getExtensionCount($fullFieldName)) $statement")
p.println()
p.println("for(i in 0..old.getExtensionCount($fullFieldName) - 1) {")
p.printlnIfWithComparison(field, "getExtension($fullFieldName, i)", withIndent = true)
p.printlnIfWithComparisonIndent(field, "getExtension($fullFieldName, i)", statement)
p.println("}")
p.println()
}
else if (field.isRequired) {
p.printlnIfWithComparison(field, "getExtension($fullFieldName)")
override fun printRequiredField() {
p.printlnIfWithComparison(field, "getExtension($fullFieldName)", statement)
}
else if (field.isOptional) {
override fun printOptionalField() {
p.println("if (old.hasExtension($fullFieldName) != new.hasExtension($fullFieldName)) return false")
p.println("if (old.hasExtension($fullFieldName)) {")
p.printlnIfWithComparison(field, "getExtension($fullFieldName)", withIndent = true)
p.printlnIfWithComparisonIndent(field, "getExtension($fullFieldName)", statement)
p.println("}")
}
p.println()
addMessageTypeToProcessIfNeeded(field)
override fun Descriptors.FieldDescriptor.getStatement(): String = "return false"
}
inner class FieldGeneratorForDiff(field: Descriptors.FieldDescriptor, p: Printer) : FieldGeneratorImpl(field, p) {
override fun Descriptors.FieldDescriptor.getStatement(): String = statementForDiff
}
inner class ExtFieldGeneratorForDiff(field: Descriptors.FieldDescriptor, p: Printer) : ExtFieldGeneratorImpl(field, p) {
override fun Descriptors.FieldDescriptor.getStatement(): String = statementForDiff
}
private val Descriptors.FieldDescriptor.statementForDiff: String
get() = "$RESULT_NAME.add(${containingType.typeName.replace(".", "")}Kind.$enumName)"
private val Descriptors.FieldDescriptor.isSkip: Boolean
get() = options.getExtension(DebugExtOptionsProtoBuf.skipInComparison)
private fun addMessageTypeToProcessIfNeeded(field: Descriptors.FieldDescriptor) {
if (field.javaType == Descriptors.FieldDescriptor.JavaType.MESSAGE &&
!doneMessages.contains(field.messageType) && !messages.contains(field.messageType)) {
messages.add(field.messageType)
if (field.javaType == Descriptors.FieldDescriptor.JavaType.MESSAGE) {
addMessageToProcessIfNeeded(field.messageType)
}
}
private fun Printer.printlnIfWithComparison(field: Descriptors.FieldDescriptor, expr: String, withIndent: Boolean = false) {
val line = when {
field.options.getExtension(DebugExtOptionsProtoBuf.stringIdInTable) ->
"if (!checkStringIdEquals(old.$expr, new.$expr)) return false"
field.options.getExtension(DebugExtOptionsProtoBuf.nameIdInTable) ->
"if (!checkNameIdEquals(old.$expr, new.$expr)) return false"
field.options.getExtension(DebugExtOptionsProtoBuf.fqNameIdInTable) ->
"if (!checkFqNameIdEquals(old.$expr, new.$expr)) return false"
field.javaType in JAVA_TYPES_WITH_INLINED_EQUALS ->
"if (old.$expr != new.$expr) return false"
else ->
"if (!checkEquals(old.$expr, new.$expr)) return false"
private fun addMessageToProcessIfNeeded(descriptor: Descriptors.Descriptor) {
if (descriptor !in allMessages) {
allMessages.add(descriptor)
messagesToProcess.add(descriptor)
}
if (withIndent) {
this.pushIndent()
}
private fun Printer.printlnIfWithComparison(field: Descriptors.FieldDescriptor, expr: String, statement: String = "return false") {
val line = when {
field.options.getExtension(DebugExtOptionsProtoBuf.stringIdInTable),
field.options.getExtension(DebugExtOptionsProtoBuf.nameIdInTable) ->
"if ($OLD_PREFIX$STRING_INDEXES_NANE[old.$expr] != $NEW_PREFIX$STRING_INDEXES_NANE[new.$expr]) $statement"
field.options.getExtension(DebugExtOptionsProtoBuf.fqNameIdInTable) ->
"if ($OLD_PREFIX$FQ_NAME_INDEXES_NANE[old.$expr] != $NEW_PREFIX$FQ_NAME_INDEXES_NANE[new.$expr]) $statement"
field.javaType in JAVA_TYPES_WITH_INLINED_EQUALS ->
"if (old.$expr != new.$expr) $statement"
else ->
"if (!$CHECK_EQAULS_NAME(old.$expr, new.$expr)) $statement"
}
this.println(line)
if (withIndent) {
this.popIndent()
}
}
private fun Descriptors.Descriptor.typeName(): String {
val outerClassName = this.file.options.javaOuterClassname.removePrefix("Debug")
val packageHeader = this.file.`package`
return outerClassName + this.fullName.removePrefix(packageHeader)
fun Printer.printlnIfWithComparisonIndent(field: Descriptors.FieldDescriptor, expr: String, statement: String = "return false") {
pushIndent()
printlnIfWithComparison(field, expr, statement)
popIndent()
}
private fun fieldToHashCode(field: Descriptors.FieldDescriptor, expr: String): String =
when {
field.options.getExtension(DebugExtOptionsProtoBuf.stringIdInTable),
field.options.getExtension(DebugExtOptionsProtoBuf.nameIdInTable) ->
"stringIndexes[$expr]"
field.options.getExtension(DebugExtOptionsProtoBuf.fqNameIdInTable) ->
"fqNameIndexes[$expr]"
field.javaType == Descriptors.FieldDescriptor.JavaType.INT ->
"$expr"
field.javaType in JAVA_TYPES_WITH_INLINED_EQUALS ->
"$expr.$HASH_CODE_NAME()"
else ->
"$expr.$HASH_CODE_NAME(stringIndexes, fqNameIndexes)"
}
private val Descriptors.Descriptor.typeName: String
get() {
val outerClassName = file.options.javaOuterClassname.removePrefix("Debug")
val packageHeader = file.`package`
return outerClassName + fullName.removePrefix(packageHeader)
}
private val Descriptors.FieldDescriptor.enumName: String
get() = (name.javaName + (if (isRepeated) "List" else "")).replace("[A-Z]".toRegex()) { "_" + it.value }.toUpperCase()
private fun Descriptors.FieldDescriptor.helperMethodName(): String {
val packageHeader = this.file.`package`
val descriptor = this.containingType
val className = descriptor.fullName.removePrefix(packageHeader).replace(".", "")
val capFieldName = this.name.toJavaName().capitalize()
return "checkEquals$className$capFieldName"
val capFieldName = this.name.javaName.capitalize()
return "$CHECK_EQAULS_NAME$className$capFieldName"
}
private fun String.toJavaName() = this.split("_").map { it.capitalize() }.join("").decapitalize()
private val String.javaName: String
get() = this.split("_").map { it.capitalize() }.join("").decapitalize()
}