[WASM] DCE implementation

This commit is contained in:
Igor Yakovlev
2021-12-29 19:28:22 +01:00
committed by TeamCityServer
parent 84caad7ba2
commit 2ec0411a7f
22 changed files with 1005 additions and 593 deletions
@@ -7,6 +7,7 @@ package org.jetbrains.kotlin.backend.wasm
import org.jetbrains.kotlin.backend.common.phaser.PhaseConfig
import org.jetbrains.kotlin.backend.common.phaser.invokeToplevel
import org.jetbrains.kotlin.backend.wasm.dce.eliminateDeadDeclarations
import org.jetbrains.kotlin.backend.wasm.ir2wasm.WasmCompiledModuleFragment
import org.jetbrains.kotlin.backend.wasm.ir2wasm.WasmModuleFragmentGenerator
import org.jetbrains.kotlin.backend.wasm.lower.markExportedDeclarations
@@ -31,6 +32,7 @@ fun compileWasm(
exportedDeclarations: Set<FqName> = emptySet(),
propertyLazyInitialization: Boolean,
emitNameSection: Boolean = false,
dceEnabled: Boolean = false,
): WasmCompilerResult {
val mainModule = depsDescriptors.mainModule
val configuration = depsDescriptors.compilerConfiguration
@@ -69,8 +71,12 @@ fun compileWasm(
wasmPhases.invokeToplevel(phaseConfig, context, moduleFragment)
if (dceEnabled) {
eliminateDeadDeclarations(listOf(moduleFragment), context)
}
val compiledWasmModule = WasmCompiledModuleFragment(context.irBuiltIns)
val codeGenerator = WasmModuleFragmentGenerator(context, compiledWasmModule)
val codeGenerator = WasmModuleFragmentGenerator(context, compiledWasmModule, allowIncompleteImplementations = dceEnabled)
codeGenerator.generateModule(moduleFragment)
val linkedModule = compiledWasmModule.linkWasmCompiledFragments()
@@ -0,0 +1,64 @@
/*
* Copyright 2010-2020 JetBrains s.r.o. and Kotlin Programming Language contributors.
* Use of this source code is governed by the Apache 2.0 license that can be found in the license/LICENSE.txt file.
*/
package org.jetbrains.kotlin.backend.wasm.dce
import org.jetbrains.kotlin.backend.wasm.WasmBackendContext
import org.jetbrains.kotlin.ir.IrElement
import org.jetbrains.kotlin.ir.backend.js.utils.*
import org.jetbrains.kotlin.ir.declarations.*
import org.jetbrains.kotlin.ir.expressions.IrBody
import org.jetbrains.kotlin.ir.visitors.IrElementVisitorVoid
import org.jetbrains.kotlin.ir.visitors.acceptChildrenVoid
import org.jetbrains.kotlin.ir.visitors.acceptVoid
import org.jetbrains.kotlin.js.config.JSConfigurationKeys
internal fun eliminateDeadDeclarations(modules: List<IrModuleFragment>, context: WasmBackendContext) {
val printReachabilityInfo =
context.configuration.getBoolean(JSConfigurationKeys.PRINT_REACHABILITY_INFO) ||
java.lang.Boolean.getBoolean("kotlin.wasm.dce.print.reachability.info")
val usefulDeclarations = WasmUsefulDeclarationProcessor(
context = context,
printReachabilityInfo = printReachabilityInfo
).collectDeclarations(rootDeclarations = buildRoots(modules, context))
val remover = WasmUselessDeclarationsRemover(usefulDeclarations)
modules.onAllFiles {
acceptVoid(remover)
}
}
private fun buildRoots(modules: List<IrModuleFragment>, context: WasmBackendContext): List<IrDeclaration> = buildList {
val declarationsCollector = object : IrElementVisitorVoid {
override fun visitElement(element: IrElement): Unit = element.acceptChildrenVoid(this)
override fun visitBody(body: IrBody): Unit = Unit // Skip
override fun visitDeclaration(declaration: IrDeclarationBase) {
super.visitDeclaration(declaration)
add(declaration)
}
}
modules.onAllFiles {
declarations.forEach { declaration ->
if (declaration.isJsExport()) {
declaration.acceptVoid(declarationsCollector)
}
}
}
add(context.irBuiltIns.throwableClass.owner)
add(context.mainCallsWrapperFunction)
add(context.fieldInitFunction)
}
private inline fun List<IrModuleFragment>.onAllFiles(body: IrFile.() -> Unit) {
forEach { module ->
module.files.forEach { file ->
file.body()
}
}
}
@@ -0,0 +1,205 @@
/*
* Copyright 2010-2021 JetBrains s.r.o. and Kotlin Programming Language contributors.
* Use of this source code is governed by the Apache 2.0 license that can be found in the license/LICENSE.txt file.
*/
package org.jetbrains.kotlin.backend.wasm.dce
import org.jetbrains.kotlin.backend.common.ir.isOverridable
import org.jetbrains.kotlin.backend.wasm.WasmBackendContext
import org.jetbrains.kotlin.backend.wasm.ir2wasm.*
import org.jetbrains.kotlin.backend.wasm.utils.*
import org.jetbrains.kotlin.ir.backend.js.dce.UsefulDeclarationProcessor
import org.jetbrains.kotlin.ir.backend.js.utils.*
import org.jetbrains.kotlin.ir.declarations.*
import org.jetbrains.kotlin.ir.expressions.*
import org.jetbrains.kotlin.ir.types.*
import org.jetbrains.kotlin.ir.util.*
import org.jetbrains.kotlin.utils.addToStdlib.firstIsInstanceOrNull
import org.jetbrains.kotlin.wasm.ir.*
internal class WasmUsefulDeclarationProcessor(
override val context: WasmBackendContext,
printReachabilityInfo: Boolean
) : UsefulDeclarationProcessor(printReachabilityInfo, removeUnusedAssociatedObjects = false) {
private val unitGetInstance: IrSimpleFunction = context.findUnitGetInstanceFunction()
override val bodyVisitor: BodyVisitorBase = object : BodyVisitorBase() {
override fun visitConst(expression: IrConst<*>, data: IrDeclaration) = when (expression.kind) {
is IrConstKind.Null -> expression.type.enqueueType(data, "expression type")
is IrConstKind.String -> context.wasmSymbols.stringGetLiteral.owner
.enqueue(data, "String literal intrinsic getter stringGetLiteral")
else -> Unit
}
private fun tryToProcessIntrinsicCall(from: IrDeclaration, call: IrCall): Boolean = when (call.symbol) {
context.wasmSymbols.unboxIntrinsic -> {
val fromType = call.getTypeArgument(0)
if (fromType != null && !fromType.isNothing() && !fromType.isNullableNothing()) {
val backingField = call.getTypeArgument(1)
?.let { context.inlineClassesUtils.getInlinedClass(it) }
?.let { getInlineClassBackingField(it) }
backingField?.enqueue(from, "backing inline class field for unboxIntrinsic")
}
true
}
context.wasmSymbols.wasmClassId,
context.wasmSymbols.wasmInterfaceId,
context.wasmSymbols.wasmRefCast -> {
call.getTypeArgument(0)?.getClass()?.enqueue(from, "generic intrinsic ${call.symbol.owner.name}")
true
}
else -> false
}
private fun tryToProcessWasmOpIntrinsicCall(from: IrDeclaration, call: IrCall, function: IrFunction): Boolean {
if (function.hasWasmNoOpCastAnnotation()) {
return true
}
val opString = function.getWasmOpAnnotation()
if (opString != null) {
val op = WasmOp.valueOf(opString)
when (op.immediates.size) {
0 -> {
if (op == WasmOp.REF_TEST) {
call.getTypeArgument(0)?.enqueueRuntimeClassOrAny(from, "REF_TEST")
}
}
1 -> {
if (op.immediates.firstOrNull() == WasmImmediateKind.STRUCT_TYPE_IDX) {
function.dispatchReceiverParameter?.type?.classOrNull?.owner?.enqueue(from, "STRUCT_TYPE_IDX")
}
}
}
return true
}
return false
}
override fun visitCall(expression: IrCall, data: IrDeclaration) {
super.visitCall(expression, data)
if (expression.symbol == context.wasmSymbols.boxIntrinsic) {
expression.getTypeArgument(0)?.enqueueRuntimeClassOrAny(data, "boxIntrinsic")
return
}
val function: IrFunction = expression.symbol.owner.realOverrideTarget
if (function.returnType == context.irBuiltIns.unitType) {
unitGetInstance.enqueue(data, "function Unit return type")
}
if (tryToProcessIntrinsicCall(data, expression)) return
if (tryToProcessWasmOpIntrinsicCall(data, expression, function)) return
val isSuperCall = expression.superQualifierSymbol != null
if (function is IrSimpleFunction && function.isOverridable && !isSuperCall) {
val klass = function.parentAsClass
if (!klass.isInterface) {
context.wasmSymbols.getVirtualMethodId.owner.enqueue(data, "call on class receiver")
} else {
klass.enqueue(data, "receiver class")
context.wasmSymbols.getInterfaceImplId.owner.enqueue(data, "call on interface receiver")
}
function.enqueue(data, "method call")
}
}
}
private fun IrType.getInlinedValueTypeIfAny(): IrType? = when (this) {
context.irBuiltIns.booleanType,
context.irBuiltIns.byteType,
context.irBuiltIns.shortType,
context.irBuiltIns.charType,
context.irBuiltIns.booleanType,
context.irBuiltIns.byteType,
context.irBuiltIns.shortType,
context.irBuiltIns.intType,
context.irBuiltIns.charType,
context.irBuiltIns.longType,
context.irBuiltIns.floatType,
context.irBuiltIns.doubleType,
context.irBuiltIns.nothingType,
context.wasmSymbols.voidType -> null
else -> when {
isBuiltInWasmRefType(this) -> null
erasedUpperBound?.isExternal == true -> null
else -> when (val ic = context.inlineClassesUtils.getInlinedClass(this)) {
null -> this
else -> context.inlineClassesUtils.getInlineClassUnderlyingType(ic).getInlinedValueTypeIfAny()
}
}
}
private fun IrType.enqueueRuntimeClassOrAny(from: IrDeclaration, info: String): Unit =
(this.getRuntimeClass ?: context.wasmSymbols.any.owner).enqueue(from, info, isContagious = false)
private fun IrType.enqueueType(from: IrDeclaration, info: String) {
getInlinedValueTypeIfAny()
?.enqueueRuntimeClassOrAny(from, info)
}
private fun IrDeclaration.enqueueParentClass() {
parentClassOrNull?.enqueue(this, "parent class", isContagious = false)
}
override fun processField(irField: IrField) {
super.processField(irField)
irField.enqueueParentClass()
irField.type.enqueueType(irField, "field types")
}
override fun processClass(irClass: IrClass) {
super.processClass(irClass)
irClass.getWasmArrayAnnotation()?.type
?.enqueueType(irClass, "array type for wasm array annotated")
if (context.inlineClassesUtils.isClassInlineLike(irClass)) {
irClass.declarations
.firstIsInstanceOrNull<IrConstructor>()
?.takeIf { it.isPrimary }
?.enqueue(irClass, "inline class primary ctor")
}
}
private fun IrValueParameter.enqueueValueParameterType(from: IrDeclaration) {
if (context.inlineClassesUtils.shouldValueParameterBeBoxed(this)) {
type.enqueueRuntimeClassOrAny(from, "function ValueParameterType")
} else {
type.enqueueType(from, "function ValueParameterType")
}
}
private fun processIrFunction(irFunction: IrFunction) {
if (irFunction.isFakeOverride) return
val isIntrinsic = irFunction.hasWasmNoOpCastAnnotation() || irFunction.getWasmOpAnnotation() != null
if (isIntrinsic) return
irFunction.getEffectiveValueParameters().forEach { it.enqueueValueParameterType(irFunction) }
irFunction.returnType.enqueueType(irFunction, "function return type")
}
override fun processSimpleFunction(irFunction: IrSimpleFunction) {
super.processSimpleFunction(irFunction)
irFunction.enqueueParentClass()
if (irFunction.isFakeOverride) {
irFunction.overriddenSymbols.forEach { overridden ->
overridden.owner.enqueue(irFunction, "original for fake-override")
}
}
processIrFunction(irFunction)
}
override fun processConstructor(irConstructor: IrConstructor) {
super.processConstructor(irConstructor)
if (!context.inlineClassesUtils.isClassInlineLike(irConstructor.parentAsClass)) {
processIrFunction(irConstructor)
}
}
override fun isExported(declaration: IrDeclaration): Boolean = declaration.isJsExport()
}
@@ -0,0 +1,44 @@
/*
* Copyright 2010-2022 JetBrains s.r.o. and Kotlin Programming Language contributors.
* Use of this source code is governed by the Apache 2.0 license that can be found in the license/LICENSE.txt file.
*/
package org.jetbrains.kotlin.backend.wasm.dce
import org.jetbrains.kotlin.ir.IrElement
import org.jetbrains.kotlin.ir.declarations.IrClass
import org.jetbrains.kotlin.ir.declarations.IrDeclaration
import org.jetbrains.kotlin.ir.declarations.IrDeclarationContainer
import org.jetbrains.kotlin.ir.declarations.IrFile
import org.jetbrains.kotlin.ir.util.transformFlat
import org.jetbrains.kotlin.ir.visitors.IrElementVisitorVoid
import org.jetbrains.kotlin.ir.visitors.acceptChildrenVoid
import org.jetbrains.kotlin.ir.visitors.acceptVoid
class WasmUselessDeclarationsRemover(
private val usefulDeclarations: Set<IrDeclaration>
) : IrElementVisitorVoid {
override fun visitElement(element: IrElement) {
element.acceptChildrenVoid(this)
}
override fun visitFile(declaration: IrFile) {
process(declaration)
}
override fun visitClass(declaration: IrClass) {
process(declaration)
}
// TODO bring back the primary constructor fix
private fun process(container: IrDeclarationContainer) {
container.declarations.transformFlat { member ->
if (member !in usefulDeclarations) {
emptyList()
} else {
member.acceptVoid(this)
null
}
}
}
}
@@ -271,7 +271,6 @@ class BodyGenerator(val context: WasmFunctionCodegenContext) : IrElementVisitorV
}
}
if (tryToGenerateIntrinsicCall(call, function)) {
if (function.returnType == irBuiltIns.unitType)
body.buildGetUnit()
@@ -284,7 +283,7 @@ class BodyGenerator(val context: WasmFunctionCodegenContext) : IrElementVisitorV
val klass = function.parentAsClass
if (!klass.isInterface) {
val classMetadata = context.getClassMetadata(klass.symbol)
val vfSlot = classMetadata.virtualMethods.map { it.function }.indexOf(function)
val vfSlot = classMetadata.virtualMethods.indexOfFirst { it.function == function }
// Dispatch receiver should be simple and without side effects at this point
// TODO: Verify
generateExpression(call.dispatchReceiver!!)
@@ -30,7 +30,7 @@ import org.jetbrains.kotlin.ir.visitors.acceptVoid
import org.jetbrains.kotlin.name.parentOrNull
import org.jetbrains.kotlin.wasm.ir.*
class DeclarationGenerator(val context: WasmModuleCodegenContext) : IrElementVisitorVoid {
class DeclarationGenerator(val context: WasmModuleCodegenContext, private val allowIncompleteImplementations: Boolean) : IrElementVisitorVoid {
// Shortcuts
private val backendContext: WasmBackendContext = context.backendContext
@@ -238,7 +238,6 @@ class DeclarationGenerator(val context: WasmModuleCodegenContext) : IrElementVis
wasmExpressionGenerator.buildRttCanon(wasmGcType)
}
val rtt = WasmGlobal(
name = "rtt_of_$nameStr",
isMutable = false,
@@ -258,13 +257,16 @@ class DeclarationGenerator(val context: WasmModuleCodegenContext) : IrElementVis
// TODO: Cache it
val interfaceMetadata = InterfaceMetadata(i, irBuiltIns)
val table = interfaceMetadata.methods.associate { method ->
val classMethod: VirtualMethodMetadata =
metadata.virtualMethods
.find { it.signature == method.signature } // TODO: Use map
?: error("Cannot find class implementation of method ${method.signature} in class ${declaration.fqNameWhenAvailable}")
val classMethod: VirtualMethodMetadata? = metadata.virtualMethods
.find { it.signature == method.signature && it.function.modality != Modality.ABSTRACT } // TODO: Use map
method.function.symbol as IrFunctionSymbol to context.referenceFunction(classMethod.function.symbol)
if (classMethod == null && !allowIncompleteImplementations) {
error("Cannot find class implementation of method ${method.signature} in class ${declaration.fqNameWhenAvailable}")
}
val matchedMethod = classMethod?.let { context.referenceFunction(it.function.symbol) }
method.function.symbol as IrFunctionSymbol to matchedMethod
}
context.registerInterfaceImplementationMethod(
interfaceImplementation,
table
@@ -74,7 +74,7 @@ class WasmCompiledModuleFragment(val irBuiltIns: IrBuiltIns) {
ReferencableElements<InterfaceImplementation, Int>()
val interfaceImplementationsMethods =
LinkedHashMap<InterfaceImplementation, Map<IrFunctionSymbol, WasmSymbol<WasmFunction>>>()
LinkedHashMap<InterfaceImplementation, Map<IrFunctionSymbol, WasmSymbol<WasmFunction>?>>()
val exports = mutableListOf<WasmExport<*>>()
@@ -83,6 +83,7 @@ class WasmCompiledModuleFragment(val irBuiltIns: IrBuiltIns) {
val jsFuns = mutableListOf<JsCodeSnippet>()
class FunWithPriority(val function: WasmFunction, val priority: String)
val initFunctions = mutableListOf<FunWithPriority>()
val scratchMemAddr = WasmSymbol<Int>()
@@ -229,28 +230,34 @@ class WasmCompiledModuleFragment(val irBuiltIns: IrBuiltIns) {
)
val interfaceTableElementsLists = interfaceMethodTables.defined.keys.associateWith {
mutableMapOf<Int, WasmSymbol<WasmFunction>>()
mutableMapOf<Int, WasmSymbol<WasmFunction>?>()
}
interfaceImplementationIds.forEach { ii: InterfaceImplementation, implId: Int ->
for ((interfaceFunction: IrFunctionSymbol, wasmFunction: WasmSymbol<WasmFunction>) in interfaceImplementationsMethods[ii]!!) {
for ((ii: InterfaceImplementation, implId: Int) in interfaceImplementationIds) {
for ((interfaceFunction: IrFunctionSymbol, wasmFunction: WasmSymbol<WasmFunction>?) in interfaceImplementationsMethods[ii]!!) {
interfaceTableElementsLists[interfaceFunction]!![implId] = wasmFunction
}
}
val interfaceTableElements = interfaceTableElementsLists.map { (interfaceFunction, methods) ->
val type = interfaceMethodTables.defined[interfaceFunction]!!.elementType
val methodTable = interfaceMethodTables.defined[interfaceFunction]!!
val type = methodTable.elementType
val functions = MutableList(methods.size) { idx ->
val wasmFunc = methods[idx]!!
val wasmFunc = methods[idx]
val expression = buildWasmExpression {
buildInstr(WasmOp.REF_FUNC, WasmImmediate.FuncIdx(wasmFunc))
if (wasmFunc != null) {
buildInstr(WasmOp.REF_FUNC, WasmImmediate.FuncIdx(wasmFunc))
} else {
//DCE could remove implementation from class, so we should to put a stub into method implementations table
buildRefNull(type.getHeapType())
}
}
WasmTable.Value.Expression(expression)
}
WasmElement(
type,
values = functions,
WasmElement.Mode.Active(interfaceMethodTables.defined[interfaceFunction]!!, offsetExpr)
WasmElement.Mode.Active(methodTable, offsetExpr)
)
}
@@ -35,7 +35,7 @@ interface WasmModuleCodegenContext : WasmBaseCodegenContext {
fun registerInterfaceImplementationMethod(
interfaceImplementation: InterfaceImplementation,
table: Map<IrFunctionSymbol, WasmSymbol<WasmFunction>>,
table: Map<IrFunctionSymbol, WasmSymbol<WasmFunction>?>,
)
fun referenceInterfaceImplementationId(interfaceImplementation: InterfaceImplementation): WasmSymbol<Int>
@@ -126,7 +126,7 @@ class WasmModuleCodegenContextImpl(
override fun registerInterfaceImplementationMethod(
interfaceImplementation: InterfaceImplementation,
table: Map<IrFunctionSymbol, WasmSymbol<WasmFunction>>
table: Map<IrFunctionSymbol, WasmSymbol<WasmFunction>?>
) {
wasmFragment.interfaceImplementationsMethods[interfaceImplementation] = table
}
@@ -13,14 +13,16 @@ import org.jetbrains.kotlin.ir.visitors.acceptVoid
class WasmModuleFragmentGenerator(
backendContext: WasmBackendContext,
wasmModuleFragment: WasmCompiledModuleFragment
wasmModuleFragment: WasmCompiledModuleFragment,
allowIncompleteImplementations: Boolean,
) {
private val declarationGenerator =
DeclarationGenerator(
WasmModuleCodegenContextImpl(
backendContext,
wasmModuleFragment
)
wasmModuleFragment,
),
allowIncompleteImplementations
)
fun generateModule(irModuleFragment: IrModuleFragment) {