[JS_IR] Use IrModuleFragment as input to lowering instead of iterable

This way we will be able to reuse existing utils
to create a lowering.

#KT-63073
This commit is contained in:
Ivan Kylchik
2023-09-21 14:39:00 +02:00
committed by Space Team
parent 1eacd5efc2
commit 0a6f711a41
5 changed files with 43 additions and 38 deletions
@@ -5,7 +5,10 @@
package org.jetbrains.kotlin.ir.backend.js package org.jetbrains.kotlin.ir.backend.js
import org.jetbrains.kotlin.backend.common.* import org.jetbrains.kotlin.backend.common.BodyLoweringPass
import org.jetbrains.kotlin.backend.common.DeclarationTransformer
import org.jetbrains.kotlin.backend.common.FileLoweringPass
import org.jetbrains.kotlin.backend.common.lower
import org.jetbrains.kotlin.backend.common.lower.* import org.jetbrains.kotlin.backend.common.lower.*
import org.jetbrains.kotlin.backend.common.lower.coroutines.AddContinuationToLocalSuspendFunctionsLowering import org.jetbrains.kotlin.backend.common.lower.coroutines.AddContinuationToLocalSuspendFunctionsLowering
import org.jetbrains.kotlin.backend.common.lower.coroutines.AddContinuationToNonLocalSuspendFunctionsLowering import org.jetbrains.kotlin.backend.common.lower.coroutines.AddContinuationToNonLocalSuspendFunctionsLowering
@@ -24,19 +27,14 @@ import org.jetbrains.kotlin.ir.backend.js.lower.coroutines.JsSuspendArityStoreLo
import org.jetbrains.kotlin.ir.backend.js.lower.coroutines.JsSuspendFunctionsLowering import org.jetbrains.kotlin.ir.backend.js.lower.coroutines.JsSuspendFunctionsLowering
import org.jetbrains.kotlin.ir.backend.js.lower.inline.* import org.jetbrains.kotlin.ir.backend.js.lower.inline.*
import org.jetbrains.kotlin.ir.backend.js.transformers.irToJs.JsGenerationGranularity import org.jetbrains.kotlin.ir.backend.js.transformers.irToJs.JsGenerationGranularity
import org.jetbrains.kotlin.ir.declarations.IrFile
import org.jetbrains.kotlin.ir.declarations.IrModuleFragment import org.jetbrains.kotlin.ir.declarations.IrModuleFragment
import org.jetbrains.kotlin.ir.interpreter.IrInterpreterConfiguration import org.jetbrains.kotlin.ir.interpreter.IrInterpreterConfiguration
import org.jetbrains.kotlin.platform.js.JsPlatforms import org.jetbrains.kotlin.platform.js.JsPlatforms
private fun DeclarationContainerLoweringPass.runOnFilesPostfix(files: Iterable<IrFile>) = files.forEach { runOnFilePostfix(it) }
private fun ClassLoweringPass.runOnFilesPostfix(moduleFragment: IrModuleFragment) = moduleFragment.files.forEach { runOnFilePostfix(it) }
private fun List<Lowering>.toCompilerPhase() = private fun List<Lowering>.toCompilerPhase() =
map { map {
@Suppress("USELESS_CAST") @Suppress("USELESS_CAST")
it.modulePhase as CompilerPhase<JsIrBackendContext, Iterable<IrModuleFragment>, Iterable<IrModuleFragment>> it.modulePhase as CompilerPhase<JsIrBackendContext, IrModuleFragment, IrModuleFragment>
}.reduce { acc, lowering -> acc.then(lowering) } }.reduce { acc, lowering -> acc.then(lowering) }
private fun makeJsModulePhase( private fun makeJsModulePhase(
@@ -44,7 +42,7 @@ private fun makeJsModulePhase(
name: String, name: String,
description: String, description: String,
prerequisite: Set<AbstractNamedCompilerPhase<JsIrBackendContext, *, *>> = emptySet() prerequisite: Set<AbstractNamedCompilerPhase<JsIrBackendContext, *, *>> = emptySet()
): SameTypeNamedCompilerPhase<JsIrBackendContext, Iterable<IrModuleFragment>> = makeCustomJsModulePhase( ): SameTypeNamedCompilerPhase<JsIrBackendContext, IrModuleFragment> = makeCustomJsModulePhase(
op = { context, modules -> lowering(context).lower(modules) }, op = { context, modules -> lowering(context).lower(modules) },
name = name, name = name,
description = description, description = description,
@@ -56,40 +54,34 @@ private fun makeCustomJsModulePhase(
description: String, description: String,
name: String, name: String,
prerequisite: Set<AbstractNamedCompilerPhase<JsIrBackendContext, *, *>> = emptySet() prerequisite: Set<AbstractNamedCompilerPhase<JsIrBackendContext, *, *>> = emptySet()
): SameTypeNamedCompilerPhase<JsIrBackendContext, Iterable<IrModuleFragment>> = SameTypeNamedCompilerPhase( ): SameTypeNamedCompilerPhase<JsIrBackendContext, IrModuleFragment> = SameTypeNamedCompilerPhase(
name = name, name = name,
description = description, description = description,
prerequisite = prerequisite, prerequisite = prerequisite,
lower = object : SameTypeCompilerPhase<JsIrBackendContext, Iterable<IrModuleFragment>> { lower = object : SameTypeCompilerPhase<JsIrBackendContext, IrModuleFragment> {
override fun invoke( override fun invoke(
phaseConfig: PhaseConfigurationService, phaseConfig: PhaseConfigurationService,
phaserState: PhaserState<Iterable<IrModuleFragment>>, phaserState: PhaserState<IrModuleFragment>,
context: JsIrBackendContext, context: JsIrBackendContext,
input: Iterable<IrModuleFragment> input: IrModuleFragment
): Iterable<IrModuleFragment> { ): IrModuleFragment {
input.forEach { module -> op(context, input)
op(context, module)
}
return input return input
} }
}, },
actions = setOf(defaultDumper.toMultiModuleAction(), validationAction.toMultiModuleAction()), actions = setOf(defaultDumper, validationAction),
) )
sealed class Lowering(val name: String) { sealed class Lowering(val name: String) {
abstract val modulePhase: SameTypeNamedCompilerPhase<JsIrBackendContext, Iterable<IrModuleFragment>> abstract val modulePhase: SameTypeNamedCompilerPhase<JsIrBackendContext, IrModuleFragment>
} }
class DeclarationLowering( class DeclarationLowering(
name: String, name: String,
description: String, description: String,
prerequisite: Set<AbstractNamedCompilerPhase<JsIrBackendContext, *, *>> = emptySet(), prerequisite: Set<AbstractNamedCompilerPhase<JsIrBackendContext, *, *>> = emptySet(),
private val factory: (JsIrBackendContext) -> DeclarationTransformer factory: (JsIrBackendContext) -> DeclarationTransformer
) : Lowering(name) { ) : Lowering(name) {
fun declarationTransformer(context: JsIrBackendContext): DeclarationTransformer {
return factory(context)
}
override val modulePhase = makeJsModulePhase(factory, name, description, prerequisite) override val modulePhase = makeJsModulePhase(factory, name, description, prerequisite)
} }
@@ -97,18 +89,14 @@ class BodyLowering(
name: String, name: String,
description: String, description: String,
prerequisite: Set<AbstractNamedCompilerPhase<JsIrBackendContext, *, *>> = emptySet(), prerequisite: Set<AbstractNamedCompilerPhase<JsIrBackendContext, *, *>> = emptySet(),
private val factory: (JsIrBackendContext) -> BodyLoweringPass factory: (JsIrBackendContext) -> BodyLoweringPass
) : Lowering(name) { ) : Lowering(name) {
fun bodyLowering(context: JsIrBackendContext): BodyLoweringPass {
return factory(context)
}
override val modulePhase = makeJsModulePhase(factory, name, description, prerequisite) override val modulePhase = makeJsModulePhase(factory, name, description, prerequisite)
} }
class ModuleLowering( class ModuleLowering(
name: String, name: String,
override val modulePhase: SameTypeNamedCompilerPhase<JsIrBackendContext, Iterable<IrModuleFragment>> override val modulePhase: SameTypeNamedCompilerPhase<JsIrBackendContext, IrModuleFragment>
) : Lowering(name) ) : Lowering(name)
private fun makeDeclarationTransformerPhase( private fun makeDeclarationTransformerPhase(
@@ -125,7 +113,7 @@ private fun makeBodyLoweringPhase(
prerequisite: Set<Lowering> = emptySet() prerequisite: Set<Lowering> = emptySet()
) = BodyLowering(name, description, prerequisite.map { it.modulePhase }.toSet(), lowering) ) = BodyLowering(name, description, prerequisite.map { it.modulePhase }.toSet(), lowering)
fun SameTypeNamedCompilerPhase<JsIrBackendContext, Iterable<IrModuleFragment>>.toModuleLowering() = ModuleLowering(this.name, this) fun SameTypeNamedCompilerPhase<JsIrBackendContext, IrModuleFragment>.toModuleLowering() = ModuleLowering(this.name, this)
private val validateIrBeforeLowering = makeCustomJsModulePhase( private val validateIrBeforeLowering = makeCustomJsModulePhase(
{ context, module -> validationCallback(context, module) }, { context, module -> validationCallback(context, module) },
@@ -981,7 +969,7 @@ val jsPhases = SameTypeNamedCompilerPhase(
name = "IrModuleLowering", name = "IrModuleLowering",
description = "IR module lowering", description = "IR module lowering",
lower = loweringList.toCompilerPhase(), lower = loweringList.toCompilerPhase(),
actions = setOf(defaultDumper.toMultiModuleAction(), validationAction.toMultiModuleAction()), actions = setOf(defaultDumper, validationAction),
nlevels = 1 nlevels = 1
) )
@@ -1045,6 +1033,6 @@ val jsOptimizationPhases = SameTypeNamedCompilerPhase(
name = "IrModuleOptimizationLowering", name = "IrModuleOptimizationLowering",
description = "IR module optimization lowering", description = "IR module optimization lowering",
lower = optimizationLoweringList.toCompilerPhase(), lower = optimizationLoweringList.toCompilerPhase(),
actions = setOf(defaultDumper.toMultiModuleAction(), validationAction.toMultiModuleAction()), actions = setOf(defaultDumper, validationAction),
nlevels = 1 nlevels = 1
) )
@@ -7,7 +7,7 @@ package org.jetbrains.kotlin.ir.backend.js
import org.jetbrains.kotlin.backend.common.linkage.issues.checkNoUnboundSymbols import org.jetbrains.kotlin.backend.common.linkage.issues.checkNoUnboundSymbols
import org.jetbrains.kotlin.backend.common.phaser.PhaseConfig import org.jetbrains.kotlin.backend.common.phaser.PhaseConfig
import org.jetbrains.kotlin.backend.common.phaser.invokeToplevel import org.jetbrains.kotlin.backend.common.phaser.PhaserState
import org.jetbrains.kotlin.config.CompilerConfiguration import org.jetbrains.kotlin.config.CompilerConfiguration
import org.jetbrains.kotlin.ir.IrBuiltIns import org.jetbrains.kotlin.ir.IrBuiltIns
import org.jetbrains.kotlin.ir.backend.js.lower.collectNativeImplementations import org.jetbrains.kotlin.ir.backend.js.lower.collectNativeImplementations
@@ -137,7 +137,14 @@ fun compileIr(
(irFactory.stageController as? WholeWorldStageController)?.let { (irFactory.stageController as? WholeWorldStageController)?.let {
lowerPreservingTags(allModules, context, phaseConfig, it) lowerPreservingTags(allModules, context, phaseConfig, it)
} ?: jsPhases.invokeToplevel(phaseConfig, context, allModules) } ?: run {
val phaserState = PhaserState<IrModuleFragment>()
loweringList.forEachIndexed { _, lowering ->
allModules.forEach { module ->
lowering.modulePhase.invoke(phaseConfig, phaserState, context, module)
}
}
}
return LoweredIr(context, moduleFragment, allModules, moduleToName) return LoweredIr(context, moduleFragment, allModules, moduleToName)
} }
@@ -76,11 +76,13 @@ fun lowerPreservingTags(
// Lower all the things // Lower all the things
controller.currentStage = 0 controller.currentStage = 0
val phaserState = PhaserState<Iterable<IrModuleFragment>>() val phaserState = PhaserState<IrModuleFragment>()
loweringList.forEachIndexed { i, lowering -> loweringList.forEachIndexed { i, lowering ->
controller.currentStage = i + 1 controller.currentStage = i + 1
lowering.modulePhase.invoke(phaseConfig, phaserState, context, modules) modules.forEach { module ->
lowering.modulePhase.invoke(phaseConfig, phaserState, context, module)
}
} }
controller.currentStage = loweringList.size + 1 controller.currentStage = loweringList.size + 1
@@ -6,6 +6,7 @@
package org.jetbrains.kotlin.ir.backend.js package org.jetbrains.kotlin.ir.backend.js
import org.jetbrains.kotlin.backend.common.phaser.PhaseConfig import org.jetbrains.kotlin.backend.common.phaser.PhaseConfig
import org.jetbrains.kotlin.backend.common.phaser.PhaserState
import org.jetbrains.kotlin.backend.common.phaser.invokeToplevel import org.jetbrains.kotlin.backend.common.phaser.invokeToplevel
import org.jetbrains.kotlin.ir.backend.js.dce.DceDumpNameCache import org.jetbrains.kotlin.ir.backend.js.dce.DceDumpNameCache
import org.jetbrains.kotlin.ir.backend.js.dce.eliminateDeadDeclarations import org.jetbrains.kotlin.ir.backend.js.dce.eliminateDeadDeclarations
@@ -24,7 +25,14 @@ fun optimizeProgramByIr(
) { ) {
val dceDumpNameCache = DceDumpNameCache() // in JS mode only DCE Graph could be dumped val dceDumpNameCache = DceDumpNameCache() // in JS mode only DCE Graph could be dumped
eliminateDeadDeclarations(modules, context, removeUnusedAssociatedObjects, dceDumpNameCache) eliminateDeadDeclarations(modules, context, removeUnusedAssociatedObjects, dceDumpNameCache)
jsOptimizationPhases.invokeToplevel(PhaseConfig(jsOptimizationPhases), context, modules)
val phaseConfig = PhaseConfig(jsOptimizationPhases)
val phaserState = PhaserState<IrModuleFragment>()
optimizationLoweringList.forEachIndexed { _, lowering ->
modules.forEach { module ->
lowering.modulePhase.invoke(phaseConfig, phaserState, context, module)
}
}
} }
fun optimizeFragmentByJsAst(fragment: JsIrProgramFragment) { fun optimizeFragmentByJsAst(fragment: JsIrProgramFragment) {
@@ -611,7 +611,7 @@ class GenerateIrRuntime {
ExternalDependenciesGenerator(symbolTable, listOf(jsLinker)).generateUnboundSymbolsAsDependencies() ExternalDependenciesGenerator(symbolTable, listOf(jsLinker)).generateUnboundSymbolsAsDependencies()
jsPhases.invokeToplevel(phaseConfig, context, listOf(module)) jsPhases.invokeToplevel(phaseConfig, context, module)
val transformer = IrModuleToJsTransformer(context, shouldReferMainFunction = false) val transformer = IrModuleToJsTransformer(context, shouldReferMainFunction = false)