FIR: Unify all references to FIR nodes from non-parents

This commit is contained in:
Denis Zharkov
2020-06-02 17:40:09 +03:00
parent 4a4dce1766
commit 7a22827af4
28 changed files with 133 additions and 109 deletions
@@ -562,7 +562,7 @@ abstract class BaseFirBuilder<T>(val baseSession: FirSession, val context: Conte
source = baseSource
rValue = value
calleeReference = nestedAccess.calleeReference
explicitReceiver = safeCallNonAssignment.checkedSubject.value
explicitReceiver = safeCallNonAssignment.checkedSubjectRef.value
}
safeCallNonAssignment.replaceRegularQualifiedAccess(
@@ -147,10 +147,12 @@ fun FirExpression.generateNotNullOrOther(
): FirWhenExpression {
val subjectName = Name.special("<$caseId>")
val subjectVariable = generateTemporaryVariable(session, baseSource, subjectName, this)
val subject = FirWhenSubject()
@OptIn(FirContractViolation::class)
val ref = FirExpressionRef<FirWhenExpression>()
val subjectExpression = buildWhenSubjectExpression {
source = baseSource
whenSubject = subject
whenRef = ref
}
return buildWhenExpression {
@@ -176,7 +178,7 @@ fun FirExpression.generateNotNullOrOther(
result = buildSingleExpressionBlock(generateResolvedAccessExpression(baseSource, subjectVariable))
}
}.also {
subject.bind(it)
ref.bind(it)
}
}
@@ -487,6 +489,7 @@ private fun FirExpression.checkReceiver(name: String?): Boolean {
return receiverName == name
}
fun FirModifiableQualifiedAccess.wrapWithSafeCall(receiver: FirExpression): FirSafeCallExpression {
// TODO: Refactor tree to make FirModifiableQualifiedAccess inherit FirQualifiedAccess
require(this is FirQualifiedAccess) {
@@ -494,13 +497,19 @@ fun FirModifiableQualifiedAccess.wrapWithSafeCall(receiver: FirExpression): FirS
}
val checkedSafeCallSubject = buildCheckedSafeCallSubject {
this.originalReceiverReference = FirSafeCallOriginalReceiverReference(receiver)
@OptIn(FirContractViolation::class)
this.originalReceiverRef = FirExpressionRef<FirExpression>().apply {
bind(receiver)
}
}
explicitReceiver = checkedSafeCallSubject
return buildSafeCallExpression {
this.receiver = receiver
this.checkedSubject = FirSafeCallCheckedSubjectReference(checkedSafeCallSubject)
@OptIn(FirContractViolation::class)
this.checkedSubjectRef = FirExpressionRef<FirCheckedSafeCallSubject>().apply {
bind(checkedSafeCallSubject)
}
this.regularQualifiedAccess = this@wrapWithSafeCall
this.source = this@wrapWithSafeCall.source
}
@@ -636,7 +636,9 @@ class ExpressionsConverter(
}
subjectExpression = subjectVariable?.initializer ?: subjectExpression
val hasSubject = subjectExpression != null
val subject = FirWhenSubject()
@OptIn(FirContractViolation::class)
val subject = FirExpressionRef<FirWhenExpression>()
whenEntryNodes.mapTo(whenEntries) { convertWhenEntry(it, subject.takeIf { hasSubject }) }
return buildWhenExpression {
source = whenExpression.toFirSourceElement()
@@ -676,15 +678,15 @@ class ExpressionsConverter(
* @see org.jetbrains.kotlin.parsing.KotlinExpressionParsing.parseWhenEntry
* @see org.jetbrains.kotlin.parsing.KotlinExpressionParsing.parseWhenEntryNotElse
*/
private fun convertWhenEntry(whenEntry: LighterASTNode, subject: FirWhenSubject?): WhenEntry {
private fun convertWhenEntry(whenEntry: LighterASTNode, whenRefWithSubject: FirExpressionRef<FirWhenExpression>?): WhenEntry {
var isElse = false
var firBlock: FirBlock = buildEmptyExpressionBlock()
val conditions = mutableListOf<FirExpression>()
whenEntry.forEachChildren {
when (it.tokenType) {
WHEN_CONDITION_EXPRESSION -> conditions += convertWhenConditionExpression(it, subject)
WHEN_CONDITION_IN_RANGE -> conditions += convertWhenConditionInRange(it, subject)
WHEN_CONDITION_IS_PATTERN -> conditions += convertWhenConditionIsPattern(it, subject)
WHEN_CONDITION_EXPRESSION -> conditions += convertWhenConditionExpression(it, whenRefWithSubject)
WHEN_CONDITION_IN_RANGE -> conditions += convertWhenConditionInRange(it, whenRefWithSubject)
WHEN_CONDITION_IS_PATTERN -> conditions += convertWhenConditionIsPattern(it, whenRefWithSubject)
ELSE_KEYWORD -> isElse = true
BLOCK -> firBlock = declarationsConverter.convertBlock(it)
else -> if (it.isExpression()) firBlock = declarationsConverter.convertBlock(it)
@@ -694,20 +696,20 @@ class ExpressionsConverter(
return WhenEntry(conditions, firBlock, isElse)
}
private fun convertWhenConditionExpression(whenCondition: LighterASTNode, subject: FirWhenSubject?): FirExpression {
private fun convertWhenConditionExpression(whenCondition: LighterASTNode, whenRefWithSubject: FirExpressionRef<FirWhenExpression>?): FirExpression {
var firExpression: FirExpression = buildErrorExpression(null, ConeSimpleDiagnostic("No expression in condition with expression", DiagnosticKind.Syntax))
whenCondition.forEachChildren {
when (it.tokenType) {
else -> if (it.isExpression()) firExpression = getAsFirExpression(it, "No expression in condition with expression")
}
}
return if (subject != null) {
return if (whenRefWithSubject != null) {
buildOperatorCall {
source = whenCondition.toFirSourceElement()
operation = FirOperation.EQ
argumentList = buildBinaryArgumentList(
buildWhenSubjectExpression {
whenSubject = subject
whenRef = whenRefWithSubject
}, firExpression
)
}
@@ -717,7 +719,7 @@ class ExpressionsConverter(
}
}
private fun convertWhenConditionInRange(whenCondition: LighterASTNode, subject: FirWhenSubject?): FirExpression {
private fun convertWhenConditionInRange(whenCondition: LighterASTNode, whenRefWithSubject: FirExpressionRef<FirWhenExpression>?): FirExpression {
var isNegate = false
var firExpression: FirExpression = buildErrorExpression(null, ConeSimpleDiagnostic("No range in condition with range", DiagnosticKind.Syntax))
var conditionSource: FirLightSourceElement? = null
@@ -731,9 +733,9 @@ class ExpressionsConverter(
}
}
val subjectExpression = if (subject != null) {
val subjectExpression = if (whenRefWithSubject != null) {
buildWhenSubjectExpression {
whenSubject = subject
whenRef = whenRefWithSubject
}
} else {
return buildErrorExpression {
@@ -750,7 +752,7 @@ class ExpressionsConverter(
)
}
private fun convertWhenConditionIsPattern(whenCondition: LighterASTNode, subject: FirWhenSubject?): FirExpression {
private fun convertWhenConditionIsPattern(whenCondition: LighterASTNode, whenRefWithSubject: FirExpressionRef<FirWhenExpression>?): FirExpression {
lateinit var firOperation: FirOperation
lateinit var firType: FirTypeRef
whenCondition.forEachChildren {
@@ -761,9 +763,9 @@ class ExpressionsConverter(
}
}
val subjectExpression = if (subject != null) {
val subjectExpression = if (whenRefWithSubject != null) {
buildWhenSubjectExpression {
whenSubject = subject
whenRef = whenRefWithSubject
}
} else {
return buildErrorExpression {
@@ -7,9 +7,9 @@ package org.jetbrains.kotlin.fir.builder
import org.jetbrains.kotlin.descriptors.Modality
import org.jetbrains.kotlin.descriptors.Visibilities
import org.jetbrains.kotlin.fir.FirExpressionRef
import org.jetbrains.kotlin.fir.FirSession
import org.jetbrains.kotlin.fir.FirSourceElement
import org.jetbrains.kotlin.fir.FirWhenSubject
import org.jetbrains.kotlin.fir.declarations.FirDeclarationOrigin
import org.jetbrains.kotlin.fir.declarations.FirVariable
import org.jetbrains.kotlin.fir.declarations.builder.buildProperty
@@ -24,14 +24,14 @@ import org.jetbrains.kotlin.fir.types.FirTypeRef
import org.jetbrains.kotlin.psi.*
internal fun KtWhenCondition.toFirWhenCondition(
subject: FirWhenSubject,
whenRefWithSibject: FirExpressionRef<FirWhenExpression>,
convert: KtExpression?.(String) -> FirExpression,
toFirOrErrorTypeRef: KtTypeReference?.() -> FirTypeRef,
): FirExpression {
val baseSource = this.toFirPsiSourceElement()
val firSubjectExpression = buildWhenSubjectExpression {
source = baseSource
whenSubject = subject
whenRef = whenRefWithSibject
}
return when (this) {
is KtWhenConditionWithExpression -> {
@@ -68,7 +68,7 @@ internal fun KtWhenCondition.toFirWhenCondition(
internal fun Array<KtWhenCondition>.toFirWhenCondition(
baseSource: FirSourceElement?,
subject: FirWhenSubject,
subject: FirExpressionRef<FirWhenExpression>,
convert: KtExpression?.(String) -> FirExpression,
toFirOrErrorTypeRef: KtTypeReference?.() -> FirTypeRef,
): FirExpression {
@@ -121,4 +121,4 @@ internal fun generateDestructuringBlock(
}
}
}
}
}
@@ -1343,7 +1343,8 @@ class RawFirBuilder(
else -> null
}
val hasSubject = subjectExpression != null
val subject = FirWhenSubject()
@OptIn(FirContractViolation::class)
val ref = FirExpressionRef<FirWhenExpression>()
return buildWhenExpression {
source = expression.toFirSourceElement()
this.subject = subjectExpression
@@ -1358,7 +1359,7 @@ class RawFirBuilder(
source = entrySource
condition = entry.conditions.toFirWhenCondition(
entrySource,
subject,
ref,
{ toFirExpression(it) },
{ toFirOrErrorType() },
)
@@ -1382,7 +1383,7 @@ class RawFirBuilder(
}
}.also {
if (hasSubject) {
subject.bind(it)
ref.bind(it)
}
}
}