save state
diff --git a/compiler/ir/backend.js/src/org/jetbrains/kotlin/ir/backend/js/BackendJsSymbols.kt b/compiler/ir/backend.js/src/org/jetbrains/kotlin/ir/backend/js/BackendJsSymbols.kt
index db0ff0f..39407509 100644
--- a/compiler/ir/backend.js/src/org/jetbrains/kotlin/ir/backend/js/BackendJsSymbols.kt
+++ b/compiler/ir/backend.js/src/org/jetbrains/kotlin/ir/backend/js/BackendJsSymbols.kt
@@ -520,7 +520,7 @@
 
     val constructCallableReferenceSymbol by CallableIds.constructCallableReference.functionSymbol()
 
-    val staticInitializationFailureWithClassName by CallableIds.staticInitializationFailureWithClassName.functionSymbol()
+    val checkStaticInitializationState by CallableIds.checkStaticInitializationState.functionSymbol()
 }
 
 private object ClassIds {
@@ -790,5 +790,5 @@
     val test = CallableId(StandardClassIds.BASE_TEST_PACKAGE, Name.identifier("test"))
     val suite = CallableId(StandardClassIds.BASE_TEST_PACKAGE, Name.identifier("suite"))
     val EmptyContinuation = CallableId(FqName.fromSegments(listOf("kotlin", "coroutines", "js", "internal")), Name.identifier("EmptyContinuation"))
-    val staticInitializationFailureWithClassName = "staticInitializationFailureWithClassName".jsCallableId
+    val checkStaticInitializationState = "checkStaticInitializationState".jsCallableId
 }
diff --git a/compiler/ir/backend.js/src/org/jetbrains/kotlin/ir/backend/js/lower/JsStaticInitializersDeclarationLowering.kt b/compiler/ir/backend.js/src/org/jetbrains/kotlin/ir/backend/js/lower/JsStaticInitializersDeclarationLowering.kt
index 7c74fb4..a95edb0 100644
--- a/compiler/ir/backend.js/src/org/jetbrains/kotlin/ir/backend/js/lower/JsStaticInitializersDeclarationLowering.kt
+++ b/compiler/ir/backend.js/src/org/jetbrains/kotlin/ir/backend/js/lower/JsStaticInitializersDeclarationLowering.kt
@@ -5,16 +5,26 @@
 
 package org.jetbrains.kotlin.ir.backend.js.lower
 
+import org.jetbrains.kotlin.backend.common.lower.createIrBuilder
+import org.jetbrains.kotlin.backend.common.lower.irCatch
+import org.jetbrains.kotlin.backend.common.lower.irIfThen
 import org.jetbrains.kotlin.backend.common.phaser.PhasePrerequisites
+import org.jetbrains.kotlin.descriptors.DescriptorVisibilities
+import org.jetbrains.kotlin.ir.IrStatement
+import org.jetbrains.kotlin.ir.UNDEFINED_OFFSET
 import org.jetbrains.kotlin.ir.backend.js.JsIrBackendContext
+import org.jetbrains.kotlin.ir.backend.js.staticInitFunction
 import org.jetbrains.kotlin.ir.backend.js.utils.getVoid
 import org.jetbrains.kotlin.ir.backend.js.utils.jsConstructorReference
-import org.jetbrains.kotlin.ir.builders.IrBuilderWithScope
-import org.jetbrains.kotlin.ir.builders.irCall
-import org.jetbrains.kotlin.ir.declarations.IrClass
+import org.jetbrains.kotlin.ir.builders.*
+import org.jetbrains.kotlin.ir.builders.declarations.buildFun
+import org.jetbrains.kotlin.ir.declarations.*
 import org.jetbrains.kotlin.ir.expressions.IrCall
 import org.jetbrains.kotlin.ir.expressions.IrExpression
+import org.jetbrains.kotlin.ir.expressions.IrGetField
 import org.jetbrains.kotlin.ir.types.IrType
+import org.jetbrains.kotlin.ir.util.*
+import org.jetbrains.kotlin.name.Name
 
 @PhasePrerequisites(
     ObjectDeclarationLowering::class,
@@ -22,13 +32,41 @@
     EnumEntryCreateGetInstancesFunsLowering::class,
 )
 class JsStaticInitializersDeclarationLowering(override val context: JsIrBackendContext) : WebStaticInitializersDeclarationLowering() {
-    override fun IrBuilderWithScope.generateStaticInitializationFailureCallWithClassName(container: IrClass): IrCall =
-        irCall(this@JsStaticInitializersDeclarationLowering.context.symbols.staticInitializationFailureWithClassName).apply {
-            arguments[0] = container.jsConstructorReference(this@JsStaticInitializersDeclarationLowering.context)
+    private fun IrBuilderWithScope.generateStaticInitializationStateCheck(getStateField: IrGetField, container: IrClass): IrCall =
+        irCall(this@JsStaticInitializersDeclarationLowering.context.symbols.checkStaticInitializationState).apply {
+            arguments[0] = getStateField
+            arguments[1] = container.jsConstructorReference(this@JsStaticInitializersDeclarationLowering.context)
         }
 
     override fun IrBuilderWithScope.undefinedOrNull(): IrExpression = this@JsStaticInitializersDeclarationLowering.context.getVoid()
 
     override val catchParameterType: IrType
         get() = context.dynamicType
+
+    protected override fun createStaticInitFunction(
+        container: IrClass,
+        origin: IrDeclarationOrigin,
+        initCalledVar: IrField,
+        initializers: List<IrStatement>
+    ): IrSimpleFunction {
+        val initFunction = context.irFactory.buildFun {
+            startOffset = UNDEFINED_OFFSET
+            endOffset = UNDEFINED_OFFSET
+            this.origin = origin
+            name = Name.identifier(STATIC_INIT_FUNCTION_NAME)
+            visibility = DescriptorVisibilities.PUBLIC
+            returnType = context.irBuiltIns.unitType
+        }
+        return initFunction.apply {
+            val builder = context.createIrBuilder(symbol, SYNTHETIC_OFFSET)
+            parent = container
+            body = context.irFactory.createBlockBody(startOffset, endOffset) {
+                with(builder) {
+                    val stateCheck = generateStaticInitializationStateCheck(irGetField(null, initCalledVar), container)
+                    statements += irIfThen(stateCheck, irReturnUnit())
+                    initializationBody(initFunction, container, initCalledVar, initializers)
+                }
+            }
+        }
+    }
 }
diff --git a/compiler/ir/backend.js/src/org/jetbrains/kotlin/ir/backend/js/lower/WebStaticInitializersDeclarationLowering.kt b/compiler/ir/backend.js/src/org/jetbrains/kotlin/ir/backend/js/lower/WebStaticInitializersDeclarationLowering.kt
index 66614b8..9c1dcda 100644
--- a/compiler/ir/backend.js/src/org/jetbrains/kotlin/ir/backend/js/lower/WebStaticInitializersDeclarationLowering.kt
+++ b/compiler/ir/backend.js/src/org/jetbrains/kotlin/ir/backend/js/lower/WebStaticInitializersDeclarationLowering.kt
@@ -15,7 +15,6 @@
 import org.jetbrains.kotlin.ir.backend.js.*
 import org.jetbrains.kotlin.ir.builders.*
 import org.jetbrains.kotlin.ir.builders.declarations.buildField
-import org.jetbrains.kotlin.ir.builders.declarations.buildFun
 import org.jetbrains.kotlin.ir.declarations.*
 import org.jetbrains.kotlin.ir.expressions.*
 import org.jetbrains.kotlin.ir.types.IrType
@@ -196,7 +195,7 @@
             )
         }
 
-    private object InitializationState {
+    protected object InitializationState {
         const val UNINITIALIZED: Int = 0
         const val INITIALIZED: Int = 1
         const val ERROR: Int = 2
@@ -217,95 +216,19 @@
         )
     }
 
-    protected abstract fun IrBuilderWithScope.generateStaticInitializationFailureCallWithClassName(container: IrClass): IrCall
-
     protected open fun IrBuilderWithScope.undefinedOrNull(): IrExpression = irNull()
 
     protected open val catchParameterType: IrType
         get() = context.irBuiltIns.throwableType
 
-    private fun createStaticInitFunction(
+    protected abstract fun createStaticInitFunction(
         container: IrClass,
         origin: IrDeclarationOrigin,
         initCalledVar: IrField,
         initializers: List<IrStatement>
-    ): IrSimpleFunction {
-        val initFunction = context.irFactory.buildFun {
-            startOffset = UNDEFINED_OFFSET
-            endOffset = UNDEFINED_OFFSET
-            this.origin = origin
-            name = Name.identifier(STATIC_INIT_FUNCTION_NAME)
-            visibility = DescriptorVisibilities.PUBLIC
-            returnType = context.irBuiltIns.unitType
-        }
-        return initFunction.apply {
-            val builder = context.createIrBuilder(symbol, SYNTHETIC_OFFSET)
-            parent = container
-            body = context.irFactory.createBlockBody(startOffset, endOffset) {
-                with(builder) {
-                    // Need a temporary variable for Wasm to transform if/then branches with br_table
-                    val initState = scope.createTemporaryVariable(
-                        irGetField(null, initCalledVar),
-                        nameHint = "initState",
-                        inventUniqueName = false,
-                    )
-                    statements += initState
-                    statements += irWhen(
-                        context.irBuiltIns.unitType,
-                        listOf(
-                            // Already initialized successfully - early branch.
-                            irBranch(
-                                irEqeqeq(irGet(initState), irInt(InitializationState.INITIALIZED)),
-                                irReturnUnit()
-                            ),
-                            // Previously attempted initialization failed with error.
-                            irBranch(
-                                irEqeqeq(irGet(initState), irInt(InitializationState.ERROR)),
-                                generateStaticInitializationFailureCallWithClassName(container)
-                            ),
-                            // Initialization hasn't been performed yet - try to initialize.
-                            irElseBranch(
-                                irBlock {
-                                    +irSetField(null, initCalledVar, irInt(InitializationState.INITIALIZED))
-                                    val allInitializers = irComposite {
-                                        val [dependencySuperInterfaces, dependencySuperClasses] =
-                                            container.dependencySuperClasses.partition { it.isInterface }
-                                        for (superClass in dependencySuperClasses + dependencySuperInterfaces) {
-                                            superClass.staticInitFunction?.let {
-                                                +irCall(it.symbol)
-                                            }
-                                        }
-                                        for (initializer in initializers) {
-                                            initializer.setDeclarationsParent(initFunction)
-                                        }
-                                        +initializers
-                                    }
-                                    val catchParameter = scope.createTemporaryVariableDeclaration(
-                                        irType = catchParameterType,
-                                        nameHint = "reason",
-                                        origin = IrDeclarationOrigin.CATCH_PARAMETER,
-                                        startOffset = UNDEFINED_OFFSET,
-                                        endOffset = UNDEFINED_OFFSET,
-                                        inventUniqueName = false,
-                                    )
-                                    val catchResult = irComposite {
-                                        +irSetField(null, initCalledVar, irInt(InitializationState.ERROR))
-                                        +irCall(this@WebStaticInitializersDeclarationLowering.context.symbols.staticInitializationFailure).apply {
-                                            arguments[0] = irCastIfNeeded(irGet(catchParameter), context.irBuiltIns.throwableType)
-                                            arguments[1] = undefinedOrNull()
-                                        }
-                                    }
-                                    +irTry(context.irBuiltIns.unitType, allInitializers, listOf(irCatch(catchParameter, catchResult)), null)
-                                }
-                            )
-                        )
-                    )
-                }
-            }
-        }
-    }
+    ): IrSimpleFunction
 
-    private val IrClass.dependencySuperClasses: List<IrClass>
+    protected val IrClass.dependencySuperClasses: List<IrClass>
         get() = superTypes
             .filter { !it.isAny() }
             .mapNotNull { it.classOrNull?.owner }
@@ -313,6 +236,44 @@
             // its initialization from the implementing class. See section §3.3 of the KEEP.
             .filter { clazz -> !clazz.isInterface || clazz.declarations.any { it.isNonAbstractInstanceMember() } }
 
+    protected fun IrBuilderWithScope.initializationBody(
+        initFunction: IrSimpleFunction,
+        container: IrClass,
+        initCalledVar: IrField,
+        initializers: List<IrStatement>
+    ) = irBlock {
+        +irSetField(null, initCalledVar, irInt(InitializationState.INITIALIZED))
+        val allInitializers = irComposite {
+            val [dependencySuperInterfaces, dependencySuperClasses] =
+                container.dependencySuperClasses.partition { it.isInterface }
+            for (superClass in dependencySuperClasses + dependencySuperInterfaces) {
+                superClass.staticInitFunction?.let {
+                    +irCall(it.symbol)
+                }
+            }
+            for (initializer in initializers) {
+                initializer.setDeclarationsParent(initFunction)
+            }
+            +initializers
+        }
+        val catchParameter = scope.createTemporaryVariableDeclaration(
+            irType = catchParameterType,
+            nameHint = "reason",
+            origin = IrDeclarationOrigin.CATCH_PARAMETER,
+            startOffset = UNDEFINED_OFFSET,
+            endOffset = UNDEFINED_OFFSET,
+            inventUniqueName = false,
+        )
+        val catchResult = irComposite {
+            +irSetField(null, initCalledVar, irInt(InitializationState.ERROR))
+            +irCall(this@WebStaticInitializersDeclarationLowering.context.symbols.staticInitializationFailure).apply {
+                arguments[0] = irCastIfNeeded(irGet(catchParameter), context.irBuiltIns.throwableType)
+                arguments[1] = undefinedOrNull()
+            }
+        }
+        +irTry(context.irBuiltIns.unitType, allInitializers, listOf(irCatch(catchParameter, catchResult)), null)
+    }
+
     private fun IrDeclaration.isNonAbstractInstanceMember(): Boolean = when (this) {
         is IrSimpleFunction if isReal && modality != Modality.ABSTRACT && dispatchReceiverParameter != null -> true
         is IrProperty if isReal && modality != Modality.ABSTRACT && (getter ?: setter)?.dispatchReceiverParameter != null -> true
diff --git a/compiler/ir/backend.wasm/src/org/jetbrains/kotlin/backend/wasm/lower/WasmStaticInitializersDeclarationLowering.kt b/compiler/ir/backend.wasm/src/org/jetbrains/kotlin/backend/wasm/lower/WasmStaticInitializersDeclarationLowering.kt
index 150240d..ed814be 100644
--- a/compiler/ir/backend.wasm/src/org/jetbrains/kotlin/backend/wasm/lower/WasmStaticInitializersDeclarationLowering.kt
+++ b/compiler/ir/backend.wasm/src/org/jetbrains/kotlin/backend/wasm/lower/WasmStaticInitializersDeclarationLowering.kt
@@ -5,18 +5,20 @@
 
 package org.jetbrains.kotlin.backend.wasm.lower
 
+import org.jetbrains.kotlin.backend.common.lower.createIrBuilder
 import org.jetbrains.kotlin.backend.common.phaser.PhasePrerequisites
 import org.jetbrains.kotlin.backend.wasm.WasmBackendContext
-import org.jetbrains.kotlin.ir.backend.js.lower.EnumEntryCreateGetInstancesFunsLowering
-import org.jetbrains.kotlin.ir.backend.js.lower.EnumEntryInstancesLowering
-import org.jetbrains.kotlin.ir.backend.js.lower.ObjectDeclarationLowering
-import org.jetbrains.kotlin.ir.backend.js.lower.WebStaticInitializersDeclarationLowering
-import org.jetbrains.kotlin.ir.builders.IrBuilderWithScope
-import org.jetbrains.kotlin.ir.builders.irCall
-import org.jetbrains.kotlin.ir.builders.kClassReference
-import org.jetbrains.kotlin.ir.declarations.IrClass
+import org.jetbrains.kotlin.descriptors.DescriptorVisibilities
+import org.jetbrains.kotlin.ir.IrStatement
+import org.jetbrains.kotlin.ir.UNDEFINED_OFFSET
+import org.jetbrains.kotlin.ir.backend.js.lower.*
+import org.jetbrains.kotlin.ir.builders.*
+import org.jetbrains.kotlin.ir.builders.declarations.buildFun
+import org.jetbrains.kotlin.ir.declarations.*
 import org.jetbrains.kotlin.ir.expressions.IrCall
 import org.jetbrains.kotlin.ir.types.starProjectedType
+import org.jetbrains.kotlin.ir.util.*
+import org.jetbrains.kotlin.name.Name
 
 @PhasePrerequisites(
     ObjectDeclarationLowering::class,
@@ -24,8 +26,58 @@
     EnumEntryCreateGetInstancesFunsLowering::class,
 )
 class WasmStaticInitializersDeclarationLowering(override val context: WasmBackendContext) : WebStaticInitializersDeclarationLowering() {
-    override fun IrBuilderWithScope.generateStaticInitializationFailureCallWithClassName(container: IrClass): IrCall =
+    private fun IrBuilderWithScope.generateStaticInitializationFailureCallWithClassName(container: IrClass): IrCall =
         irCall(this@WasmStaticInitializersDeclarationLowering.context.symbols.staticInitializationFailureWithClassName).apply {
             arguments[0] = kClassReference(container.symbol.starProjectedType)
         }
+
+    protected override fun createStaticInitFunction(
+        container: IrClass,
+        origin: IrDeclarationOrigin,
+        initCalledVar: IrField,
+        initializers: List<IrStatement>
+    ): IrSimpleFunction {
+        val initFunction = context.irFactory.buildFun {
+            startOffset = UNDEFINED_OFFSET
+            endOffset = UNDEFINED_OFFSET
+            this.origin = origin
+            name = Name.identifier(STATIC_INIT_FUNCTION_NAME)
+            visibility = DescriptorVisibilities.PUBLIC
+            returnType = context.irBuiltIns.unitType
+        }
+        return initFunction.apply {
+            val builder = context.createIrBuilder(symbol, SYNTHETIC_OFFSET)
+            parent = container
+            body = context.irFactory.createBlockBody(startOffset, endOffset) {
+                with(builder) {
+                    // Need a temporary variable for Wasm to transform if/then branches with br_table
+                    val initState = scope.createTemporaryVariable(
+                        irGetField(null, initCalledVar),
+                        nameHint = "initState",
+                        inventUniqueName = false,
+                    )
+                    statements += initState
+                    statements += irWhen(
+                        context.irBuiltIns.unitType,
+                        listOf(
+                            // Already initialized successfully - early branch.
+                            irBranch(
+                                irEqeqeq(irGet(initState), irInt(InitializationState.INITIALIZED)),
+                                irReturnUnit()
+                            ),
+                            // Previously attempted initialization failed with error.
+                            irBranch(
+                                irEqeqeq(irGet(initState), irInt(InitializationState.ERROR)),
+                                generateStaticInitializationFailureCallWithClassName(container)
+                            ),
+                            // Initialization hasn't been performed yet - try to initialize.
+                            irElseBranch(
+                                initializationBody(initFunction, container, initCalledVar, initializers)
+                            )
+                        )
+                    )
+                }
+            }
+        }
+    }
 }
diff --git a/libraries/stdlib/js/runtime/staticInitialization.kt b/libraries/stdlib/js/runtime/staticInitialization.kt
index b73dc1e..d7a3c9b 100644
--- a/libraries/stdlib/js/runtime/staticInitialization.kt
+++ b/libraries/stdlib/js/runtime/staticInitialization.kt
@@ -8,7 +8,13 @@
 import kotlin.internal.UsedFromCompilerGeneratedCode
 import kotlin.internal.staticInitializationFailure
 
+private const val INITIALIZATION_STATE_INITIALIZED: Int = 1
+private const val INITIALIZATION_STATE_ERROR: Int = 2
+
 @UsedFromCompilerGeneratedCode
-internal fun staticInitializationFailureWithClassName(ctor: Ctor?) {
-    staticInitializationFailure(null, ctor?.`$metadata$`?.simpleName)
+internal fun checkStaticInitializationState(state: Int, ctor: Ctor?): Boolean {
+    if (state == INITIALIZATION_STATE_ERROR) {
+        staticInitializationFailure(null, ctor?.`$metadata$`?.simpleName)
+    }
+    return state == INITIALIZATION_STATE_INITIALIZED
 }