[WasmJs] Fix kotlin.test AfterTest for test with Promise return type

Fix #KT-61888
This commit is contained in:
Igor Yakovlev
2024-01-16 20:32:17 +01:00
committed by Space Team
parent 78b666bb40
commit b6ef1e08e1
8 changed files with 56 additions and 28 deletions
@@ -35,6 +35,8 @@ interface JsCommonBackendContext : CommonBackendContext {
val coroutineSymbols: JsCommonCoroutineSymbols val coroutineSymbols: JsCommonCoroutineSymbols
val jsPromiseSymbol: IrClassSymbol?
val catchAllThrowableType: IrType val catchAllThrowableType: IrType
get() = irBuiltIns.throwableType get() = irBuiltIns.throwableType
@@ -196,6 +196,9 @@ class JsIrBackendContext(
override val coroutineSymbols = override val coroutineSymbols =
JsCommonCoroutineSymbols(symbolTable, module, this) JsCommonCoroutineSymbols(symbolTable, module, this)
override val jsPromiseSymbol: IrClassSymbol?
get() = intrinsics.promiseClassSymbol
override val enumEntries = getIrClass(ENUMS_PACKAGE_FQNAME.child(Name.identifier("EnumEntries"))) override val enumEntries = getIrClass(ENUMS_PACKAGE_FQNAME.child(Name.identifier("EnumEntries")))
override val createEnumEntries = getFunctions(ENUMS_PACKAGE_FQNAME.child(Name.identifier("enumEntries"))) override val createEnumEntries = getFunctions(ENUMS_PACKAGE_FQNAME.child(Name.identifier("enumEntries")))
.find { it.valueParameters.firstOrNull()?.type?.isFunctionType == false } .find { it.valueParameters.firstOrNull()?.type?.isFunctionType == false }
@@ -131,7 +131,7 @@ fun compileIr(
} }
// TODO should be done incrementally // TODO should be done incrementally
generateJsTests(context, allModules.last()) generateJsTests(context, allModules.last(), groupByPackage = false)
(irFactory.stageController as? WholeWorldStageController)?.let { (irFactory.stageController as? WholeWorldStageController)?.let {
lowerPreservingTags(allModules, context, phaseConfig, it) lowerPreservingTags(allModules, context, phaseConfig, it)
@@ -58,7 +58,7 @@ class JsIrCompilerWithIC(
moveBodilessDeclarationsToSeparatePlace(context, it) moveBodilessDeclarationsToSeparatePlace(context, it)
} }
generateJsTests(context, mainModule) generateJsTests(context, mainModule, groupByPackage = false)
lowerPreservingTags(allModules, context, phaseConfig, context.irFactory.stageController as WholeWorldStageController) lowerPreservingTags(allModules, context, phaseConfig, context.irFactory.stageController as WholeWorldStageController)
@@ -12,7 +12,6 @@ import org.jetbrains.kotlin.descriptors.DescriptorVisibilities
import org.jetbrains.kotlin.descriptors.Modality import org.jetbrains.kotlin.descriptors.Modality
import org.jetbrains.kotlin.ir.UNDEFINED_OFFSET import org.jetbrains.kotlin.ir.UNDEFINED_OFFSET
import org.jetbrains.kotlin.ir.backend.js.JsCommonBackendContext import org.jetbrains.kotlin.ir.backend.js.JsCommonBackendContext
import org.jetbrains.kotlin.ir.backend.js.JsIrBackendContext
import org.jetbrains.kotlin.ir.backend.js.ir.JsIrBuilder import org.jetbrains.kotlin.ir.backend.js.ir.JsIrBuilder
import org.jetbrains.kotlin.ir.backend.js.utils.getJsModule import org.jetbrains.kotlin.ir.backend.js.utils.getJsModule
import org.jetbrains.kotlin.ir.backend.js.utils.getJsNameOrKotlinName import org.jetbrains.kotlin.ir.backend.js.utils.getJsNameOrKotlinName
@@ -21,32 +20,36 @@ import org.jetbrains.kotlin.ir.builders.irCall
import org.jetbrains.kotlin.ir.builders.irString import org.jetbrains.kotlin.ir.builders.irString
import org.jetbrains.kotlin.ir.declarations.* import org.jetbrains.kotlin.ir.declarations.*
import org.jetbrains.kotlin.ir.expressions.IrBlockBody import org.jetbrains.kotlin.ir.expressions.IrBlockBody
import org.jetbrains.kotlin.ir.expressions.IrDynamicOperator
import org.jetbrains.kotlin.ir.expressions.IrExpression import org.jetbrains.kotlin.ir.expressions.IrExpression
import org.jetbrains.kotlin.ir.expressions.IrTypeOperator
import org.jetbrains.kotlin.ir.expressions.impl.IrConstructorCallImpl import org.jetbrains.kotlin.ir.expressions.impl.IrConstructorCallImpl
import org.jetbrains.kotlin.ir.expressions.impl.IrDynamicMemberExpressionImpl import org.jetbrains.kotlin.ir.expressions.impl.IrTypeOperatorCallImpl
import org.jetbrains.kotlin.ir.expressions.impl.IrDynamicOperatorExpressionImpl
import org.jetbrains.kotlin.ir.symbols.IrClassSymbol
import org.jetbrains.kotlin.ir.symbols.IrSimpleFunctionSymbol import org.jetbrains.kotlin.ir.symbols.IrSimpleFunctionSymbol
import org.jetbrains.kotlin.ir.types.IrSimpleType import org.jetbrains.kotlin.ir.types.IrSimpleType
import org.jetbrains.kotlin.ir.types.defaultType
import org.jetbrains.kotlin.ir.types.impl.IrSimpleTypeImpl import org.jetbrains.kotlin.ir.types.impl.IrSimpleTypeImpl
import org.jetbrains.kotlin.ir.util.* import org.jetbrains.kotlin.ir.util.*
import org.jetbrains.kotlin.name.FqName import org.jetbrains.kotlin.name.FqName
import org.jetbrains.kotlin.name.Name import org.jetbrains.kotlin.name.Name
import org.jetbrains.kotlin.types.checker.SimpleClassicTypeSystemContext.asSimpleType
import org.jetbrains.kotlin.utils.filterIsInstanceAnd import org.jetbrains.kotlin.utils.filterIsInstanceAnd
import org.jetbrains.kotlin.utils.findIsInstanceAnd import org.jetbrains.kotlin.utils.findIsInstanceAnd
import org.jetbrains.kotlin.utils.memoryOptimizedMap import org.jetbrains.kotlin.utils.memoryOptimizedMap
fun generateJsTests(context: JsIrBackendContext, moduleFragment: IrModuleFragment) { fun generateJsTests(context: JsCommonBackendContext, moduleFragment: IrModuleFragment, groupByPackage: Boolean) {
val generator = TestGenerator(context, false) val generator = TestGenerator(
context = context,
groupByPackage = groupByPackage
)
moduleFragment.files.toList().forEach { moduleFragment.files.toList().forEach {
generator.lower(it) generator.lower(it)
} }
} }
class TestGenerator(val context: JsCommonBackendContext, val groupByPackage: Boolean) : FileLoweringPass { private class TestGenerator(
val context: JsCommonBackendContext,
private val groupByPackage: Boolean,
) : FileLoweringPass {
override fun lower(irFile: IrFile) { override fun lower(irFile: IrFile) {
// Additional copy to prevent ConcurrentModificationException // Additional copy to prevent ConcurrentModificationException
@@ -186,9 +189,21 @@ class TestGenerator(val context: JsCommonBackendContext, val groupByPackage: Boo
return return
} }
val returnType = testFun.returnType as? IrSimpleType
if (context is JsIrBackendContext && (testFun.returnType as? IrSimpleType)?.isPromise == true) { val promiseSymbol = context.jsPromiseSymbol
val refType = IrSimpleTypeImpl(context.ir.symbols.functionN(0), false, emptyList(), emptyList()) if (promiseSymbol != null && returnType != null && returnType.isPromise) {
val promiseCastedIfNeeded = if (returnType.classifier == context.jsPromiseSymbol) {
returnStatement.value
} else {
IrTypeOperatorCallImpl(
startOffset = UNDEFINED_OFFSET,
endOffset = UNDEFINED_OFFSET,
type = testFun.returnType,
operator = IrTypeOperator.CAST,
typeOperand = promiseSymbol.defaultType,
argument = returnStatement.value
)
}
val afterFunction = context.irFactory.buildFun { val afterFunction = context.irFactory.buildFun {
this.name = Name.identifier("${irClass.name.asString()} after test fun") this.name = Name.identifier("${irClass.name.asString()} after test fun")
@@ -207,16 +222,17 @@ class TestGenerator(val context: JsCommonBackendContext, val groupByPackage: Boo
) )
} }
val refType = IrSimpleTypeImpl(context.ir.symbols.functionN(0), false, emptyList(), emptyList())
val finallyLambda = JsIrBuilder.buildFunctionExpression(refType, afterFunction) val finallyLambda = JsIrBuilder.buildFunctionExpression(refType, afterFunction)
val finally = promiseSymbol.owner.declarations
.findIsInstanceAnd<IrSimpleFunction> { it.name.asString() == "finally" }!!
val finally = val returnValue = JsIrBuilder.buildCall(
IrDynamicMemberExpressionImpl(UNDEFINED_OFFSET, UNDEFINED_OFFSET, context.dynamicType, "finally", returnStatement.value) finally.symbol
).apply {
val returnValue = this.dispatchReceiver = promiseCastedIfNeeded
IrDynamicOperatorExpressionImpl(UNDEFINED_OFFSET, UNDEFINED_OFFSET, context.dynamicType, IrDynamicOperator.INVOKE).also { putValueArgument(0, finallyLambda)
it.receiver = finally }
it.arguments.add(finallyLambda)
}
body.statements += JsIrBuilder.buildReturn( body.statements += JsIrBuilder.buildReturn(
fn.symbol, fn.symbol,
@@ -271,6 +287,7 @@ class TestGenerator(val context: JsCommonBackendContext, val groupByPackage: Boo
private val IrSimpleType.isPromise: Boolean private val IrSimpleType.isPromise: Boolean
get() { get() {
if (this.classifier == context.jsPromiseSymbol) return true
val klass = classifier.owner as? IrClass ?: return false val klass = classifier.owner as? IrClass ?: return false
return klass.isExternal && klass.getJsModule() == null && klass.getJsNameOrKotlinName().asString() == "Promise" return klass.isExternal && klass.getJsModule() == null && klass.getJsNameOrKotlinName().asString() == "Promise"
} }
@@ -22,6 +22,7 @@ import org.jetbrains.kotlin.ir.builders.declarations.buildFun
import org.jetbrains.kotlin.ir.declarations.* import org.jetbrains.kotlin.ir.declarations.*
import org.jetbrains.kotlin.ir.declarations.impl.IrExternalPackageFragmentImpl import org.jetbrains.kotlin.ir.declarations.impl.IrExternalPackageFragmentImpl
import org.jetbrains.kotlin.ir.declarations.impl.IrFileImpl import org.jetbrains.kotlin.ir.declarations.impl.IrFileImpl
import org.jetbrains.kotlin.ir.symbols.IrClassSymbol
import org.jetbrains.kotlin.ir.symbols.IrSimpleFunctionSymbol import org.jetbrains.kotlin.ir.symbols.IrSimpleFunctionSymbol
import org.jetbrains.kotlin.ir.symbols.impl.DescriptorlessExternalPackageFragmentSymbol import org.jetbrains.kotlin.ir.symbols.impl.DescriptorlessExternalPackageFragmentSymbol
import org.jetbrains.kotlin.ir.types.IrTypeSystemContext import org.jetbrains.kotlin.ir.types.IrTypeSystemContext
@@ -73,6 +74,9 @@ class WasmBackendContext(
override val coroutineSymbols = override val coroutineSymbols =
JsCommonCoroutineSymbols(symbolTable, module,this) JsCommonCoroutineSymbols(symbolTable, module,this)
override val jsPromiseSymbol: IrClassSymbol?
get() = if (isWasmJsTarget) wasmSymbols.jsRelatedSymbols.jsPromise else null
val innerClassesSupport = JsInnerClassesSupport(mapping, irFactory) val innerClassesSupport = JsInnerClassesSupport(mapping, irFactory)
override val internalPackageFqn = FqName("kotlin.wasm") override val internalPackageFqn = FqName("kotlin.wasm")
@@ -385,6 +385,8 @@ class WasmSymbols(
val externRefIsNull = getInternalFunction("wasm_externref_is_null") val externRefIsNull = getInternalFunction("wasm_externref_is_null")
val jsPromise = getIrClass(FqName("kotlin.js.Promise"))
internal val throwAsJsException: IrSimpleFunctionSymbol = internal val throwAsJsException: IrSimpleFunctionSymbol =
getInternalFunction("throwAsJsException") getInternalFunction("throwAsJsException")
} }
@@ -7,17 +7,17 @@ package org.jetbrains.kotlin.backend.wasm.lower
import org.jetbrains.kotlin.backend.common.lower.createIrBuilder import org.jetbrains.kotlin.backend.common.lower.createIrBuilder
import org.jetbrains.kotlin.backend.wasm.WasmBackendContext import org.jetbrains.kotlin.backend.wasm.WasmBackendContext
import org.jetbrains.kotlin.ir.backend.js.lower.TestGenerator
import org.jetbrains.kotlin.ir.builders.irCall import org.jetbrains.kotlin.ir.builders.irCall
import org.jetbrains.kotlin.ir.backend.js.lower.generateJsTests
import org.jetbrains.kotlin.ir.declarations.IrModuleFragment import org.jetbrains.kotlin.ir.declarations.IrModuleFragment
import org.jetbrains.kotlin.ir.expressions.IrBlockBody import org.jetbrains.kotlin.ir.expressions.IrBlockBody
fun generateWasmTests(context: WasmBackendContext, moduleFragment: IrModuleFragment) { fun generateWasmTests(context: WasmBackendContext, moduleFragment: IrModuleFragment) {
val generator = TestGenerator(context, true) generateJsTests(
context = context,
moduleFragment.files.toList().forEach { moduleFragment = moduleFragment,
generator.lower(it) groupByPackage = true
} )
if (context.testEntryPoints.isEmpty()) if (context.testEntryPoints.isEmpty())
return return