[LL API] Correct used IR symbols in code fragments

As 'FirCodeFragments' are converted to IR independently of its context,
in some cases duplicate (and not quite correct) symbols for local
classes and functions are created.

Until properly fixed in fir2ir, here we replace such duplicates with
original symbols.
This commit is contained in:
Yan Zhulanow
2023-08-08 02:51:53 +09:00
committed by Space Team
parent ee7e6b0fce
commit 5a67b0d7dc
39 changed files with 611 additions and 9 deletions
@@ -41,23 +41,28 @@ import org.jetbrains.kotlin.diagnostics.KtPsiDiagnostic
import org.jetbrains.kotlin.diagnostics.Severity
import org.jetbrains.kotlin.fir.backend.*
import org.jetbrains.kotlin.fir.backend.jvm.*
import org.jetbrains.kotlin.fir.declarations.FirDeclaration
import org.jetbrains.kotlin.fir.declarations.FirFile
import org.jetbrains.kotlin.fir.declarations.FirResolvePhase
import org.jetbrains.kotlin.fir.pipeline.applyIrGenerationExtensions
import org.jetbrains.kotlin.fir.pipeline.signatureComposerForJvmFir2Ir
import org.jetbrains.kotlin.fir.psi
import org.jetbrains.kotlin.fir.references.FirReference
import org.jetbrains.kotlin.fir.references.FirThisReference
import org.jetbrains.kotlin.fir.references.toResolvedSymbol
import org.jetbrains.kotlin.fir.symbols.FirBasedSymbol
import org.jetbrains.kotlin.ir.IrElement
import org.jetbrains.kotlin.ir.ObsoleteDescriptorBasedAPI
import org.jetbrains.kotlin.ir.PsiIrFileEntry
import org.jetbrains.kotlin.ir.backend.jvm.serialization.JvmIrMangler
import org.jetbrains.kotlin.ir.declarations.IrClass
import org.jetbrains.kotlin.ir.declarations.IrConstructor
import org.jetbrains.kotlin.ir.declarations.IrFunction
import org.jetbrains.kotlin.ir.declarations.IrModuleFragment
import org.jetbrains.kotlin.ir.declarations.*
import org.jetbrains.kotlin.ir.declarations.impl.IrFactoryImpl
import org.jetbrains.kotlin.ir.expressions.*
import org.jetbrains.kotlin.ir.symbols.IrSymbol
import org.jetbrains.kotlin.ir.util.StubGeneratorExtensions
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.load.kotlin.toSourceElement
import org.jetbrains.kotlin.name.ClassId
import org.jetbrains.kotlin.name.FqName
@@ -67,6 +72,7 @@ import org.jetbrains.kotlin.psi2ir.generators.fragments.EvaluatorFragmentInfo
import org.jetbrains.kotlin.resolve.source.PsiSourceFile
import org.jetbrains.kotlin.serialization.deserialization.descriptors.DeserializedContainerSource
import org.jetbrains.kotlin.utils.addToStdlib.runIf
import java.util.Collections
internal class KtFirCompilerFacility(
override val analysisSession: KtFirAnalysisSession
@@ -147,6 +153,8 @@ internal class KtFirCompilerFacility(
initializedIrBuiltIns = null
)
patchCodeFragmentIr(fir2IrResult)
ProgressManager.checkCanceled()
val irGeneratorExtensions = IrGenerationExtension.getInstances(project)
@@ -200,6 +208,22 @@ internal class KtFirCompilerFacility(
}
}
private fun patchCodeFragmentIr(fir2IrResult: Fir2IrResult) {
fun isCodeFragmentFile(irFile: IrFile): Boolean {
val firFiles = (irFile.metadata as? FirMetadataSource.File)?.files ?: return false
return firFiles.any { it.psi is KtCodeFragment }
}
val (irCodeFragmentFiles, irOrdinaryFiles) = fir2IrResult.irModuleFragment.files.partition(::isCodeFragmentFile)
// Collect original declarations from the context files
val collectingVisitor = IrDeclarationMappingCollectingVisitor()
irOrdinaryFiles.forEach { it.acceptVoid(collectingVisitor) }
// Replace duplicate symbols with the original ones
val patchingVisitor = IrDeclarationPatchingVisitor(collectingVisitor.mappings)
irCodeFragmentFiles.forEach { it.acceptVoid(patchingVisitor) }
}
private fun getFullyResolvedFirFile(file: KtFile): FirFile {
val firFile = file.getOrBuildFirFile(firResolveSession)
LLFirWholeFileResolveTarget(firFile).resolve(FirResolvePhase.BODY_RESOLVE)
@@ -368,4 +392,96 @@ internal class KtFirCompilerFacility(
ideCodegenSettings = ideCodegenSettings,
)
}
}
private class IrDeclarationMappingCollectingVisitor : IrElementVisitorVoid {
private val collectedMappings = HashMap<FirDeclaration, IrDeclaration>()
val mappings: Map<FirDeclaration, IrDeclaration>
get() = Collections.unmodifiableMap(collectedMappings)
override fun visitElement(element: IrElement) {
element.acceptChildrenVoid(this)
}
override fun visitDeclaration(declaration: IrDeclarationBase) {
dumpDeclaration(declaration)
super.visitDeclaration(declaration)
}
private fun dumpDeclaration(declaration: IrDeclaration) {
if (declaration is IrMetadataSourceOwner) {
val fir = (declaration.metadata as? FirMetadataSource)?.fir
if (fir != null) {
collectedMappings.putIfAbsent(fir, declaration)
}
}
}
}
private class IrDeclarationPatchingVisitor(private val mapping: Map<FirDeclaration, IrDeclaration>) : IrElementVisitorVoid {
override fun visitElement(element: IrElement) {
element.acceptChildrenVoid(this)
}
override fun visitFieldAccess(expression: IrFieldAccessExpression) {
patchIfNeeded(expression.symbol) { expression.symbol = it }
patchIfNeeded(expression.superQualifierSymbol) { expression.superQualifierSymbol = it }
super.visitFieldAccess(expression)
}
override fun visitValueAccess(expression: IrValueAccessExpression) {
patchIfNeeded(expression.symbol) { expression.symbol = it }
super.visitValueAccess(expression)
}
override fun visitGetEnumValue(expression: IrGetEnumValue) {
patchIfNeeded(expression.symbol) { expression.symbol = it }
super.visitGetEnumValue(expression)
}
override fun visitGetObjectValue(expression: IrGetObjectValue) {
patchIfNeeded(expression.symbol) { expression.symbol = it }
super.visitGetObjectValue(expression)
}
override fun visitCall(expression: IrCall) {
patchIfNeeded(expression.symbol) { expression.symbol = it }
patchIfNeeded(expression.superQualifierSymbol) { expression.superQualifierSymbol = it }
super.visitCall(expression)
}
override fun visitConstructorCall(expression: IrConstructorCall) {
patchIfNeeded(expression.symbol) { expression.symbol = it }
super.visitConstructorCall(expression)
}
override fun visitPropertyReference(expression: IrPropertyReference) {
patchIfNeeded(expression.symbol) { expression.symbol = it }
patchIfNeeded(expression.getter) { expression.getter = it }
patchIfNeeded(expression.setter) { expression.setter = it }
super.visitPropertyReference(expression)
}
override fun visitFunctionReference(expression: IrFunctionReference) {
patchIfNeeded(expression.symbol) { expression.symbol = it }
patchIfNeeded(expression.reflectionTarget) { expression.reflectionTarget = it }
super.visitFunctionReference(expression)
}
override fun visitClassReference(expression: IrClassReference) {
patchIfNeeded(expression.symbol) { expression.symbol = it }
super.visitClassReference(expression)
}
private inline fun <reified T : IrSymbol> patchIfNeeded(irSymbol: T?, patcher: (T) -> Unit) {
if (irSymbol != null) {
val irDeclaration = irSymbol.owner as? IrMetadataSourceOwner ?: return
val firDeclaration = (irDeclaration.metadata as? FirMetadataSource)?.fir ?: return
val correctedIrSymbol = mapping[firDeclaration]?.symbol as? T ?: return
if (correctedIrSymbol != irSymbol) {
patcher(correctedIrSymbol)
}
}
}
}
@@ -172,6 +172,48 @@ public class FirIdeNormalAnalysisSourceModuleCompilerFacilityTestGenerated exten
runTest("analysis/analysis-api/testData/components/compilerFacility/compilation/codeFragments/capturing/localFunction.kt");
}
@Test
@TestMetadata("localFunctionContainingClassClosure.kt")
public void testLocalFunctionContainingClassClosure() throws Exception {
runTest("analysis/analysis-api/testData/components/compilerFacility/compilation/codeFragments/capturing/localFunctionContainingClassClosure.kt");
}
@Test
@TestMetadata("localFunctionExtensionReceiverClosure.kt")
public void testLocalFunctionExtensionReceiverClosure() throws Exception {
runTest("analysis/analysis-api/testData/components/compilerFacility/compilation/codeFragments/capturing/localFunctionExtensionReceiverClosure.kt");
}
@Test
@TestMetadata("localFunctionLambdaParameterClosure.kt")
public void testLocalFunctionLambdaParameterClosure() throws Exception {
runTest("analysis/analysis-api/testData/components/compilerFacility/compilation/codeFragments/capturing/localFunctionLambdaParameterClosure.kt");
}
@Test
@TestMetadata("localFunctionLocalClosure.kt")
public void testLocalFunctionLocalClosure() throws Exception {
runTest("analysis/analysis-api/testData/components/compilerFacility/compilation/codeFragments/capturing/localFunctionLocalClosure.kt");
}
@Test
@TestMetadata("localFunctionLocalClosureMutating.kt")
public void testLocalFunctionLocalClosureMutating() throws Exception {
runTest("analysis/analysis-api/testData/components/compilerFacility/compilation/codeFragments/capturing/localFunctionLocalClosureMutating.kt");
}
@Test
@TestMetadata("localFunctionMultipleCapturing.kt")
public void testLocalFunctionMultipleCapturing() throws Exception {
runTest("analysis/analysis-api/testData/components/compilerFacility/compilation/codeFragments/capturing/localFunctionMultipleCapturing.kt");
}
@Test
@TestMetadata("localFunctionParameterClosure.kt")
public void testLocalFunctionParameterClosure() throws Exception {
runTest("analysis/analysis-api/testData/components/compilerFacility/compilation/codeFragments/capturing/localFunctionParameterClosure.kt");
}
@Test
@TestMetadata("localMutated.kt")
public void testLocalMutated() throws Exception {