- Create independent instances of MandatoryMethodTrasformer.

- Properly encapsulate LabelNormalizationMethodTransformer state.
This commit is contained in:
Dmitry Petrov
2015-07-28 12:36:34 +03:00
parent 36c88da93a
commit 641a59dcf2
5 changed files with 118 additions and 114 deletions
@@ -375,7 +375,7 @@ public class MethodInliner {
node = prepareNode(node, finallyDeepShift); node = prepareNode(node, finallyDeepShift);
try { try {
MandatoryMethodTransformer.INSTANCE$.transform("fake", node); new MandatoryMethodTransformer().transform("fake", node);
} }
catch (Throwable e) { catch (Throwable e) {
throw wrapException(e, node, "couldn't inline method call"); throw wrapException(e, node, "couldn't inline method call");
@@ -20,136 +20,144 @@ import org.jetbrains.kotlin.codegen.optimization.transformer.MethodTransformer
import org.jetbrains.org.objectweb.asm.Label import org.jetbrains.org.objectweb.asm.Label
import org.jetbrains.org.objectweb.asm.tree.* import org.jetbrains.org.objectweb.asm.tree.*
public object LabelNormalizationMethodTransformer : MethodTransformer() { public class LabelNormalizationMethodTransformer : MethodTransformer() {
val newLabelNodes = hashMapOf<LabelNode, LabelNode>()
val removedLabelNodes = hashSetOf<LabelNode>()
public override fun transform(internalClassName: String, methodNode: MethodNode) { public override fun transform(internalClassName: String, methodNode: MethodNode) {
newLabelNodes.clear() TransformerForMethod(methodNode).transform()
removedLabelNodes.clear()
with(methodNode.instructions) {
insertBefore(getFirst(), LabelNode(Label()))
insert(getLast(), LabelNode(Label()))
}
rewriteLabelInsns(methodNode)
if (removedLabelNodes.isEmpty()) return
rewriteInsns(methodNode)
rewriteTryCatchBlocks(methodNode)
rewriteLocalVars(methodNode)
} }
private fun rewriteLabelInsns(methodNode: MethodNode) { private class TransformerForMethod(val methodNode: MethodNode) {
var prevLabelNode: LabelNode? = null val instructions = methodNode.instructions
var thisNode = methodNode.instructions.getFirst() val newLabelNodes = hashMapOf<Label, LabelNode>()
while (thisNode != null) {
if (thisNode is LabelNode) { public fun transform() {
if (prevLabelNode != null) { if (rewriteLabelInstructions()) {
newLabelNodes[thisNode] = prevLabelNode rewriteNonLabelInstructions()
removedLabelNodes.add(thisNode) rewriteTryCatchBlocks()
thisNode = methodNode.instructions.removeNodeGetNext(thisNode) rewriteLocalVars()
}
}
private fun rewriteLabelInstructions(): Boolean {
var removedAnyLabels = false
var thisNode = instructions.first
while (thisNode != null) {
if (thisNode is LabelNode) {
val prevNode = thisNode.previous
if (prevNode is LabelNode) {
newLabelNodes[thisNode.label] = prevNode
removedAnyLabels = true
thisNode = instructions.removeNodeGetNext(thisNode)
}
else {
newLabelNodes[thisNode.label] = thisNode
thisNode = thisNode.next
}
} }
else { else {
prevLabelNode = thisNode thisNode = thisNode.next
newLabelNodes[thisNode] = thisNode
thisNode = thisNode.getNext()
} }
} }
else { return removedAnyLabels
prevLabelNode = null }
thisNode = thisNode.getNext()
private fun rewriteNonLabelInstructions() {
var thisNode = instructions.first
while (thisNode != null) {
thisNode = when (thisNode) {
is JumpInsnNode ->
rewriteJumpInsn(thisNode)
is LineNumberNode ->
rewriteLineNumberNode(thisNode)
is LookupSwitchInsnNode ->
rewriteLookupSwitchInsn(thisNode)
is TableSwitchInsnNode ->
rewriteTableSwitchInsn(thisNode)
is FrameNode ->
rewriteFrameNode(thisNode)
else ->
thisNode.next
}
} }
} }
}
private fun rewriteInsns(methodNode: MethodNode) { private fun rewriteLineNumberNode(oldLineNode: LineNumberNode): AbstractInsnNode? =
var thisNode = methodNode.instructions.getFirst() instructions.replaceNodeGetNext(oldLineNode, oldLineNode.rewriteLabels())
while (thisNode != null) {
thisNode = when (thisNode) {
is JumpInsnNode ->
rewriteJumpInsn(methodNode, thisNode)
is LineNumberNode ->
rewriteLineNumberNode(methodNode, thisNode)
is LookupSwitchInsnNode ->
rewriteLookupSwitchInsn(methodNode, thisNode)
is TableSwitchInsnNode ->
rewriteTableSwitchInsn(methodNode, thisNode)
is FrameNode ->
rewriteFrameNode(methodNode, thisNode)
else ->
thisNode.getNext()
}
}
}
private fun rewriteLineNumberNode(methodNode: MethodNode, oldLineNode: LineNumberNode): AbstractInsnNode? { private fun rewriteJumpInsn(oldJumpNode: JumpInsnNode): AbstractInsnNode? =
if (isRemoved(oldLineNode.start)) { instructions.replaceNodeGetNext(oldJumpNode, oldJumpNode.rewriteLabels())
val newLineNode = oldLineNode.clone(newLabelNodes)
return methodNode.instructions.replaceNodeGetNext(oldLineNode, newLineNode)
}
else {
return oldLineNode.getNext()
}
}
private fun rewriteJumpInsn(methodNode: MethodNode, oldJumpNode: JumpInsnNode): AbstractInsnNode? { private fun rewriteLookupSwitchInsn(oldSwitchNode: LookupSwitchInsnNode): AbstractInsnNode? =
if (isRemoved(oldJumpNode.label)) { instructions.replaceNodeGetNext(oldSwitchNode, oldSwitchNode.rewriteLabels())
val newJumpNode = oldJumpNode.clone(newLabelNodes)
return methodNode.instructions.replaceNodeGetNext(oldJumpNode, newJumpNode)
}
else {
return oldJumpNode.getNext()
}
}
private fun rewriteLookupSwitchInsn(methodNode: MethodNode, oldSwitchNode: LookupSwitchInsnNode): AbstractInsnNode? = private fun rewriteTableSwitchInsn(oldSwitchNode: TableSwitchInsnNode): AbstractInsnNode? =
methodNode.instructions.replaceNodeGetNext(oldSwitchNode, oldSwitchNode.clone(newLabelNodes)) instructions.replaceNodeGetNext(oldSwitchNode, oldSwitchNode.rewriteLabels())
private fun rewriteTableSwitchInsn(methodNode: MethodNode, oldSwitchNode: TableSwitchInsnNode): AbstractInsnNode? = private fun rewriteFrameNode(oldFrameNode: FrameNode): AbstractInsnNode? =
methodNode.instructions.replaceNodeGetNext(oldSwitchNode, oldSwitchNode.clone(newLabelNodes)) instructions.replaceNodeGetNext(oldFrameNode, oldFrameNode.rewriteLabels())
private fun rewriteFrameNode(methodNode: MethodNode, oldFrameNode: FrameNode): AbstractInsnNode? = private fun rewriteTryCatchBlocks() {
methodNode.instructions.replaceNodeGetNext(oldFrameNode, oldFrameNode.clone(newLabelNodes)) methodNode.tryCatchBlocks = methodNode.tryCatchBlocks.map { oldTcb ->
private fun rewriteTryCatchBlocks(methodNode: MethodNode) {
methodNode.tryCatchBlocks = methodNode.tryCatchBlocks.map { oldTcb ->
if (isRemoved(oldTcb.start) || isRemoved(oldTcb.end) || isRemoved(oldTcb.handler)) {
val newTcb = TryCatchBlockNode(getNew(oldTcb.start), getNew(oldTcb.end), getNew(oldTcb.handler), oldTcb.type) val newTcb = TryCatchBlockNode(getNew(oldTcb.start), getNew(oldTcb.end), getNew(oldTcb.handler), oldTcb.type)
newTcb.visibleTypeAnnotations = oldTcb.visibleTypeAnnotations newTcb.visibleTypeAnnotations = oldTcb.visibleTypeAnnotations
newTcb.invisibleTypeAnnotations = oldTcb.invisibleTypeAnnotations newTcb.invisibleTypeAnnotations = oldTcb.invisibleTypeAnnotations
newTcb newTcb
} }
else { }
oldTcb
private fun rewriteLocalVars() {
methodNode.localVariables = methodNode.localVariables.map { oldVar ->
LocalVariableNode(
oldVar.name,
oldVar.desc,
oldVar.signature,
getNew(oldVar.start),
getNew(oldVar.end),
oldVar.index
)
} }
} }
}
private fun rewriteLocalVars(methodNode: MethodNode) { private fun LineNumberNode.rewriteLabels(): AbstractInsnNode =
methodNode.localVariables = methodNode.localVariables.map { oldVar -> LineNumberNode(line, getNewOrOld(start))
if (isRemoved(oldVar.start) || isRemoved(oldVar.end)) {
LocalVariableNode(oldVar.name, oldVar.desc, oldVar.signature, getNew(oldVar.start), getNew(oldVar.end), oldVar.index) private fun JumpInsnNode.rewriteLabels(): AbstractInsnNode =
} JumpInsnNode(opcode, getNew(label))
else {
oldVar private fun LookupSwitchInsnNode.rewriteLabels(): AbstractInsnNode {
} val switchNode = LookupSwitchInsnNode(getNew(dflt), keys.toIntArray(), emptyArray())
switchNode.labels = labels.map { getNew(it) }
return switchNode
} }
}
private fun isRemoved(labelNode: LabelNode): Boolean = removedLabelNodes.contains(labelNode) private fun TableSwitchInsnNode.rewriteLabels(): AbstractInsnNode {
private fun getNew(oldLabelNode: LabelNode): LabelNode = newLabelNodes[oldLabelNode]!! val switchNode = TableSwitchInsnNode(min, max, getNew(dflt))
switchNode.labels = labels.map { getNew(it) }
return switchNode
}
private fun FrameNode.rewriteLabels(): AbstractInsnNode {
val frameNode = FrameNode(type, 0, emptyArray(), 0, emptyArray())
frameNode.local = local.map { if (it is LabelNode) getNewOrOld(it) else it }
frameNode.stack = stack.map { if (it is LabelNode) getNewOrOld(it) else it }
return frameNode
}
private fun getNew(oldLabelNode: LabelNode): LabelNode =
newLabelNodes[oldLabelNode.label]!!
private fun getNewOrOld(oldLabelNode: LabelNode): LabelNode =
newLabelNodes[oldLabelNode.label] ?: oldLabelNode
}
} }
private fun InsnList.replaceNodeGetNext(oldNode: AbstractInsnNode, newNode: AbstractInsnNode): AbstractInsnNode? { private fun InsnList.replaceNodeGetNext(oldNode: AbstractInsnNode, newNode: AbstractInsnNode): AbstractInsnNode? {
insertBefore(oldNode, newNode) insertBefore(oldNode, newNode)
remove(oldNode) remove(oldNode)
return newNode.getNext() return newNode.next
} }
private fun InsnList.removeNodeGetNext(oldNode: AbstractInsnNode): AbstractInsnNode? { private fun InsnList.removeNodeGetNext(oldNode: AbstractInsnNode): AbstractInsnNode? {
val next = oldNode.getNext() val next = oldNode.next
remove(oldNode) remove(oldNode)
return next return next
} }
@@ -20,9 +20,12 @@ import org.jetbrains.kotlin.codegen.optimization.fixStack.FixStackMethodTransfor
import org.jetbrains.kotlin.codegen.optimization.transformer.MethodTransformer import org.jetbrains.kotlin.codegen.optimization.transformer.MethodTransformer
import org.jetbrains.org.objectweb.asm.tree.MethodNode import org.jetbrains.org.objectweb.asm.tree.MethodNode
public object MandatoryMethodTransformer : MethodTransformer() { public class MandatoryMethodTransformer : MethodTransformer() {
private val labelNormalization = LabelNormalizationMethodTransformer()
private val fixStack = FixStackMethodTransformer()
public override fun transform(internalClassName: String, methodNode: MethodNode) { public override fun transform(internalClassName: String, methodNode: MethodNode) {
LabelNormalizationMethodTransformer.transform(internalClassName, methodNode) labelNormalization.transform(internalClassName, methodNode)
FixStackMethodTransformer.transform(internalClassName, methodNode) fixStack.transform(internalClassName, methodNode)
} }
} }
@@ -37,6 +37,8 @@ import java.util.List;
public class OptimizationMethodVisitor extends MethodVisitor { public class OptimizationMethodVisitor extends MethodVisitor {
private static final int MEMORY_LIMIT_BY_METHOD_MB = 50; private static final int MEMORY_LIMIT_BY_METHOD_MB = 50;
private static final MethodTransformer MANDATORY_METHOD_TRANSFORMER = new MandatoryMethodTransformer();
private static final MethodTransformer[] OPTIMIZATION_TRANSFORMERS = new MethodTransformer[] { private static final MethodTransformer[] OPTIMIZATION_TRANSFORMERS = new MethodTransformer[] {
new RedundantNullCheckMethodTransformer(), new RedundantNullCheckMethodTransformer(),
new RedundantBoxingMethodTransformer(), new RedundantBoxingMethodTransformer(),
@@ -75,7 +77,7 @@ public class OptimizationMethodVisitor extends MethodVisitor {
super.visitEnd(); super.visitEnd();
if (shouldBeTransformed(methodNode)) { if (shouldBeTransformed(methodNode)) {
MandatoryMethodTransformer.INSTANCE$.transform("fake", methodNode); MANDATORY_METHOD_TRANSFORMER.transform("fake", methodNode);
if (canBeOptimized(methodNode) && !disableOptimization) { if (canBeOptimized(methodNode) && !disableOptimization) {
for (MethodTransformer transformer : OPTIMIZATION_TRANSFORMERS) { for (MethodTransformer transformer : OPTIMIZATION_TRANSFORMERS) {
transformer.transform("fake", methodNode); transformer.transform("fake", methodNode);
@@ -16,24 +16,15 @@
package org.jetbrains.kotlin.codegen.optimization.fixStack package org.jetbrains.kotlin.codegen.optimization.fixStack
import com.intellij.util.containers.Stack
import org.jetbrains.kotlin.codegen.inline.InlineCodegenUtil import org.jetbrains.kotlin.codegen.inline.InlineCodegenUtil
import org.jetbrains.kotlin.codegen.optimization.common.InsnSequence import org.jetbrains.kotlin.codegen.optimization.common.InsnSequence
import org.jetbrains.kotlin.codegen.optimization.common.MethodAnalyzer
import org.jetbrains.kotlin.codegen.optimization.common.OptimizationBasicInterpreter
import org.jetbrains.kotlin.codegen.optimization.fixStack.forEachPseudoInsn
import org.jetbrains.kotlin.codegen.optimization.transformer.MethodTransformer import org.jetbrains.kotlin.codegen.optimization.transformer.MethodTransformer
import org.jetbrains.kotlin.codegen.pseudoInsns.PseudoInsn import org.jetbrains.kotlin.codegen.pseudoInsns.PseudoInsn
import org.jetbrains.kotlin.codegen.pseudoInsns.parsePseudoInsnOrNull import org.jetbrains.kotlin.codegen.pseudoInsns.parsePseudoInsnOrNull
import org.jetbrains.org.objectweb.asm.Opcodes import org.jetbrains.org.objectweb.asm.tree.AbstractInsnNode
import org.jetbrains.org.objectweb.asm.tree.* import org.jetbrains.org.objectweb.asm.tree.MethodNode
import org.jetbrains.org.objectweb.asm.tree.analysis.BasicValue
import org.jetbrains.org.objectweb.asm.tree.analysis.Frame
import org.jetbrains.org.objectweb.asm.tree.analysis.Interpreter
import java.util.*
import kotlin.properties.Delegates
public object FixStackMethodTransformer : MethodTransformer() { public class FixStackMethodTransformer : MethodTransformer() {
public override fun transform(internalClassName: String, methodNode: MethodNode) { public override fun transform(internalClassName: String, methodNode: MethodNode) {
val context = FixStackContext(methodNode) val context = FixStackContext(methodNode)