Backend: remove codegen factory from generation state

use it explicitly. This is a step in attempt to abstract dependencies
on PSI in the GenerationState and related places.
This commit is contained in:
Ilya Chernikov
2022-02-15 12:54:20 +03:00
committed by teamcity
parent 018782f0c7
commit da41fddabb
15 changed files with 37 additions and 50 deletions
@@ -25,7 +25,7 @@ import org.jetbrains.kotlin.psi.KtFile;
import java.util.Collection; import java.util.Collection;
public class KotlinCodegenFacade { public class KotlinCodegenFacade {
public static void compileCorrectFiles(@NotNull GenerationState state) { public static void compileCorrectFiles(@NotNull GenerationState state, CodegenFactory codegenFactory) {
ProgressIndicatorAndCompilationCanceledStatus.checkCanceled(); ProgressIndicatorAndCompilationCanceledStatus.checkCanceled();
state.beforeCompile(); state.beforeCompile();
@@ -33,11 +33,11 @@ public class KotlinCodegenFacade {
ProgressIndicatorAndCompilationCanceledStatus.checkCanceled(); ProgressIndicatorAndCompilationCanceledStatus.checkCanceled();
CodegenFactory.IrConversionInput psi2irInput = CodegenFactory.IrConversionInput.Companion.fromGenerationState(state); CodegenFactory.IrConversionInput psi2irInput = CodegenFactory.IrConversionInput.Companion.fromGenerationState(state);
CodegenFactory.BackendInput backendInput = state.getCodegenFactory().convertToIr(psi2irInput); CodegenFactory.BackendInput backendInput = codegenFactory.convertToIr(psi2irInput);
ProgressIndicatorAndCompilationCanceledStatus.checkCanceled(); ProgressIndicatorAndCompilationCanceledStatus.checkCanceled();
state.getCodegenFactory().generateModule(state, backendInput); codegenFactory.generateModule(state, backendInput);
CodegenFactory.Companion.doCheckCancelled(state); CodegenFactory.Companion.doCheckCancelled(state);
state.getFactory().done(); state.getFactory().done();
@@ -62,7 +62,6 @@ class GenerationState private constructor(
val files: List<KtFile>, val files: List<KtFile>,
val configuration: CompilerConfiguration, val configuration: CompilerConfiguration,
val generateDeclaredClassFilter: GenerateClassFilter, val generateDeclaredClassFilter: GenerateClassFilter,
val codegenFactory: CodegenFactory,
val targetId: TargetId?, val targetId: TargetId?,
moduleName: String?, moduleName: String?,
val outDirectory: File?, val outDirectory: File?,
@@ -86,10 +85,6 @@ class GenerationState private constructor(
fun generateDeclaredClassFilter(v: GenerateClassFilter) = fun generateDeclaredClassFilter(v: GenerateClassFilter) =
apply { generateDeclaredClassFilter = v } apply { generateDeclaredClassFilter = v }
private var codegenFactory: CodegenFactory = DefaultCodegenFactory
fun codegenFactory(v: CodegenFactory) =
apply { codegenFactory = v }
private var targetId: TargetId? = null private var targetId: TargetId? = null
fun targetId(v: TargetId?) = fun targetId(v: TargetId?) =
apply { targetId = v } apply { targetId = v }
@@ -135,7 +130,7 @@ class GenerationState private constructor(
fun build() = fun build() =
GenerationState( GenerationState(
project, builderFactory, module, bindingContext, files, configuration, project, builderFactory, module, bindingContext, files, configuration,
generateDeclaredClassFilter, codegenFactory, targetId, generateDeclaredClassFilter, targetId,
moduleName, outDirectory, onIndependentPartCompilationEnd, wantsDiagnostics, moduleName, outDirectory, onIndependentPartCompilationEnd, wantsDiagnostics,
jvmBackendClassResolver, isIrBackend, ignoreErrors, jvmBackendClassResolver, isIrBackend, ignoreErrors,
diagnosticReporter ?: DiagnosticReporterFactory.createReporter(), diagnosticReporter ?: DiagnosticReporterFactory.createReporter(),
@@ -310,8 +310,6 @@ object FirKotlinToJvmBytecodeCompiler {
(projectEnvironment as VfsBasedProjectEnvironment).project, ClassBuilderFactories.BINARIES, (projectEnvironment as VfsBasedProjectEnvironment).project, ClassBuilderFactories.BINARIES,
moduleFragment.descriptor, dummyBindingContext, ktFiles, moduleFragment.descriptor, dummyBindingContext, ktFiles,
moduleConfiguration moduleConfiguration
).codegenFactory(
codegenFactory
).withModule( ).withModule(
module module
).onIndependentPartCompilationEnd( ).onIndependentPartCompilationEnd(
@@ -133,7 +133,7 @@ object KotlinToJVMBytecodeCompiler {
val outputs = ArrayList<GenerationState>(chunk.size) val outputs = ArrayList<GenerationState>(chunk.size)
for (input in codegenInputs) { for (input in codegenInputs) {
outputs += runCodegen(input, input.state, result.bindingContext, diagnosticsReporter, environment.configuration) outputs += runCodegen(input, input.state, codegenFactory, result.bindingContext, diagnosticsReporter, environment.configuration)
} }
return writeOutputs(environment.project, projectConfiguration, chunk, outputs, mainClassFqName) return writeOutputs(environment.project, projectConfiguration, chunk, outputs, mainClassFqName)
@@ -206,7 +206,7 @@ object KotlinToJVMBytecodeCompiler {
environment, environment.configuration, result, environment.getSourceFiles(), null, codegenFactory, backendInput, environment, environment.configuration, result, environment.getSourceFiles(), null, codegenFactory, backendInput,
diagnosticsReporter diagnosticsReporter
) )
return runCodegen(input, input.state, result.bindingContext, diagnosticsReporter, environment.configuration) return runCodegen(input, input.state, codegenFactory, result.bindingContext, diagnosticsReporter, environment.configuration)
} }
private fun convertToIr(environment: KotlinCoreEnvironment, result: AnalysisResult): Pair<CodegenFactory, CodegenFactory.BackendInput> { private fun convertToIr(environment: KotlinCoreEnvironment, result: AnalysisResult): Pair<CodegenFactory, CodegenFactory.BackendInput> {
@@ -321,7 +321,6 @@ object KotlinToJVMBytecodeCompiler {
sourceFiles, sourceFiles,
configuration configuration
) )
.codegenFactory(codegenFactory)
.withModule(module) .withModule(module)
.onIndependentPartCompilationEnd(createOutputFilesFlushingCallbackIfPossible(configuration)) .onIndependentPartCompilationEnd(createOutputFilesFlushingCallbackIfPossible(configuration))
.diagnosticReporter(diagnosticsReporter) .diagnosticReporter(diagnosticsReporter)
@@ -343,6 +342,7 @@ object KotlinToJVMBytecodeCompiler {
private fun runCodegen( private fun runCodegen(
codegenInput: CodegenFactory.CodegenInput, codegenInput: CodegenFactory.CodegenInput,
state: GenerationState, state: GenerationState,
codegenFactory: CodegenFactory,
bindingContext: BindingContext, bindingContext: BindingContext,
diagnosticsReporter: BaseDiagnosticsCollector, diagnosticsReporter: BaseDiagnosticsCollector,
configuration: CompilerConfiguration, configuration: CompilerConfiguration,
@@ -352,7 +352,7 @@ object KotlinToJVMBytecodeCompiler {
val performanceManager = configuration[CLIConfigurationKeys.PERF_MANAGER] val performanceManager = configuration[CLIConfigurationKeys.PERF_MANAGER]
performanceManager?.notifyIRGenerationStarted() performanceManager?.notifyIRGenerationStarted()
state.codegenFactory.invokeCodegen(codegenInput) codegenFactory.invokeCodegen(codegenInput)
CodegenFactory.doCheckCancelled(state) CodegenFactory.doCheckCancelled(state)
state.factory.done() state.factory.done()
@@ -222,8 +222,6 @@ fun generateCodeFromIr(
(environment.projectEnvironment as VfsBasedProjectEnvironment).project, ClassBuilderFactories.BINARIES, (environment.projectEnvironment as VfsBasedProjectEnvironment).project, ClassBuilderFactories.BINARIES,
input.irModuleFragment.descriptor, dummyBindingContext, emptyList()/* !! */, input.irModuleFragment.descriptor, dummyBindingContext, emptyList()/* !! */,
input.configuration input.configuration
).codegenFactory(
codegenFactory
).targetId( ).targetId(
input.targetId input.targetId
).moduleName( ).moduleName(
@@ -39,9 +39,9 @@ class ClassicJvmBackendFacade(
analysisResult.bindingContext, analysisResult.bindingContext,
psiFiles.toList(), psiFiles.toList(),
configuration configuration
).codegenFactory(DefaultCodegenFactory).build() ).build()
KotlinCodegenFacade.compileCorrectFiles(generationState) KotlinCodegenFacade.compileCorrectFiles(generationState, DefaultCodegenFactory)
javaCompilerFacade.compileJavaFiles(module, configuration, generationState.factory) javaCompilerFacade.compileJavaFiles(module, configuration, generationState.factory)
return BinaryArtifacts.Jvm(generationState.factory) return BinaryArtifacts.Jvm(generationState.factory)
} }
@@ -6,6 +6,7 @@
package org.jetbrains.kotlin.test.backend.ir package org.jetbrains.kotlin.test.backend.ir
import org.jetbrains.kotlin.backend.jvm.JvmIrCodegenFactory import org.jetbrains.kotlin.backend.jvm.JvmIrCodegenFactory
import org.jetbrains.kotlin.codegen.CodegenFactory
import org.jetbrains.kotlin.codegen.state.GenerationState import org.jetbrains.kotlin.codegen.state.GenerationState
import org.jetbrains.kotlin.descriptors.DeclarationDescriptor import org.jetbrains.kotlin.descriptors.DeclarationDescriptor
import org.jetbrains.kotlin.ir.backend.js.KotlinFileSerializedData import org.jetbrains.kotlin.ir.backend.js.KotlinFileSerializedData
@@ -33,6 +34,7 @@ sealed class IrBackendInput : ResultingArtifact.BackendInput<IrBackendInput>() {
data class JvmIrBackendInput( data class JvmIrBackendInput(
val state: GenerationState, val state: GenerationState,
val codegenFactory: JvmIrCodegenFactory,
val backendInput: JvmIrCodegenFactory.JvmIrBackendInput val backendInput: JvmIrCodegenFactory.JvmIrBackendInput
) : IrBackendInput() { ) : IrBackendInput() {
override val irModuleFragment: IrModuleFragment override val irModuleFragment: IrModuleFragment
@@ -28,9 +28,8 @@ class JvmIrBackendFacade(
"JvmIrBackendFacade expects IrBackendInput.JvmIrBackendInput as input" "JvmIrBackendFacade expects IrBackendInput.JvmIrBackendInput as input"
} }
val state = inputArtifact.state val state = inputArtifact.state
val codegenFactory = state.codegenFactory as JvmIrCodegenFactory
try { try {
codegenFactory.generateModule(state, inputArtifact.backendInput) inputArtifact.codegenFactory.generateModule(state, inputArtifact.backendInput)
} catch (e: BackendException) { } catch (e: BackendException) {
if (CodegenTestDirectives.IGNORE_ERRORS in module.directives) { if (CodegenTestDirectives.IGNORE_ERRORS in module.directives) {
return null return null
@@ -57,14 +57,14 @@ class ClassicFrontend2IrConverter(
val state = GenerationState.Builder( val state = GenerationState.Builder(
project, ClassBuilderFactories.TEST, analysisResult.moduleDescriptor, analysisResult.bindingContext, project, ClassBuilderFactories.TEST, analysisResult.moduleDescriptor, analysisResult.bindingContext,
files, configuration files, configuration
).codegenFactory(codegenFactory) ).isIrBackend(true)
.isIrBackend(true)
.ignoreErrors(CodegenTestDirectives.IGNORE_ERRORS in module.directives) .ignoreErrors(CodegenTestDirectives.IGNORE_ERRORS in module.directives)
.diagnosticReporter(DiagnosticReporterFactory.createReporter()) .diagnosticReporter(DiagnosticReporterFactory.createReporter())
.build() .build()
return IrBackendInput.JvmIrBackendInput( return IrBackendInput.JvmIrBackendInput(
state, state,
codegenFactory,
codegenFactory.convertToIr(CodegenFactory.IrConversionInput.fromGenerationState(state)) codegenFactory.convertToIr(CodegenFactory.IrConversionInput.fromGenerationState(state))
) )
} }
@@ -64,8 +64,6 @@ class Fir2IrResultsConverter(
project, ClassBuilderFactories.TEST, project, ClassBuilderFactories.TEST,
container.get(), dummyBindingContext, ktFiles, container.get(), dummyBindingContext, ktFiles,
configuration configuration
).codegenFactory(
codegenFactory
).isIrBackend( ).isIrBackend(
true true
).jvmBackendClassResolver( ).jvmBackendClassResolver(
@@ -76,6 +74,7 @@ class Fir2IrResultsConverter(
return IrBackendInput.JvmIrBackendInput( return IrBackendInput.JvmIrBackendInput(
generationState, generationState,
codegenFactory,
JvmIrCodegenFactory.JvmIrBackendInput( JvmIrCodegenFactory.JvmIrBackendInput(
irModuleFragment, irModuleFragment,
symbolTable, symbolTable,
@@ -125,8 +125,6 @@ object GenerationUtils {
val generationState = GenerationState.Builder( val generationState = GenerationState.Builder(
project, classBuilderFactory, moduleFragment.descriptor, dummyBindingContext, files, configuration project, classBuilderFactory, moduleFragment.descriptor, dummyBindingContext, files, configuration
).codegenFactory(
codegenFactory
).isIrBackend( ).isIrBackend(
true true
).jvmBackendClassResolver( ).jvmBackendClassResolver(
@@ -191,13 +189,14 @@ object GenerationUtils {
val generationState = GenerationState.Builder( val generationState = GenerationState.Builder(
project, classBuilderFactory, analysisResult.moduleDescriptor, analysisResult.bindingContext, project, classBuilderFactory, analysisResult.moduleDescriptor, analysisResult.bindingContext,
files, configuration files, configuration
).codegenFactory(
if (isIrBackend)
JvmIrCodegenFactory(configuration, configuration.get(CLIConfigurationKeys.PHASE_CONFIG))
else DefaultCodegenFactory
).isIrBackend(isIrBackend).apply(configureGenerationState).build() ).isIrBackend(isIrBackend).apply(configureGenerationState).build()
if (analysisResult.shouldGenerateCode) { if (analysisResult.shouldGenerateCode) {
KotlinCodegenFacade.compileCorrectFiles(generationState) KotlinCodegenFacade.compileCorrectFiles(
generationState,
if (isIrBackend)
JvmIrCodegenFactory(configuration, configuration.get(CLIConfigurationKeys.PHASE_CONFIG))
else DefaultCodegenFactory
)
} }
return generationState return generationState
} }
@@ -13,10 +13,7 @@ import org.jetbrains.kotlin.cli.common.messages.CompilerMessageSeverity
import org.jetbrains.kotlin.cli.common.messages.MessageRenderer import org.jetbrains.kotlin.cli.common.messages.MessageRenderer
import org.jetbrains.kotlin.cli.common.messages.OutputMessageUtil import org.jetbrains.kotlin.cli.common.messages.OutputMessageUtil
import org.jetbrains.kotlin.cli.common.messages.PrintingMessageCollector import org.jetbrains.kotlin.cli.common.messages.PrintingMessageCollector
import org.jetbrains.kotlin.codegen.ClassBuilder import org.jetbrains.kotlin.codegen.*
import org.jetbrains.kotlin.codegen.ClassBuilderFactory
import org.jetbrains.kotlin.codegen.ClassBuilderMode
import org.jetbrains.kotlin.codegen.KotlinCodegenFacade
import org.jetbrains.kotlin.codegen.state.GenerationState import org.jetbrains.kotlin.codegen.state.GenerationState
import org.jetbrains.kotlin.compilerRunner.OutputItemsCollector import org.jetbrains.kotlin.compilerRunner.OutputItemsCollector
import org.jetbrains.kotlin.config.CommonConfigurationKeys import org.jetbrains.kotlin.config.CommonConfigurationKeys
@@ -76,7 +73,7 @@ class JvmAbiAnalysisHandlerExtension(
files.toList(), files.toList(),
compilerConfiguration compilerConfiguration
).targetId(targetId).build() ).targetId(targetId).build()
KotlinCodegenFacade.compileCorrectFiles(generationState) KotlinCodegenFacade.compileCorrectFiles(generationState, DefaultCodegenFactory)
val outputDir = compilerConfiguration.get(JVMConfigurationKeys.OUTPUT_DIRECTORY)!! val outputDir = compilerConfiguration.get(JVMConfigurationKeys.OUTPUT_DIRECTORY)!!
val outputs = ArrayList<AbiOutput>() val outputs = ArrayList<AbiOutput>()
@@ -277,15 +277,15 @@ abstract class AbstractKapt3Extension(
compilerConfiguration compilerConfiguration
).targetId(targetId) ).targetId(targetId)
.isIrBackend(isIrBackend) .isIrBackend(isIrBackend)
.codegenFactory( .build()
val (classFilesCompilationTime) = measureTimeMillis {
KotlinCodegenFacade.compileCorrectFiles(
generationState,
if (isIrBackend) if (isIrBackend)
JvmIrCodegenFactory(compilerConfiguration, compilerConfiguration.get(CLIConfigurationKeys.PHASE_CONFIG)) JvmIrCodegenFactory(compilerConfiguration, compilerConfiguration.get(CLIConfigurationKeys.PHASE_CONFIG))
else DefaultCodegenFactory else DefaultCodegenFactory
) )
.build()
val (classFilesCompilationTime) = measureTimeMillis {
KotlinCodegenFacade.compileCorrectFiles(generationState)
} }
val compiledClasses = builderFactory.compiledClasses val compiledClasses = builderFactory.compiledClasses
@@ -197,7 +197,6 @@ open class KJvmReplCompilerBase<AnalyzerT : ReplCodeAnalyzerBase>(
sourceFiles, sourceFiles,
compilationState.environment.configuration compilationState.environment.configuration
) )
.codegenFactory(codegenFactory)
.build() .build()
codegenFactory.generateModule( codegenFactory.generateModule(
@@ -246,16 +246,17 @@ private fun generate(
analysisResult.bindingContext, analysisResult.bindingContext,
sourceFiles, sourceFiles,
kotlinCompilerConfiguration kotlinCompilerConfiguration
).codegenFactory(
if (kotlinCompilerConfiguration.getBoolean(JVMConfigurationKeys.IR))
JvmIrCodegenFactory(
kotlinCompilerConfiguration,
kotlinCompilerConfiguration.get(CLIConfigurationKeys.PHASE_CONFIG),
) else DefaultCodegenFactory
).diagnosticReporter( ).diagnosticReporter(
diagnosticsReporter diagnosticsReporter
).build().also { ).build().also {
KotlinCodegenFacade.compileCorrectFiles(it) KotlinCodegenFacade.compileCorrectFiles(
it,
if (kotlinCompilerConfiguration.getBoolean(JVMConfigurationKeys.IR))
JvmIrCodegenFactory(
kotlinCompilerConfiguration,
kotlinCompilerConfiguration.get(CLIConfigurationKeys.PHASE_CONFIG),
) else DefaultCodegenFactory
)
FirDiagnosticsCompilerResultsReporter.reportToMessageCollector( FirDiagnosticsCompilerResultsReporter.reportToMessageCollector(
diagnosticsReporter, diagnosticsReporter,
messageCollector, messageCollector,