[FIR] Don't force calculation of return type of overrides during status resolution

This commit is contained in:
Dmitriy Novozhilov
2022-03-15 11:10:35 +03:00
committed by teamcity
parent effa3f197a
commit ae0ce57b2c
6 changed files with 34 additions and 20 deletions
@@ -10,7 +10,11 @@ import org.jetbrains.kotlin.fir.symbols.impl.FirCallableSymbol
import org.jetbrains.kotlin.fir.types.FirResolvedTypeRef import org.jetbrains.kotlin.fir.types.FirResolvedTypeRef
abstract class ReturnTypeCalculator { abstract class ReturnTypeCalculator {
abstract fun tryCalculateReturnType(declaration: FirTypedDeclaration): FirResolvedTypeRef abstract fun tryCalculateReturnTypeOrNull(declaration: FirTypedDeclaration): FirResolvedTypeRef?
fun tryCalculateReturnType(declaration: FirTypedDeclaration): FirResolvedTypeRef {
return tryCalculateReturnTypeOrNull(declaration)!!
}
fun tryCalculateReturnType(symbol: FirCallableSymbol<*>): FirResolvedTypeRef { fun tryCalculateReturnType(symbol: FirCallableSymbol<*>): FirResolvedTypeRef {
return tryCalculateReturnType(symbol.fir) return tryCalculateReturnType(symbol.fir)
@@ -14,7 +14,7 @@ import org.jetbrains.kotlin.fir.types.FirResolvedTypeRef
import org.jetbrains.kotlin.fir.types.builder.buildErrorTypeRef import org.jetbrains.kotlin.fir.types.builder.buildErrorTypeRef
object ReturnTypeCalculatorForFullBodyResolve : ReturnTypeCalculator() { object ReturnTypeCalculatorForFullBodyResolve : ReturnTypeCalculator() {
override fun tryCalculateReturnType(declaration: FirTypedDeclaration): FirResolvedTypeRef { override fun tryCalculateReturnTypeOrNull(declaration: FirTypedDeclaration): FirResolvedTypeRef? {
val returnTypeRef = declaration.returnTypeRef val returnTypeRef = declaration.returnTypeRef
if (returnTypeRef is FirResolvedTypeRef) return returnTypeRef if (returnTypeRef is FirResolvedTypeRef) return returnTypeRef
if (declaration.origin.fromSupertypes) { if (declaration.origin.fromSupertypes) {
@@ -16,7 +16,7 @@ import org.jetbrains.kotlin.fir.types.FirResolvedTypeRef
import org.jetbrains.kotlin.fir.types.FirTypeRef import org.jetbrains.kotlin.fir.types.FirTypeRef
abstract class FakeOverrideTypeCalculator { abstract class FakeOverrideTypeCalculator {
abstract fun computeReturnType(declaration: FirTypedDeclaration): FirTypeRef abstract fun computeReturnType(declaration: FirTypedDeclaration): FirTypeRef?
object DoNothing : FakeOverrideTypeCalculator() { object DoNothing : FakeOverrideTypeCalculator() {
override fun computeReturnType(declaration: FirTypedDeclaration): FirTypeRef { override fun computeReturnType(declaration: FirTypedDeclaration): FirTypeRef {
@@ -25,17 +25,17 @@ abstract class FakeOverrideTypeCalculator {
} }
object Forced : FakeOverrideTypeCalculator() { object Forced : FakeOverrideTypeCalculator() {
override fun computeReturnType(declaration: FirTypedDeclaration): FirResolvedTypeRef { override fun computeReturnType(declaration: FirTypedDeclaration): FirResolvedTypeRef? {
val fakeOverrideSubstitution = declaration.attributes.fakeOverrideSubstitution val fakeOverrideSubstitution = declaration.attributes.fakeOverrideSubstitution
?: return declaration.returnTypeRef as FirResolvedTypeRef ?: return declaration.returnTypeRef as? FirResolvedTypeRef
synchronized(fakeOverrideSubstitution) { synchronized(fakeOverrideSubstitution) {
if (declaration.attributes.fakeOverrideSubstitution == null) { if (declaration.attributes.fakeOverrideSubstitution == null) {
return declaration.returnTypeRef as FirResolvedTypeRef return declaration.returnTypeRef as FirResolvedTypeRef
} }
declaration.attributes.fakeOverrideSubstitution = null
val (substitutor, baseSymbol) = fakeOverrideSubstitution val (substitutor, baseSymbol) = fakeOverrideSubstitution
val baseDeclaration = baseSymbol.fir as FirTypedDeclaration val baseDeclaration = baseSymbol.fir as FirTypedDeclaration
val baseReturnType = computeReturnType(baseDeclaration).type val baseReturnType = computeReturnType(baseDeclaration)?.type ?: return null
declaration.attributes.fakeOverrideSubstitution = null
val coneType = substitutor.substituteOrSelf(baseReturnType) val coneType = substitutor.substituteOrSelf(baseReturnType)
val returnType = declaration.returnTypeRef.resolvedTypeFromPrototype(coneType) val returnType = declaration.returnTypeRef.resolvedTypeFromPrototype(coneType)
declaration.replaceReturnTypeRef(returnType) declaration.replaceReturnTypeRef(returnType)
@@ -13,6 +13,7 @@ import org.jetbrains.kotlin.fir.declarations.FirProperty
import org.jetbrains.kotlin.fir.declarations.FirPropertyAccessor import org.jetbrains.kotlin.fir.declarations.FirPropertyAccessor
import org.jetbrains.kotlin.fir.declarations.FirSimpleFunction import org.jetbrains.kotlin.fir.declarations.FirSimpleFunction
import org.jetbrains.kotlin.fir.declarations.utils.visibility import org.jetbrains.kotlin.fir.declarations.utils.visibility
import org.jetbrains.kotlin.fir.resolve.transformers.ReturnTypeCalculator
import org.jetbrains.kotlin.fir.scopes.impl.buildSubstitutorForOverridesCheck import org.jetbrains.kotlin.fir.scopes.impl.buildSubstitutorForOverridesCheck
import org.jetbrains.kotlin.fir.scopes.impl.similarFunctionsOrBothProperties import org.jetbrains.kotlin.fir.scopes.impl.similarFunctionsOrBothProperties
import org.jetbrains.kotlin.fir.symbols.impl.FirCallableSymbol import org.jetbrains.kotlin.fir.symbols.impl.FirCallableSymbol
@@ -28,7 +29,8 @@ import java.util.*
class FirOverrideService(val session: FirSession) : FirSessionComponent { class FirOverrideService(val session: FirSession) : FirSessionComponent {
fun <D : FirCallableSymbol<*>> selectMostSpecificInEachOverridableGroup( fun <D : FirCallableSymbol<*>> selectMostSpecificInEachOverridableGroup(
members: Collection<MemberWithBaseScope<D>>, members: Collection<MemberWithBaseScope<D>>,
overrideChecker: FirOverrideChecker overrideChecker: FirOverrideChecker,
returnTypeCalculator: ReturnTypeCalculator
): Collection<MemberWithBaseScope<D>> { ): Collection<MemberWithBaseScope<D>> {
if (members.size <= 1) return members if (members.size <= 1) return members
val queue = LinkedList(members) val queue = LinkedList(members)
@@ -46,10 +48,10 @@ class FirOverrideService(val session: FirSession) : FirSessionComponent {
continue continue
} }
val mostSpecific = selectMostSpecificMember(overridableGroup) val mostSpecific = selectMostSpecificMember(overridableGroup, returnTypeCalculator)
overridableGroup.filterNotTo(conflictedHandles) { overridableGroup.filterNotTo(conflictedHandles) {
isMoreSpecific(mostSpecific.member, it.member) isMoreSpecific(mostSpecific.member, it.member, returnTypeCalculator)
} }
if (conflictedHandles.isNotEmpty()) { if (conflictedHandles.isNotEmpty()) {
@@ -87,7 +89,10 @@ class FirOverrideService(val session: FirSession) : FirSessionComponent {
return result return result
} }
fun <D : FirCallableSymbol<*>> selectMostSpecificMember(overridables: Collection<MemberWithBaseScope<D>>): MemberWithBaseScope<D> { fun <D : FirCallableSymbol<*>> selectMostSpecificMember(
overridables: Collection<MemberWithBaseScope<D>>,
returnTypeCalculator: ReturnTypeCalculator
): MemberWithBaseScope<D> {
require(overridables.isNotEmpty()) { "Should have at least one overridable symbol" } require(overridables.isNotEmpty()) { "Should have at least one overridable symbol" }
if (overridables.size == 1) { if (overridables.size == 1) {
return overridables.first() return overridables.first()
@@ -97,12 +102,12 @@ class FirOverrideService(val session: FirSession) : FirSessionComponent {
var transitivelyMostSpecific: MemberWithBaseScope<D> = overridables.first() var transitivelyMostSpecific: MemberWithBaseScope<D> = overridables.first()
for (candidate in overridables) { for (candidate in overridables) {
if (overridables.all { isMoreSpecific(candidate.member, it.member) }) { if (overridables.all { isMoreSpecific(candidate.member, it.member, returnTypeCalculator) }) {
candidates.add(candidate) candidates.add(candidate)
} }
if (isMoreSpecific(candidate.member, transitivelyMostSpecific.member) && if (isMoreSpecific(candidate.member, transitivelyMostSpecific.member, returnTypeCalculator) &&
!isMoreSpecific(transitivelyMostSpecific.member, candidate.member) !isMoreSpecific(transitivelyMostSpecific.member, candidate.member, returnTypeCalculator)
) { ) {
transitivelyMostSpecific = candidate transitivelyMostSpecific = candidate
} }
@@ -123,7 +128,8 @@ class FirOverrideService(val session: FirSession) : FirSessionComponent {
private fun isMoreSpecific( private fun isMoreSpecific(
a: FirCallableSymbol<*>, a: FirCallableSymbol<*>,
b: FirCallableSymbol<*> b: FirCallableSymbol<*>,
returnTypeCalculator: ReturnTypeCalculator
): Boolean { ): Boolean {
val aFir = a.fir val aFir = a.fir
val bFir = b.fir val bFir = b.fir
@@ -132,8 +138,8 @@ class FirOverrideService(val session: FirSession) : FirSessionComponent {
val substitutor = buildSubstitutorForOverridesCheck(aFir, bFir, session) ?: return false val substitutor = buildSubstitutorForOverridesCheck(aFir, bFir, session) ?: return false
// NB: these lines throw CCE in modularized tests when changed to just .coneType (FirImplicitTypeRef) // NB: these lines throw CCE in modularized tests when changed to just .coneType (FirImplicitTypeRef)
val aReturnType = a.fir.returnTypeRef.coneTypeSafe<ConeKotlinType>()?.let(substitutor::substituteOrSelf) ?: return false val aReturnType = returnTypeCalculator.tryCalculateReturnTypeOrNull(a.fir)?.type?.let(substitutor::substituteOrSelf) ?: return false
val bReturnType = b.fir.returnTypeRef.coneTypeSafe<ConeKotlinType>() ?: return false val bReturnType = returnTypeCalculator.tryCalculateReturnTypeOrNull(b.fir)?.type ?: return false
val typeCheckerState = session.typeContext.newTypeCheckerState( val typeCheckerState = session.typeContext.newTypeCheckerState(
errorTypesEqualToAnything = false, errorTypesEqualToAnything = false,
@@ -15,6 +15,7 @@ import org.jetbrains.kotlin.fir.declarations.utils.isExpect
import org.jetbrains.kotlin.fir.declarations.utils.modality import org.jetbrains.kotlin.fir.declarations.utils.modality
import org.jetbrains.kotlin.fir.declarations.utils.visibility import org.jetbrains.kotlin.fir.declarations.utils.visibility
import org.jetbrains.kotlin.fir.resolve.substitution.ConeSubstitutor import org.jetbrains.kotlin.fir.resolve.substitution.ConeSubstitutor
import org.jetbrains.kotlin.fir.resolve.transformers.ReturnTypeCalculatorForFullBodyResolve
import org.jetbrains.kotlin.fir.scopes.* import org.jetbrains.kotlin.fir.scopes.*
import org.jetbrains.kotlin.fir.scopes.impl.FirIntersectionOverrideStorage.ContextForIntersectionOverrideConstruction import org.jetbrains.kotlin.fir.scopes.impl.FirIntersectionOverrideStorage.ContextForIntersectionOverrideConstruction
import org.jetbrains.kotlin.fir.scopes.impl.FirTypeIntersectionScopeContext.ResultOfIntersection import org.jetbrains.kotlin.fir.scopes.impl.FirTypeIntersectionScopeContext.ResultOfIntersection
@@ -163,7 +164,10 @@ class FirTypeIntersectionScopeContext(
}.takeIf { it.isNotEmpty() } ?: extractBothWaysWithPrivate }.takeIf { it.isNotEmpty() } ?: extractBothWaysWithPrivate
val baseMembersForIntersection = extractedOverrides.calcBaseMembersForIntersectionOverride() val baseMembersForIntersection = extractedOverrides.calcBaseMembersForIntersectionOverride()
if (baseMembersForIntersection.size > 1) { if (baseMembersForIntersection.size > 1) {
val (mostSpecific, scopeForMostSpecific) = overrideService.selectMostSpecificMember(baseMembersForIntersection) val (mostSpecific, scopeForMostSpecific) = overrideService.selectMostSpecificMember(
baseMembersForIntersection,
ReturnTypeCalculatorForFullBodyResolve
)
val intersectionOverrideContext = ContextForIntersectionOverrideConstruction( val intersectionOverrideContext = ContextForIntersectionOverrideConstruction(
mostSpecific, mostSpecific,
this, this,
@@ -216,7 +220,7 @@ class FirTypeIntersectionScopeContext(
// we should just take most specific member without creating intersection // we should just take most specific member without creating intersection
// A typical sample here is inheritance of the same class in different places of hierarchy // A typical sample here is inheritance of the same class in different places of hierarchy
if (unwrappedMemberSet.size == 1) { if (unwrappedMemberSet.size == 1) {
return listOf(overrideService.selectMostSpecificMember(this)) return listOf(overrideService.selectMostSpecificMember(this, ReturnTypeCalculatorForFullBodyResolve))
} }
val baseMembers = mutableSetOf<S>() val baseMembers = mutableSetOf<S>()
@@ -215,7 +215,7 @@ private class ReturnTypeCalculatorWithJump(
var outerTowerDataContexts: FirRegularTowerDataContexts? = null var outerTowerDataContexts: FirRegularTowerDataContexts? = null
override fun tryCalculateReturnType(declaration: FirTypedDeclaration): FirResolvedTypeRef { override fun tryCalculateReturnTypeOrNull(declaration: FirTypedDeclaration): FirResolvedTypeRef {
if (declaration is FirValueParameter && declaration.returnTypeRef is FirImplicitTypeRef) { if (declaration is FirValueParameter && declaration.returnTypeRef is FirImplicitTypeRef) {
// TODO? // TODO?
declaration.transformReturnTypeRef( declaration.transformReturnTypeRef(