[KT-58817] Implicit receiver for code completion

Prior to this commit code completion for .kts files was supported only
partially. Script has an implicit receiver - base class representing a
script itself, the point where its basic API resides. The knowledge of
this receiver was missing.

There are at least two context where this knowledge is crucial:
1. Code highlighting (worked fine)
2. Code completion (failed)

`FirScriptConfiguratorExtension` is responsible for filling
`FirScriptBuilder` with base script class (in addition to other
properties). See usages of [1] in its implementation.

The thing is that resolution during the completion works a bit
different. Instead of converting the entire `KtFile` it's interested
in `KtScript` only. See usages of [2].
`RawFirBuilder.Visitor.visitScript` is where implicit receivers were
missing.

Code completion and inspections applied to `.kts` files already have
`FirScript` and don't require its full reconstruction with expensive
`FirScriptConfiguratorExtension`. `RawFirBuilder.Visitor` was modified
to support sometimes already existing `FirElement`.

----------------------------------------------------------------
[1]: ScriptCompilationConfiguration.baseClass
[2]: RawFirBuilder.Visitor.convertScript
This commit is contained in:
Andrei Klunnyi
2023-05-17 15:21:28 +02:00
committed by Space Team
parent 2a1d4a42ae
commit 6535278bd3
5 changed files with 133 additions and 105 deletions
@@ -24,8 +24,8 @@ internal fun buildFileFirAnnotation(
val builder = object : RawFirBuilder(session, baseScopeProvider) { val builder = object : RawFirBuilder(session, baseScopeProvider) {
inner class VisitorWithReplacement : Visitor() { inner class VisitorWithReplacement : Visitor() {
override fun convertElement(element: KtElement): FirElement? = override fun convertElement(element: KtElement, original: FirElement?): FirElement? =
super.convertElement(replacementApplier?.tryReplace(element) ?: element) super.convertElement(replacementApplier?.tryReplace(element) ?: element, original)
} }
} }
builder.context.packageFqName = fileAnnotation.containingKtFile.packageFqName builder.context.packageFqName = fileAnnotation.containingKtFile.packageFqName
@@ -177,8 +177,8 @@ internal class RawFirNonLocalDeclarationBuilder private constructor(
} }
private inner class VisitorWithReplacement(private val containingClass: FirRegularClass?) : Visitor() { private inner class VisitorWithReplacement(private val containingClass: FirRegularClass?) : Visitor() {
override fun convertElement(element: KtElement): FirElement? = override fun convertElement(element: KtElement, original: FirElement?): FirElement? =
super.convertElement(replacementApplier?.tryReplace(element) ?: element) super.convertElement(replacementApplier?.tryReplace(element) ?: element, original)
override fun convertProperty( override fun convertProperty(
property: KtProperty, property: KtProperty,
@@ -226,7 +226,7 @@ internal class RawFirNonLocalDeclarationBuilder private constructor(
return ConstructorConversionParams(superTypeCallEntry, selfType, typeParameters) return ConstructorConversionParams(superTypeCallEntry, selfType, typeParameters)
} }
override fun visitSecondaryConstructor(constructor: KtSecondaryConstructor, data: Unit?): FirElement { override fun visitSecondaryConstructor(constructor: KtSecondaryConstructor, data: FirElement?): FirElement {
val classOrObject = constructor.getContainingClassOrObject() val classOrObject = constructor.getContainingClassOrObject()
val params = extractContructorConversionParams(classOrObject, constructor) val params = extractContructorConversionParams(classOrObject, constructor)
val delegatedTypeRef = (originalDeclaration as FirConstructor).delegatedConstructor?.constructedTypeRef val delegatedTypeRef = (originalDeclaration as FirConstructor).delegatedConstructor?.constructedTypeRef
@@ -260,10 +260,10 @@ internal class RawFirNonLocalDeclarationBuilder private constructor(
return newConstructor return newConstructor
} }
override fun visitPrimaryConstructor(constructor: KtPrimaryConstructor, data: Unit?): FirElement = override fun visitPrimaryConstructor(constructor: KtPrimaryConstructor, data: FirElement?): FirElement =
processPrimaryConstructor(constructor.getContainingClassOrObject(), constructor) processPrimaryConstructor(constructor.getContainingClassOrObject(), constructor)
override fun visitEnumEntry(enumEntry: KtEnumEntry, data: Unit?): FirElement { override fun visitEnumEntry(enumEntry: KtEnumEntry, data: FirElement?): FirElement {
val owner = containingClass ?: buildErrorWithAttachment("Enum entry outside of class") { val owner = containingClass ?: buildErrorWithAttachment("Enum entry outside of class") {
withPsiEntry("enumEntry", enumEntry, baseSession.llFirModuleData.ktModule) withPsiEntry("enumEntry", enumEntry, baseSession.llFirModuleData.ktModule)
} }
@@ -310,17 +310,17 @@ internal class RawFirNonLocalDeclarationBuilder private constructor(
// Constructor outside of class, syntax error, we should not do anything // Constructor outside of class, syntax error, we should not do anything
originalDeclaration originalDeclaration
} else { } else {
visitor.convertElement(declarationToBuild) visitor.convertElement(declarationToBuild, originalDeclaration)
} }
} }
is KtClassOrObject -> { is KtClassOrObject -> {
when { when {
originalDeclaration is FirConstructor -> visitor.processPrimaryConstructor(declarationToBuild, null) originalDeclaration is FirConstructor -> visitor.processPrimaryConstructor(declarationToBuild, null)
originalDeclaration is FirField -> visitor.processField(declarationToBuild, originalDeclaration) originalDeclaration is FirField -> visitor.processField(declarationToBuild, originalDeclaration)
else -> visitor.convertElement(declarationToBuild) else -> visitor.convertElement(declarationToBuild, originalDeclaration)
} }
} }
else -> visitor.convertElement(declarationToBuild) else -> visitor.convertElement(declarationToBuild, originalDeclaration)
} as FirDeclaration } as FirDeclaration
} }
@@ -5,7 +5,7 @@
package org.jetbrains.kotlin.analysis.low.level.api.fir.lazy.resolve package org.jetbrains.kotlin.analysis.low.level.api.fir.lazy.resolve
import org.jetbrains.kotlin.fir.* import org.jetbrains.kotlin.fir.FirSession
import org.jetbrains.kotlin.fir.builder.RawFirBuilder import org.jetbrains.kotlin.fir.builder.RawFirBuilder
import org.jetbrains.kotlin.fir.scopes.FirScopeProvider import org.jetbrains.kotlin.fir.scopes.FirScopeProvider
import org.jetbrains.kotlin.fir.types.FirUserTypeRef import org.jetbrains.kotlin.fir.types.FirUserTypeRef
@@ -17,7 +17,7 @@ internal fun buildFirUserTypeRef(
baseScopeProvider: FirScopeProvider baseScopeProvider: FirScopeProvider
): FirUserTypeRef { ): FirUserTypeRef {
val builder = object : RawFirBuilder(session, baseScopeProvider) { val builder = object : RawFirBuilder(session, baseScopeProvider) {
fun build(): FirUserTypeRef = Visitor().visitTypeReference(typeReference, Unit) as FirUserTypeRef fun build(): FirUserTypeRef = Visitor().visitTypeReference(typeReference, null) as FirUserTypeRef
} }
builder.context.packageFqName = typeReference.containingKtFile.packageFqName builder.context.packageFqName = typeReference.containingKtFile.packageFqName
return builder.build() return builder.build()
@@ -75,15 +75,15 @@ open class RawFirBuilder(
} }
fun buildFirFile(file: KtFile): FirFile { fun buildFirFile(file: KtFile): FirFile {
return runOnStubs { file.accept(Visitor(), Unit) as FirFile } return runOnStubs { file.accept(Visitor(), null) as FirFile }
} }
fun buildAnnotationCall(annotation: KtAnnotationEntry): FirAnnotationCall { fun buildAnnotationCall(annotation: KtAnnotationEntry): FirAnnotationCall {
return Visitor().visitAnnotationEntry(annotation, Unit) as FirAnnotationCall return Visitor().visitAnnotationEntry(annotation, null) as FirAnnotationCall
} }
fun buildTypeReference(reference: KtTypeReference): FirTypeRef { fun buildTypeReference(reference: KtTypeReference): FirTypeRef {
return reference.accept(Visitor(), Unit) as FirTypeRef return reference.accept(Visitor(), null) as FirTypeRef
} }
override fun PsiElement.toFirSourceElement(kind: KtFakeSourceElementKind?): KtPsiSourceElement { override fun PsiElement.toFirSourceElement(kind: KtFakeSourceElementKind?): KtPsiSourceElement {
@@ -162,7 +162,7 @@ open class RawFirBuilder(
} }
} }
protected open inner class Visitor : KtVisitor<FirElement, Unit>() { protected open inner class Visitor : KtVisitor<FirElement, FirElement?>() {
private inline fun <reified R : FirElement> KtElement?.convertSafe(): R? = private inline fun <reified R : FirElement> KtElement?.convertSafe(): R? =
this?.let(::convertElement) as? R this?.let(::convertElement) as? R
@@ -214,8 +214,8 @@ open class RawFirBuilder(
}) })
} }
open fun convertElement(element: KtElement): FirElement? = open fun convertElement(element: KtElement, original: FirElement? = null): FirElement? =
element.accept(this@Visitor, Unit) element.accept(this@Visitor, original)
open fun convertProperty( open fun convertProperty(
property: KtProperty, ownerRegularOrAnonymousObjectSymbol: FirClassSymbol<*>?, property: KtProperty, ownerRegularOrAnonymousObjectSymbol: FirClassSymbol<*>?,
@@ -349,7 +349,7 @@ open class RawFirBuilder(
private fun KtExpression?.toFirBlock(): FirBlock = private fun KtExpression?.toFirBlock(): FirBlock =
when (this) { when (this) {
is KtBlockExpression -> is KtBlockExpression ->
accept(this@Visitor, Unit) as FirBlock accept(this@Visitor, null) as FirBlock
null -> null ->
buildEmptyExpressionBlock() buildEmptyExpressionBlock()
else -> { else -> {
@@ -369,7 +369,7 @@ open class RawFirBuilder(
if (hasBody()) { if (hasBody()) {
buildOrLazyBlock { buildOrLazyBlock {
if (hasBlockBody()) { if (hasBlockBody()) {
val block = bodyBlockExpression?.accept(this@Visitor, Unit) as? FirBlock val block = bodyBlockExpression?.accept(this@Visitor, null) as? FirBlock
val contractDescription = when { val contractDescription = when {
!hasContractEffectList() -> block?.let(::processLegacyContractDescription) !hasContractEffectList() -> block?.let(::processLegacyContractDescription)
else -> null else -> null
@@ -394,7 +394,7 @@ open class RawFirBuilder(
val name = this.getArgumentName()?.asName val name = this.getArgumentName()?.asName
val firExpression = when (val expression = this.getArgumentExpression()) { val firExpression = when (val expression = this.getArgumentExpression()) {
is KtConstantExpression, is KtStringTemplateExpression -> { is KtConstantExpression, is KtStringTemplateExpression -> {
expression.accept(this@Visitor, Unit) as FirExpression expression.accept(this@Visitor, null) as FirExpression
} }
else -> { else -> {
@@ -792,7 +792,7 @@ open class RawFirBuilder(
} }
} }
override fun visitTypeParameter(parameter: KtTypeParameter, data: Unit?): FirElement { override fun visitTypeParameter(parameter: KtTypeParameter, data: FirElement?): FirElement {
throw AssertionError("KtTypeParameter should be process via extractTypeParameter") throw AssertionError("KtTypeParameter should be process via extractTypeParameter")
} }
@@ -1028,7 +1028,7 @@ open class RawFirBuilder(
private fun KtClassOrObject.obtainDispatchReceiverForConstructor(): ConeClassLikeType? = private fun KtClassOrObject.obtainDispatchReceiverForConstructor(): ConeClassLikeType? =
if (hasModifier(INNER_KEYWORD)) dispatchReceiverForInnerClassConstructor() else null if (hasModifier(INNER_KEYWORD)) dispatchReceiverForInnerClassConstructor() else null
override fun visitKtFile(file: KtFile, data: Unit): FirElement { override fun visitKtFile(file: KtFile, data: FirElement?): FirElement {
context.packageFqName = when (mode) { context.packageFqName = when (mode) {
BodyBuildingMode.NORMAL -> file.packageFqNameByTree BodyBuildingMode.NORMAL -> file.packageFqNameByTree
BodyBuildingMode.LAZY_BODIES -> file.packageFqName BodyBuildingMode.LAZY_BODIES -> file.packageFqName
@@ -1109,9 +1109,16 @@ open class RawFirBuilder(
} }
} }
override fun visitScript(script: KtScript, data: Unit?): FirElement { // TODO: if no original FirScript is passed here, the invalid script could be constructed, consider throwing an error
val fileName = script.containingKtFile.name override fun visitScript(script: KtScript, data: FirElement?): FirElement {
return convertScript(script, fileName) val ktFile = script.containingKtFile
val fileName = ktFile.name
return convertScript(script, fileName) {
(data as? FirScript)?.let {
contextReceivers.addAll(it.contextReceivers)
parameters.addAll(it.parameters)
}
}
} }
protected fun KtEnumEntry.toFirEnumEntry( protected fun KtEnumEntry.toFirEnumEntry(
@@ -1211,7 +1218,7 @@ open class RawFirBuilder(
} }
} }
override fun visitClassOrObject(classOrObject: KtClassOrObject, data: Unit): FirElement { override fun visitClassOrObject(classOrObject: KtClassOrObject, data: FirElement?): FirElement {
// NB: enum entry nested classes are considered local by FIR design (see discussion in KT-45115) // NB: enum entry nested classes are considered local by FIR design (see discussion in KT-45115)
val isLocal = classOrObject.isLocal || classOrObject.getStrictParentOfType<KtEnumEntry>() != null val isLocal = classOrObject.isLocal || classOrObject.getStrictParentOfType<KtEnumEntry>() != null
val classIsExpect = classOrObject.hasExpectModifier() || context.containerIsExpect val classIsExpect = classOrObject.hasExpectModifier() || context.containerIsExpect
@@ -1363,7 +1370,7 @@ open class RawFirBuilder(
} }
} }
override fun visitObjectLiteralExpression(expression: KtObjectLiteralExpression, data: Unit): FirElement { override fun visitObjectLiteralExpression(expression: KtObjectLiteralExpression, data: FirElement?): FirElement {
return withChildClassName(SpecialNames.ANONYMOUS, forceLocalContext = true, isExpect = false) { return withChildClassName(SpecialNames.ANONYMOUS, forceLocalContext = true, isExpect = false) {
var delegatedFieldsMap: Map<Int, FirFieldSymbol>? var delegatedFieldsMap: Map<Int, FirFieldSymbol>?
buildAnonymousObjectExpression { buildAnonymousObjectExpression {
@@ -1413,7 +1420,7 @@ open class RawFirBuilder(
} }
} }
override fun visitTypeAlias(typeAlias: KtTypeAlias, data: Unit): FirElement { override fun visitTypeAlias(typeAlias: KtTypeAlias, data: FirElement?): FirElement {
val typeAliasIsExpect = typeAlias.hasExpectModifier() || context.containerIsExpect val typeAliasIsExpect = typeAlias.hasExpectModifier() || context.containerIsExpect
return withChildClassName(typeAlias.nameAsSafeName, isExpect = typeAliasIsExpect) { return withChildClassName(typeAlias.nameAsSafeName, isExpect = typeAliasIsExpect) {
buildTypeAlias { buildTypeAlias {
@@ -1433,7 +1440,7 @@ open class RawFirBuilder(
} }
} }
override fun visitNamedFunction(function: KtNamedFunction, data: Unit): FirElement { override fun visitNamedFunction(function: KtNamedFunction, data: FirElement?): FirElement {
val typeReference = function.typeReference val typeReference = function.typeReference
val returnType = if (function.hasBlockBody()) { val returnType = if (function.hasBlockBody()) {
typeReference.toFirOrUnitType() typeReference.toFirOrUnitType()
@@ -1553,12 +1560,12 @@ open class RawFirBuilder(
private fun KtContractEffectList.extractRawEffects(destination: MutableList<FirExpression>) { private fun KtContractEffectList.extractRawEffects(destination: MutableList<FirExpression>) {
getContractEffects().mapTo(destination) { effect -> getContractEffects().mapTo(destination) { effect ->
buildOrLazyExpression(effect.toFirSourceElement()) { buildOrLazyExpression(effect.toFirSourceElement()) {
effect.getExpression().accept(this@Visitor, Unit) as FirExpression effect.getExpression().accept(this@Visitor, null) as FirExpression
} }
} }
} }
override fun visitLambdaExpression(expression: KtLambdaExpression, data: Unit): FirElement { override fun visitLambdaExpression(expression: KtLambdaExpression, data: FirElement?): FirElement {
val literal = expression.functionLiteral val literal = expression.functionLiteral
val literalSource = literal.toFirSourceElement() val literalSource = literal.toFirSourceElement()
val implicitTypeRefSource = literal.toFirSourceElement(KtFakeSourceElementKind.ImplicitTypeRef) val implicitTypeRefSource = literal.toFirSourceElement(KtFakeSourceElementKind.ImplicitTypeRef)
@@ -1889,7 +1896,7 @@ open class RawFirBuilder(
} }
} }
override fun visitAnonymousInitializer(initializer: KtAnonymousInitializer, data: Unit): FirElement { override fun visitAnonymousInitializer(initializer: KtAnonymousInitializer, data: FirElement?): FirElement {
return buildAnonymousInitializer { return buildAnonymousInitializer {
source = initializer.toFirSourceElement() source = initializer.toFirSourceElement()
moduleData = baseModuleData moduleData = baseModuleData
@@ -1899,14 +1906,14 @@ open class RawFirBuilder(
} }
} }
override fun visitProperty(property: KtProperty, data: Unit): FirElement { override fun visitProperty(property: KtProperty, data: FirElement?): FirElement {
return property.toFirProperty( return property.toFirProperty(
ownerRegularOrAnonymousObjectSymbol = null, ownerRegularOrAnonymousObjectSymbol = null,
context = context context = context
) )
} }
override fun visitTypeReference(typeReference: KtTypeReference, data: Unit): FirElement { override fun visitTypeReference(typeReference: KtTypeReference, data: FirElement?): FirElement {
val typeElement = typeReference.typeElement val typeElement = typeReference.typeElement
val source = typeReference.toFirSourceElement() val source = typeReference.toFirSourceElement()
val isNullable = typeElement is KtNullableType val isNullable = typeElement is KtNullableType
@@ -2023,7 +2030,7 @@ open class RawFirBuilder(
return firTypeBuilder.build() return firTypeBuilder.build()
} }
override fun visitAnnotationEntry(annotationEntry: KtAnnotationEntry, data: Unit): FirElement { override fun visitAnnotationEntry(annotationEntry: KtAnnotationEntry, data: FirElement?): FirElement {
return buildAnnotationCall { return buildAnnotationCall {
source = annotationEntry.toFirSourceElement() source = annotationEntry.toFirSourceElement()
useSiteTarget = annotationEntry.useSiteTarget?.getAnnotationUseSiteTarget() useSiteTarget = annotationEntry.useSiteTarget?.getAnnotationUseSiteTarget()
@@ -2038,7 +2045,7 @@ open class RawFirBuilder(
} }
} }
override fun visitTypeProjection(typeProjection: KtTypeProjection, data: Unit): FirElement { override fun visitTypeProjection(typeProjection: KtTypeProjection, data: FirElement?): FirElement {
val projectionKind = typeProjection.projectionKind val projectionKind = typeProjection.projectionKind
val projectionSource = typeProjection.toFirSourceElement() val projectionSource = typeProjection.toFirSourceElement()
if (projectionKind == KtProjectionKind.STAR) { if (projectionKind == KtProjectionKind.STAR) {
@@ -2066,7 +2073,7 @@ open class RawFirBuilder(
} }
} }
override fun visitBlockExpression(expression: KtBlockExpression, data: Unit): FirElement { override fun visitBlockExpression(expression: KtBlockExpression, data: FirElement?): FirElement {
return configureBlockWithoutBuilding(expression).build() return configureBlockWithoutBuilding(expression).build()
} }
@@ -2086,7 +2093,7 @@ open class RawFirBuilder(
} }
} }
override fun visitSimpleNameExpression(expression: KtSimpleNameExpression, data: Unit): FirElement { override fun visitSimpleNameExpression(expression: KtSimpleNameExpression, data: FirElement?): FirElement {
val qualifiedSource = when { val qualifiedSource = when {
expression.getQualifiedExpressionForSelector() != null -> expression.parent expression.getQualifiedExpressionForSelector() != null -> expression.parent
else -> expression else -> expression
@@ -2107,10 +2114,10 @@ open class RawFirBuilder(
) )
} }
override fun visitConstantExpression(expression: KtConstantExpression, data: Unit): FirElement = override fun visitConstantExpression(expression: KtConstantExpression, data: FirElement?): FirElement =
generateConstantExpressionByLiteral(expression) generateConstantExpressionByLiteral(expression)
override fun visitStringTemplateExpression(expression: KtStringTemplateExpression, data: Unit): FirElement { override fun visitStringTemplateExpression(expression: KtStringTemplateExpression, data: FirElement?): FirElement {
return expression.entries.toInterpolatingCall( return expression.entries.toInterpolatingCall(
expression, expression,
getElementType = { element -> getElementType = { element ->
@@ -2128,14 +2135,14 @@ open class RawFirBuilder(
) )
} }
override fun visitReturnExpression(expression: KtReturnExpression, data: Unit): FirElement { override fun visitReturnExpression(expression: KtReturnExpression, data: FirElement?): FirElement {
val source = expression.toFirSourceElement(KtFakeSourceElementKind.ImplicitUnit) val source = expression.toFirSourceElement(KtFakeSourceElementKind.ImplicitUnit)
val result = expression.returnedExpression?.toFirExpression("Incorrect return expression") val result = expression.returnedExpression?.toFirExpression("Incorrect return expression")
?: buildUnitExpression { this.source = source } ?: buildUnitExpression { this.source = source }
return result.toReturn(source, expression.getTargetLabel()?.getReferencedName(), fromKtReturnExpression = true) return result.toReturn(source, expression.getTargetLabel()?.getReferencedName(), fromKtReturnExpression = true)
} }
override fun visitTryExpression(expression: KtTryExpression, data: Unit): FirElement { override fun visitTryExpression(expression: KtTryExpression, data: FirElement?): FirElement {
return buildTryExpression { return buildTryExpression {
source = expression.toFirSourceElement() source = expression.toFirSourceElement()
tryBlock = expression.tryBlock.toFirBlock() tryBlock = expression.tryBlock.toFirBlock()
@@ -2176,7 +2183,7 @@ open class RawFirBuilder(
} }
} }
override fun visitIfExpression(expression: KtIfExpression, data: Unit): FirElement { override fun visitIfExpression(expression: KtIfExpression, data: FirElement?): FirElement {
return buildWhenExpression { return buildWhenExpression {
source = expression.toFirSourceElement() source = expression.toFirSourceElement()
@@ -2207,7 +2214,7 @@ open class RawFirBuilder(
} }
} }
override fun visitWhenExpression(expression: KtWhenExpression, data: Unit): FirElement { override fun visitWhenExpression(expression: KtWhenExpression, data: FirElement?): FirElement {
val ktSubjectExpression = expression.subjectExpression val ktSubjectExpression = expression.subjectExpression
val subjectExpression = when (ktSubjectExpression) { val subjectExpression = when (ktSubjectExpression) {
is KtVariableDeclaration -> ktSubjectExpression.initializer is KtVariableDeclaration -> ktSubjectExpression.initializer
@@ -2305,7 +2312,7 @@ open class RawFirBuilder(
return !(type == KtNodeTypes.FOR || type == KtNodeTypes.WHILE || type == KtNodeTypes.DO_WHILE) return !(type == KtNodeTypes.FOR || type == KtNodeTypes.WHILE || type == KtNodeTypes.DO_WHILE)
} }
override fun visitDoWhileExpression(expression: KtDoWhileExpression, data: Unit): FirElement { override fun visitDoWhileExpression(expression: KtDoWhileExpression, data: FirElement?): FirElement {
val target: FirLoopTarget val target: FirLoopTarget
return FirDoWhileLoopBuilder().apply { return FirDoWhileLoopBuilder().apply {
source = expression.toFirSourceElement() source = expression.toFirSourceElement()
@@ -2315,7 +2322,7 @@ open class RawFirBuilder(
}.configure(target) { expression.body.toFirBlock() } }.configure(target) { expression.body.toFirBlock() }
} }
override fun visitWhileExpression(expression: KtWhileExpression, data: Unit): FirElement { override fun visitWhileExpression(expression: KtWhileExpression, data: FirElement?): FirElement {
val target: FirLoopTarget val target: FirLoopTarget
return FirWhileLoopBuilder().apply { return FirWhileLoopBuilder().apply {
source = expression.toFirSourceElement() source = expression.toFirSourceElement()
@@ -2326,7 +2333,7 @@ open class RawFirBuilder(
}.configure(target) { expression.body.toFirBlock() } }.configure(target) { expression.body.toFirBlock() }
} }
override fun visitForExpression(expression: KtForExpression, data: Unit?): FirElement { override fun visitForExpression(expression: KtForExpression, data: FirElement?): FirElement {
val rangeExpression = expression.loopRange.toFirExpression("No range in for loop") val rangeExpression = expression.loopRange.toFirExpression("No range in for loop")
val ktParameter = expression.loopParameter val ktParameter = expression.loopParameter
val fakeSource = expression.toKtPsiSourceElement(KtFakeSourceElementKind.DesugaredForLoop) val fakeSource = expression.toKtPsiSourceElement(KtFakeSourceElementKind.DesugaredForLoop)
@@ -2398,19 +2405,19 @@ open class RawFirBuilder(
} }
} }
override fun visitBreakExpression(expression: KtBreakExpression, data: Unit): FirElement { override fun visitBreakExpression(expression: KtBreakExpression, data: FirElement?): FirElement {
return FirBreakExpressionBuilder().apply { return FirBreakExpressionBuilder().apply {
source = expression.toFirSourceElement() source = expression.toFirSourceElement()
}.bindLabel(expression).build() }.bindLabel(expression).build()
} }
override fun visitContinueExpression(expression: KtContinueExpression, data: Unit): FirElement { override fun visitContinueExpression(expression: KtContinueExpression, data: FirElement?): FirElement {
return FirContinueExpressionBuilder().apply { return FirContinueExpressionBuilder().apply {
source = expression.toFirSourceElement() source = expression.toFirSourceElement()
}.bindLabel(expression).build() }.bindLabel(expression).build()
} }
override fun visitBinaryExpression(expression: KtBinaryExpression, data: Unit): FirElement { override fun visitBinaryExpression(expression: KtBinaryExpression, data: FirElement?): FirElement {
val operationToken = expression.operationToken val operationToken = expression.operationToken
if (operationToken == IDENTIFIER) { if (operationToken == IDENTIFIER) {
@@ -2478,7 +2485,7 @@ open class RawFirBuilder(
} }
} }
override fun visitBinaryWithTypeRHSExpression(expression: KtBinaryExpressionWithTypeRHS, data: Unit): FirElement { override fun visitBinaryWithTypeRHSExpression(expression: KtBinaryExpressionWithTypeRHS, data: FirElement?): FirElement {
return buildTypeOperatorCall { return buildTypeOperatorCall {
source = expression.toFirSourceElement() source = expression.toFirSourceElement()
operation = expression.operationReference.getReferencedNameElementType().toFirOperation() operation = expression.operationReference.getReferencedNameElementType().toFirOperation()
@@ -2487,7 +2494,7 @@ open class RawFirBuilder(
} }
} }
override fun visitIsExpression(expression: KtIsExpression, data: Unit): FirElement { override fun visitIsExpression(expression: KtIsExpression, data: FirElement?): FirElement {
return buildTypeOperatorCall { return buildTypeOperatorCall {
source = expression.toFirSourceElement() source = expression.toFirSourceElement()
operation = if (expression.isNegated) FirOperation.NOT_IS else FirOperation.IS operation = if (expression.isNegated) FirOperation.NOT_IS else FirOperation.IS
@@ -2496,7 +2503,7 @@ open class RawFirBuilder(
} }
} }
override fun visitUnaryExpression(expression: KtUnaryExpression, data: Unit): FirElement { override fun visitUnaryExpression(expression: KtUnaryExpression, data: FirElement?): FirElement {
val operationToken = expression.operationToken val operationToken = expression.operationToken
val argument = expression.baseExpression val argument = expression.baseExpression
val conventionCallName = operationToken.toUnaryName() val conventionCallName = operationToken.toUnaryName()
@@ -2577,7 +2584,7 @@ open class RawFirBuilder(
} }
} }
override fun visitCallExpression(expression: KtCallExpression, data: Unit): FirElement { override fun visitCallExpression(expression: KtCallExpression, data: FirElement?): FirElement {
val source = expression.toFirSourceElement() val source = expression.toFirSourceElement()
val (calleeReference, explicitReceiver, isImplicitInvoke) = splitToCalleeAndReceiver(expression.calleeExpression, source) val (calleeReference, explicitReceiver, isImplicitInvoke) = splitToCalleeAndReceiver(expression.calleeExpression, source)
@@ -2604,7 +2611,7 @@ open class RawFirBuilder(
}.build() }.build()
} }
override fun visitArrayAccessExpression(expression: KtArrayAccessExpression, data: Unit): FirElement { override fun visitArrayAccessExpression(expression: KtArrayAccessExpression, data: FirElement?): FirElement {
val arrayExpression = expression.arrayExpression val arrayExpression = expression.arrayExpression
val setArgument = context.arraySetArgument.remove(expression) val setArgument = context.arraySetArgument.remove(expression)
return buildFunctionCall { return buildFunctionCall {
@@ -2627,7 +2634,7 @@ open class RawFirBuilder(
}.pullUpSafeCallIfNecessary() }.pullUpSafeCallIfNecessary()
} }
override fun visitQualifiedExpression(expression: KtQualifiedExpression, data: Unit): FirElement { override fun visitQualifiedExpression(expression: KtQualifiedExpression, data: FirElement?): FirElement {
val receiver = expression.receiverExpression.toFirExpression("Incorrect receiver expression") val receiver = expression.receiverExpression.toFirExpression("Incorrect receiver expression")
val selector = expression.selectorExpression val selector = expression.selectorExpression
@@ -2663,7 +2670,7 @@ open class RawFirBuilder(
return firSelector return firSelector
} }
override fun visitThisExpression(expression: KtThisExpression, data: Unit): FirElement { override fun visitThisExpression(expression: KtThisExpression, data: FirElement?): FirElement {
return buildThisReceiverExpression { return buildThisReceiverExpression {
val sourceElement = expression.toFirSourceElement() val sourceElement = expression.toFirSourceElement()
source = sourceElement source = sourceElement
@@ -2674,7 +2681,7 @@ open class RawFirBuilder(
} }
} }
override fun visitSuperExpression(expression: KtSuperExpression, data: Unit): FirElement { override fun visitSuperExpression(expression: KtSuperExpression, data: FirElement?): FirElement {
val superType = expression.superTypeQualifier val superType = expression.superTypeQualifier
val theSource = expression.toFirSourceElement() val theSource = expression.toFirSourceElement()
return buildPropertyAccessExpression { return buildPropertyAccessExpression {
@@ -2687,7 +2694,7 @@ open class RawFirBuilder(
} }
} }
override fun visitParenthesizedExpression(expression: KtParenthesizedExpression, data: Unit): FirElement { override fun visitParenthesizedExpression(expression: KtParenthesizedExpression, data: FirElement?): FirElement {
context.forwardLabelUsagePermission(expression, expression.expression) context.forwardLabelUsagePermission(expression, expression.expression)
return expression.expression?.accept(this, data) return expression.expression?.accept(this, data)
?: buildErrorExpression( ?: buildErrorExpression(
@@ -2696,7 +2703,7 @@ open class RawFirBuilder(
) )
} }
override fun visitLabeledExpression(expression: KtLabeledExpression, data: Unit): FirElement { override fun visitLabeledExpression(expression: KtLabeledExpression, data: FirElement?): FirElement {
val label = expression.getTargetLabel() val label = expression.getTargetLabel()
var errorLabelSource: KtSourceElement? = null var errorLabelSource: KtSourceElement? = null
@@ -2714,7 +2721,7 @@ open class RawFirBuilder(
return buildExpressionWithErrorLabel(result, errorLabelSource, expression.toFirSourceElement()) return buildExpressionWithErrorLabel(result, errorLabelSource, expression.toFirSourceElement())
} }
override fun visitAnnotatedExpression(expression: KtAnnotatedExpression, data: Unit): FirElement { override fun visitAnnotatedExpression(expression: KtAnnotatedExpression, data: FirElement?): FirElement {
val baseExpression = expression.baseExpression val baseExpression = expression.baseExpression
context.forwardLabelUsagePermission(expression, baseExpression) context.forwardLabelUsagePermission(expression, baseExpression)
val rawResult = baseExpression?.accept(this, data) val rawResult = baseExpression?.accept(this, data)
@@ -2727,14 +2734,14 @@ open class RawFirBuilder(
return result return result
} }
override fun visitThrowExpression(expression: KtThrowExpression, data: Unit): FirElement { override fun visitThrowExpression(expression: KtThrowExpression, data: FirElement?): FirElement {
return buildThrowExpression { return buildThrowExpression {
source = expression.toFirSourceElement() source = expression.toFirSourceElement()
exception = expression.thrownExpression.toFirExpression("Nothing to throw") exception = expression.thrownExpression.toFirExpression("Nothing to throw")
} }
} }
override fun visitDestructuringDeclaration(multiDeclaration: KtDestructuringDeclaration, data: Unit): FirElement { override fun visitDestructuringDeclaration(multiDeclaration: KtDestructuringDeclaration, data: FirElement?): FirElement {
val baseVariable = generateTemporaryVariable( val baseVariable = generateTemporaryVariable(
baseModuleData, baseModuleData,
multiDeclaration.toFirSourceElement(), multiDeclaration.toFirSourceElement(),
@@ -2753,14 +2760,14 @@ open class RawFirBuilder(
} }
} }
override fun visitClassLiteralExpression(expression: KtClassLiteralExpression, data: Unit): FirElement { override fun visitClassLiteralExpression(expression: KtClassLiteralExpression, data: FirElement?): FirElement {
return buildGetClassCall { return buildGetClassCall {
source = expression.toFirSourceElement() source = expression.toFirSourceElement()
argumentList = buildUnaryArgumentList(expression.receiverExpression.toFirExpression("No receiver in class literal")) argumentList = buildUnaryArgumentList(expression.receiverExpression.toFirExpression("No receiver in class literal"))
} }
} }
override fun visitCallableReferenceExpression(expression: KtCallableReferenceExpression, data: Unit): FirElement { override fun visitCallableReferenceExpression(expression: KtCallableReferenceExpression, data: FirElement?): FirElement {
return buildCallableReferenceAccess { return buildCallableReferenceAccess {
source = expression.toFirSourceElement() source = expression.toFirSourceElement()
calleeReference = buildSimpleNamedReference { calleeReference = buildSimpleNamedReference {
@@ -2772,7 +2779,7 @@ open class RawFirBuilder(
} }
} }
override fun visitCollectionLiteralExpression(expression: KtCollectionLiteralExpression, data: Unit): FirElement { override fun visitCollectionLiteralExpression(expression: KtCollectionLiteralExpression, data: FirElement?): FirElement {
return buildArrayOfCall { return buildArrayOfCall {
source = expression.toFirSourceElement() source = expression.toFirSourceElement()
argumentList = buildArgumentList { argumentList = buildArgumentList {
@@ -2783,7 +2790,7 @@ open class RawFirBuilder(
} }
} }
override fun visitExpression(expression: KtExpression, data: Unit): FirElement { override fun visitExpression(expression: KtExpression, data: FirElement?): FirElement {
return buildExpressionStub { return buildExpressionStub {
source = expression.toFirSourceElement() source = expression.toFirSourceElement()
} }
@@ -41,38 +41,16 @@ import kotlin.script.experimental.host.StringScriptSource
class FirScriptConfiguratorExtensionImpl( class FirScriptConfiguratorExtensionImpl(
session: FirSession, session: FirSession,
// TODO: left here because it seems it will be needed soon, remove supression if used or remove the param if it is not the case // TODO: left here because it seems it will be needed soon, remove supression if used or remove the param if it is not the case
@Suppress("UNUSED_PARAMETER") hostConfiguration: ScriptingHostConfiguration @Suppress("UNUSED_PARAMETER") hostConfiguration: ScriptingHostConfiguration,
) : FirScriptConfiguratorExtension(session) { ) : FirScriptConfiguratorExtension(session) {
@OptIn(SymbolInternals::class) @OptIn(SymbolInternals::class)
override fun FirScriptBuilder.configure(fileBuilder: FirFileBuilder) { override fun FirScriptBuilder.configure(fileBuilder: FirFileBuilder) {
val sourceFile = fileBuilder.sourceFile ?: return
// TODO: rewrite/extract decision logic for clarity withConfigurationIfAny(sourceFile) { configuration ->
val compilationConfiguration = session.scriptDefinitionProviderService?.let { providerService -> // TODO: rewrite/extract decision logic for clarity
fileBuilder.sourceFile?.toSourceCode()?.let { script -> configuration[ScriptCompilationConfiguration.baseClass]?.let { baseClass ->
val ktFile = (script as? KtFileScriptSource)?.ktFile ?: error("only PSI scripts are supported at the moment")
providerService.configurationProvider?.getScriptConfigurationResult(ktFile)?.valueOrNull()?.configuration
?: providerService.definitionProvider?.findDefinition(script)?.compilationConfiguration
} ?: providerService.definitionProvider?.getDefaultDefinition()?.compilationConfiguration
}
if (compilationConfiguration != null) {
compilationConfiguration[ScriptCompilationConfiguration.defaultImports]?.forEach { defaultImport ->
val trimmed = defaultImport.trim()
val endsWithStar = trimmed.endsWith("*")
val stripped = if (endsWithStar) trimmed.substring(0, trimmed.length - 2) else trimmed
val fqName = FqName.fromSegments(stripped.split("."))
fileBuilder.imports += buildImport {
fileBuilder.sourceFile?.project()?.let {
val dummyElement = KtPsiFactory(it, markGenerated = true).createColon()
source = KtFakeSourceElement(dummyElement, KtFakeSourceElementKind.ImplicitImport)
}
importedFqName = fqName
isAllUnder = endsWithStar
}
}
compilationConfiguration[ScriptCompilationConfiguration.baseClass]?.let { baseClass ->
val baseClassFqn = FqName.fromSegments(baseClass.typeName.split(".")) val baseClassFqn = FqName.fromSegments(baseClass.typeName.split("."))
contextReceivers.add(buildContextReceiverWithFqName(baseClassFqn)) contextReceivers.add(buildContextReceiverWithFqName(baseClassFqn))
@@ -88,8 +66,8 @@ class FirScriptConfiguratorExtensionImpl(
origin = FirDeclarationOrigin.ScriptCustomization origin = FirDeclarationOrigin.ScriptCustomization
// TODO: copy type parameters? // TODO: copy type parameters?
returnTypeRef = baseCtorParameter.returnTypeRef returnTypeRef = baseCtorParameter.returnTypeRef
this.name = baseCtorParameter.name name = baseCtorParameter.name
this.symbol = FirPropertySymbol(this.name) symbol = FirPropertySymbol(name)
status = FirDeclarationStatusImpl(Visibilities.Local, Modality.FINAL) status = FirDeclarationStatusImpl(Visibilities.Local, Modality.FINAL)
isLocal = true isLocal = true
isVar = false isVar = false
@@ -98,12 +76,10 @@ class FirScriptConfiguratorExtensionImpl(
} }
} }
} }
configuration[ScriptCompilationConfiguration.implicitReceivers]?.forEach { implicitReceiver ->
compilationConfiguration[ScriptCompilationConfiguration.implicitReceivers]?.forEach { implicitReceiver ->
contextReceivers.add(buildContextReceiverWithFqName(FqName.fromSegments(implicitReceiver.typeName.split(".")))) contextReceivers.add(buildContextReceiverWithFqName(FqName.fromSegments(implicitReceiver.typeName.split("."))))
} }
configuration[ScriptCompilationConfiguration.providedProperties]?.forEach { propertyName, propertyType ->
compilationConfiguration[ScriptCompilationConfiguration.providedProperties]?.forEach { propertyName, propertyType ->
val typeRef = buildUserTypeRef { val typeRef = buildUserTypeRef {
isMarkedNullable = propertyType.isNullable isMarkedNullable = propertyType.isNullable
propertyType.typeName.split(".").forEach { propertyType.typeName.split(".").forEach {
@@ -115,21 +91,53 @@ class FirScriptConfiguratorExtensionImpl(
moduleData = session.moduleData moduleData = session.moduleData
origin = FirDeclarationOrigin.ScriptCustomization origin = FirDeclarationOrigin.ScriptCustomization
returnTypeRef = typeRef returnTypeRef = typeRef
this.name = Name.identifier(propertyName) name = Name.identifier(propertyName)
this.symbol = FirPropertySymbol(this.name) symbol = FirPropertySymbol(name)
status = FirDeclarationStatusImpl(Visibilities.Local, Modality.FINAL) status = FirDeclarationStatusImpl(Visibilities.Local, Modality.FINAL)
isLocal = true isLocal = true
isVar = false isVar = false
} }
) )
} }
configuration[ScriptCompilationConfiguration.annotationsForSamWithReceivers]?.forEach {
_knownAnnotationsForSamWithReceiver.add(it.typeName)
}
compilationConfiguration[ScriptCompilationConfiguration.annotationsForSamWithReceivers]?.forEach { configuration[ScriptCompilationConfiguration.defaultImports]?.forEach { defaultImport ->
val trimmed = defaultImport.trim()
val endsWithStar = trimmed.endsWith("*")
val stripped = if (endsWithStar) trimmed.substring(0, trimmed.length - 2) else trimmed
val fqName = FqName.fromSegments(stripped.split("."))
fileBuilder.imports += buildImport {
fileBuilder.sourceFile?.project()?.let {
val dummyElement = KtPsiFactory(it, markGenerated = true).createColon()
source = KtFakeSourceElement(dummyElement, KtFakeSourceElementKind.ImplicitImport)
}
importedFqName = fqName
isAllUnder = endsWithStar
}
}
configuration[ScriptCompilationConfiguration.annotationsForSamWithReceivers]?.forEach {
_knownAnnotationsForSamWithReceiver.add(it.typeName) _knownAnnotationsForSamWithReceiver.add(it.typeName)
} }
} }
} }
private fun withConfigurationIfAny(file: KtSourceFile, body: (ScriptCompilationConfiguration) -> Unit) {
val configuration = session.scriptDefinitionProviderService?.let { providerService ->
val sourceCode = file.toSourceCode()
val ktFile = sourceCode?.originalKtFile()
with(providerService) {
ktFile?.let { configurationFor(it) }
?: sourceCode?.let { configurationFor(it) }
?: defaultConfiguration()
}
}
configuration?.let { body.invoke(it) }
}
private fun buildContextReceiverWithFqName(baseClassFqn: FqName) = private fun buildContextReceiverWithFqName(baseClassFqn: FqName) =
buildContextReceiver { buildContextReceiver {
typeRef = buildUserTypeRef { typeRef = buildUserTypeRef {
@@ -156,6 +164,19 @@ class FirScriptConfiguratorExtensionImpl(
private fun KtSourceFile.project(): Project? = (toSourceCode() as? KtFileScriptSource)?.ktFile?.project private fun KtSourceFile.project(): Project? = (toSourceCode() as? KtFileScriptSource)?.ktFile?.project
private fun SourceCode.originalKtFile(): KtFile =
(this as? KtFileScriptSource)?.ktFile?.originalFile as? KtFile
?: error("only PSI scripts are supported at the moment")
private fun FirScriptDefinitionProviderService.configurationFor(file: KtFile): ScriptCompilationConfiguration? =
configurationProvider?.getScriptConfigurationResult(file)?.valueOrNull()?.configuration
private fun FirScriptDefinitionProviderService.configurationFor(sourceCode: SourceCode): ScriptCompilationConfiguration? =
definitionProvider?.findDefinition(sourceCode)?.compilationConfiguration
private fun FirScriptDefinitionProviderService.defaultConfiguration(): ScriptCompilationConfiguration? =
definitionProvider?.getDefaultDefinition()?.compilationConfiguration
fun KtSourceFile.toSourceCode(): SourceCode? = when (this) { fun KtSourceFile.toSourceCode(): SourceCode? = when (this) {
is KtPsiSourceFile -> (psiFile as? KtFile)?.let(::KtFileScriptSource) ?: VirtualFileScriptSource(psiFile.virtualFile) is KtPsiSourceFile -> (psiFile as? KtFile)?.let(::KtFileScriptSource) ?: VirtualFileScriptSource(psiFile.virtualFile)
is KtVirtualFileSourceFile -> VirtualFileScriptSource(virtualFile) is KtVirtualFileSourceFile -> VirtualFileScriptSource(virtualFile)