- 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) private class TransformerForMethod(val methodNode: MethodNode) {
if (removedLabelNodes.isEmpty()) return val instructions = methodNode.instructions
val newLabelNodes = hashMapOf<Label, LabelNode>()
rewriteInsns(methodNode) public fun transform() {
rewriteTryCatchBlocks(methodNode) if (rewriteLabelInstructions()) {
rewriteLocalVars(methodNode) rewriteNonLabelInstructions()
rewriteTryCatchBlocks()
rewriteLocalVars()
}
} }
private fun rewriteLabelInsns(methodNode: MethodNode) { private fun rewriteLabelInstructions(): Boolean {
var prevLabelNode: LabelNode? = null var removedAnyLabels = false
var thisNode = methodNode.instructions.getFirst() var thisNode = instructions.first
while (thisNode != null) { while (thisNode != null) {
if (thisNode is LabelNode) { if (thisNode is LabelNode) {
if (prevLabelNode != null) { val prevNode = thisNode.previous
newLabelNodes[thisNode] = prevLabelNode if (prevNode is LabelNode) {
removedLabelNodes.add(thisNode) newLabelNodes[thisNode.label] = prevNode
thisNode = methodNode.instructions.removeNodeGetNext(thisNode) removedAnyLabels = true
thisNode = instructions.removeNodeGetNext(thisNode)
} }
else { else {
prevLabelNode = thisNode newLabelNodes[thisNode.label] = thisNode
newLabelNodes[thisNode] = thisNode thisNode = thisNode.next
thisNode = thisNode.getNext()
} }
} }
else { else {
prevLabelNode = null thisNode = thisNode.next
thisNode = thisNode.getNext()
} }
} }
return removedAnyLabels
} }
private fun rewriteInsns(methodNode: MethodNode) { private fun rewriteNonLabelInstructions() {
var thisNode = methodNode.instructions.getFirst() var thisNode = instructions.first
while (thisNode != null) { while (thisNode != null) {
thisNode = when (thisNode) { thisNode = when (thisNode) {
is JumpInsnNode -> is JumpInsnNode ->
rewriteJumpInsn(methodNode, thisNode) rewriteJumpInsn(thisNode)
is LineNumberNode -> is LineNumberNode ->
rewriteLineNumberNode(methodNode, thisNode) rewriteLineNumberNode(thisNode)
is LookupSwitchInsnNode -> is LookupSwitchInsnNode ->
rewriteLookupSwitchInsn(methodNode, thisNode) rewriteLookupSwitchInsn(thisNode)
is TableSwitchInsnNode -> is TableSwitchInsnNode ->
rewriteTableSwitchInsn(methodNode, thisNode) rewriteTableSwitchInsn(thisNode)
is FrameNode -> is FrameNode ->
rewriteFrameNode(methodNode, thisNode) rewriteFrameNode(thisNode)
else -> else ->
thisNode.getNext() thisNode.next
} }
} }
} }
private fun rewriteLineNumberNode(methodNode: MethodNode, oldLineNode: LineNumberNode): AbstractInsnNode? { private fun rewriteLineNumberNode(oldLineNode: LineNumberNode): AbstractInsnNode? =
if (isRemoved(oldLineNode.start)) { instructions.replaceNodeGetNext(oldLineNode, oldLineNode.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 rewriteJumpInsn(oldJumpNode: JumpInsnNode): AbstractInsnNode? =
if (isRemoved(oldJumpNode.label)) { instructions.replaceNodeGetNext(oldJumpNode, oldJumpNode.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 rewriteLookupSwitchInsn(oldSwitchNode: LookupSwitchInsnNode): AbstractInsnNode? =
methodNode.instructions.replaceNodeGetNext(oldSwitchNode, oldSwitchNode.clone(newLabelNodes)) instructions.replaceNodeGetNext(oldSwitchNode, oldSwitchNode.rewriteLabels())
private fun rewriteTableSwitchInsn(methodNode: MethodNode, oldSwitchNode: TableSwitchInsnNode): AbstractInsnNode? = private fun rewriteTableSwitchInsn(oldSwitchNode: TableSwitchInsnNode): AbstractInsnNode? =
methodNode.instructions.replaceNodeGetNext(oldSwitchNode, oldSwitchNode.clone(newLabelNodes)) instructions.replaceNodeGetNext(oldSwitchNode, oldSwitchNode.rewriteLabels())
private fun rewriteFrameNode(methodNode: MethodNode, oldFrameNode: FrameNode): AbstractInsnNode? = private fun rewriteFrameNode(oldFrameNode: FrameNode): AbstractInsnNode? =
methodNode.instructions.replaceNodeGetNext(oldFrameNode, oldFrameNode.clone(newLabelNodes)) instructions.replaceNodeGetNext(oldFrameNode, oldFrameNode.rewriteLabels())
private fun rewriteTryCatchBlocks(methodNode: MethodNode) { private fun rewriteTryCatchBlocks() {
methodNode.tryCatchBlocks = methodNode.tryCatchBlocks.map { oldTcb -> 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: MethodNode) { private fun rewriteLocalVars() {
methodNode.localVariables = methodNode.localVariables.map { oldVar -> methodNode.localVariables = methodNode.localVariables.map { oldVar ->
if (isRemoved(oldVar.start) || isRemoved(oldVar.end)) { LocalVariableNode(
LocalVariableNode(oldVar.name, oldVar.desc, oldVar.signature, getNew(oldVar.start), getNew(oldVar.end), oldVar.index) oldVar.name,
} oldVar.desc,
else { oldVar.signature,
oldVar getNew(oldVar.start),
} getNew(oldVar.end),
oldVar.index
)
} }
} }
private fun isRemoved(labelNode: LabelNode): Boolean = removedLabelNodes.contains(labelNode) private fun LineNumberNode.rewriteLabels(): AbstractInsnNode =
private fun getNew(oldLabelNode: LabelNode): LabelNode = newLabelNodes[oldLabelNode]!! LineNumberNode(line, getNewOrOld(start))
private fun JumpInsnNode.rewriteLabels(): AbstractInsnNode =
JumpInsnNode(opcode, getNew(label))
private fun LookupSwitchInsnNode.rewriteLabels(): AbstractInsnNode {
val switchNode = LookupSwitchInsnNode(getNew(dflt), keys.toIntArray(), emptyArray())
switchNode.labels = labels.map { getNew(it) }
return switchNode
}
private fun TableSwitchInsnNode.rewriteLabels(): AbstractInsnNode {
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)