FIR IDE: common super type of given KtExpression's

This commit is contained in:
Jinseong Jeon
2021-07-15 14:06:04 -07:00
committed by Ilya Kirillov
parent 5690b4d8c2
commit f19a501cc7
4 changed files with 49 additions and 14 deletions
@@ -13,6 +13,7 @@ import org.jetbrains.kotlin.psi.KtExpression
import org.jetbrains.kotlin.psi.KtTypeReference import org.jetbrains.kotlin.psi.KtTypeReference
public abstract class KtPsiTypeProvider : KtAnalysisSessionComponent() { public abstract class KtPsiTypeProvider : KtAnalysisSessionComponent() {
public abstract fun commonSuperType(expressions: Collection<KtExpression>, context: KtExpression, mode: TypeMappingMode): PsiType?
public abstract fun getPsiTypeForKtExpression(expression: KtExpression, mode: TypeMappingMode): PsiType public abstract fun getPsiTypeForKtExpression(expression: KtExpression, mode: TypeMappingMode): PsiType
public abstract fun getPsiTypeForKtDeclaration(ktDeclaration: KtDeclaration, mode: TypeMappingMode): PsiType public abstract fun getPsiTypeForKtDeclaration(ktDeclaration: KtDeclaration, mode: TypeMappingMode): PsiType
public abstract fun getPsiTypeForKtTypeReference(ktTypeReference: KtTypeReference, mode: TypeMappingMode): PsiType public abstract fun getPsiTypeForKtTypeReference(ktTypeReference: KtTypeReference, mode: TypeMappingMode): PsiType
@@ -23,6 +24,13 @@ public abstract class KtPsiTypeProvider : KtAnalysisSessionComponent() {
} }
public interface KtPsiTypeProviderMixIn : KtAnalysisSessionMixIn { public interface KtPsiTypeProviderMixIn : KtAnalysisSessionMixIn {
public fun commonSuperType(
expressions: Collection<KtExpression>,
context: KtExpression,
mode: TypeMappingMode = TypeMappingMode.DEFAULT
): PsiType? =
analysisSession.psiTypeProvider.commonSuperType(expressions, context, mode)
public fun KtExpression.getPsiType(mode: TypeMappingMode = TypeMappingMode.DEFAULT): PsiType = public fun KtExpression.getPsiType(mode: TypeMappingMode = TypeMappingMode.DEFAULT): PsiType =
analysisSession.psiTypeProvider.getPsiTypeForKtExpression(this, mode) analysisSession.psiTypeProvider.getPsiTypeForKtExpression(this, mode)
@@ -18,6 +18,7 @@ import org.jetbrains.kotlin.idea.fir.low.level.api.api.getOrBuildFirSafe
import org.jetbrains.kotlin.idea.frontend.api.components.KtExpressionTypeProvider import org.jetbrains.kotlin.idea.frontend.api.components.KtExpressionTypeProvider
import org.jetbrains.kotlin.idea.frontend.api.fir.KtFirAnalysisSession import org.jetbrains.kotlin.idea.frontend.api.fir.KtFirAnalysisSession
import org.jetbrains.kotlin.idea.frontend.api.fir.utils.getReferencedElementType import org.jetbrains.kotlin.idea.frontend.api.fir.utils.getReferencedElementType
import org.jetbrains.kotlin.idea.frontend.api.fir.utils.unwrap
import org.jetbrains.kotlin.idea.frontend.api.tokens.ValidityToken import org.jetbrains.kotlin.idea.frontend.api.tokens.ValidityToken
import org.jetbrains.kotlin.idea.frontend.api.types.KtClassErrorType import org.jetbrains.kotlin.idea.frontend.api.types.KtClassErrorType
import org.jetbrains.kotlin.idea.frontend.api.types.KtType import org.jetbrains.kotlin.idea.frontend.api.types.KtType
@@ -43,15 +44,6 @@ internal class KtFirExpressionTypeProvider(
} }
} }
private fun KtExpression.unwrap(): KtExpression {
return when (this) {
is KtLabeledExpression -> baseExpression?.unwrap()
is KtAnnotatedExpression -> baseExpression?.unwrap()
is KtObjectLiteralExpression -> objectDeclaration
else -> null
} ?: this
}
override fun getExpectedType(expression: PsiElement): KtType? { override fun getExpectedType(expression: PsiElement): KtType? {
val expectedType = getExpectedTypeByReturnExpression(expression) val expectedType = getExpectedTypeByReturnExpression(expression)
?: getExpressionTypeByIfOrBooleanCondition(expression) ?: getExpressionTypeByIfOrBooleanCondition(expression)
@@ -12,6 +12,7 @@ import org.jetbrains.kotlin.fir.expressions.FirExpression
import org.jetbrains.kotlin.fir.expressions.FirGetClassCall import org.jetbrains.kotlin.fir.expressions.FirGetClassCall
import org.jetbrains.kotlin.fir.expressions.FirStatement import org.jetbrains.kotlin.fir.expressions.FirStatement
import org.jetbrains.kotlin.fir.references.FirNamedReference import org.jetbrains.kotlin.fir.references.FirNamedReference
import org.jetbrains.kotlin.fir.typeContext
import org.jetbrains.kotlin.fir.types.* import org.jetbrains.kotlin.fir.types.*
import org.jetbrains.kotlin.idea.fir.low.level.api.api.getOrBuildFir import org.jetbrains.kotlin.idea.fir.low.level.api.api.getOrBuildFir
import org.jetbrains.kotlin.idea.fir.low.level.api.api.getOrBuildFirOfType import org.jetbrains.kotlin.idea.fir.low.level.api.api.getOrBuildFirOfType
@@ -32,15 +33,36 @@ internal class KtFirPsiTypeProvider(
override val analysisSession: KtFirAnalysisSession, override val analysisSession: KtFirAnalysisSession,
override val token: ValidityToken, override val token: ValidityToken,
) : KtPsiTypeProvider(), KtFirAnalysisSessionComponent { ) : KtPsiTypeProvider(), KtFirAnalysisSessionComponent {
override fun commonSuperType(
expressions: Collection<KtExpression>,
context: KtExpression,
mode: TypeMappingMode
): PsiType? = withValidityAssertion {
val unitType = analysisSession.rootModuleSession.builtinTypes.unitType.type
analysisSession.rootModuleSession.typeContext
.commonSuperTypeOrNull(expressions.map { e -> e.getConeType(unitType) { it } })
?.asPsiType(mode, context)
}
override fun getPsiTypeForKtExpression( override fun getPsiTypeForKtExpression(
expression: KtExpression, expression: KtExpression,
mode: TypeMappingMode, mode: TypeMappingMode,
): PsiType = withValidityAssertion { ): PsiType = withValidityAssertion {
when (val fir = expression.getOrBuildFir(firResolveState)) { expression.getConeType(PsiType.VOID) {
is FirExpression -> fir.typeRef.coneType.asPsiType(mode, expression) it.asPsiType(mode, expression)
is FirNamedReference -> fir.getReferencedElementType().asPsiType(mode, expression) }
is FirStatement -> PsiType.VOID }
else -> throwUnexpectedFirElementError(fir, expression)
private inline fun <T> KtExpression.getConeType(
defaultType: T,
coneTypeConverter: (ConeKotlinType) -> T,
): T = withValidityAssertion {
when (val fir = this.getOrBuildFir(firResolveState)) {
is FirExpression -> coneTypeConverter(fir.typeRef.coneType)
is FirNamedReference -> coneTypeConverter(fir.getReferencedElementType())
is FirStatement -> defaultType
else -> throwUnexpectedFirElementError(fir, this)
} }
} }
@@ -23,6 +23,19 @@ import org.jetbrains.kotlin.idea.frontend.api.symbols.markers.KtSimpleConstantVa
import org.jetbrains.kotlin.idea.frontend.api.symbols.markers.KtUnsupportedConstantValue import org.jetbrains.kotlin.idea.frontend.api.symbols.markers.KtUnsupportedConstantValue
import org.jetbrains.kotlin.idea.frontend.api.types.KtTypeNullability import org.jetbrains.kotlin.idea.frontend.api.types.KtTypeNullability
import org.jetbrains.kotlin.idea.references.FirReferenceResolveHelper import org.jetbrains.kotlin.idea.references.FirReferenceResolveHelper
import org.jetbrains.kotlin.psi.KtAnnotatedExpression
import org.jetbrains.kotlin.psi.KtExpression
import org.jetbrains.kotlin.psi.KtLabeledExpression
import org.jetbrains.kotlin.psi.KtObjectLiteralExpression
internal fun KtExpression.unwrap(): KtExpression {
return when (this) {
is KtLabeledExpression -> baseExpression?.unwrap()
is KtAnnotatedExpression -> baseExpression?.unwrap()
is KtObjectLiteralExpression -> objectDeclaration
else -> null
} ?: this
}
internal fun FirNamedReference.getReferencedElementType(): ConeKotlinType { internal fun FirNamedReference.getReferencedElementType(): ConeKotlinType {
val symbols = when (this) { val symbols = when (this) {