[FIR] Properly collect overriddens for method enhancement

If some java class has multiple supertypes then we need to collect
  overriddens from all those types directly, even if superTypeScope
  (which is FirTypeIntersectionScope in this case) returns only
  one symbol from one of this types (not intersection one)

This is needed to proper enhancement in cases when some type occurs
  multiple times in supertypes graph with different nullability
  of arguments:

class ConcurrentHashMap<K, V> : AbstractMap<K!, V!>, MutableMap<K, V>

If we try to find method `get(key: K): V` supertype scope returns
  `AbstractMap.get(key: K!): V!` (because it actually overrides
  `MutableMap(key: K): V?`), but we need to get both symbols to
  properly enhance types for `ConcurrentHashMap.remove`
This commit is contained in:
Dmitriy Novozhilov
2021-11-19 18:16:07 +03:00
parent 01c0cf80d0
commit 9807c67ae4
42 changed files with 393 additions and 301 deletions
@@ -19,7 +19,7 @@ abstract class AbstractFirUseSiteMemberScope(
val classId: ClassId,
session: FirSession,
overrideChecker: FirOverrideChecker,
protected val superTypesScope: FirTypeScope,
val superTypesScope: FirTypeScope,
protected val declaredMemberScope: FirContainingNamesAwareScope
) : AbstractFirOverrideScope(session, overrideChecker) {
@@ -159,8 +159,7 @@ class FirTypeIntersectionScope private constructor(
}.withScope(scopeForMostSpecific)
}
private fun <S : FirCallableSymbol<*>>
MutableList<MemberWithBaseScope<S>>.calcBaseMembersForIntersectionOverride(): List<MemberWithBaseScope<S>> {
private fun <S : FirCallableSymbol<*>> List<MemberWithBaseScope<S>>.calcBaseMembersForIntersectionOverride(): List<MemberWithBaseScope<S>> {
if (size == 1) return this
val unwrappedMemberSet = mutableSetOf<MemberWithBaseScope<S>>()
for ((member, scope) in this) {
@@ -195,8 +194,9 @@ class FirTypeIntersectionScope private constructor(
}
}
}
removeIf { (member, _) -> member.fir.unwrapSubstitutionOverrides().symbol in baseMembers }
return this
val result = this.toMutableList()
result.removeIf { (member, _) -> member.fir.unwrapSubstitutionOverrides().symbol in baseMembers }
return result
}
private fun <D : FirCallableSymbol<*>> chooseIntersectionOverrideModality(
@@ -514,7 +514,7 @@ class FirTypeIntersectionScope private constructor(
}
@Suppress("UNCHECKED_CAST")
private fun <S : FirCallableSymbol<*>> getDirectOverriddenSymbols(symbol: S): Collection<MemberWithBaseScope<S>> {
fun <S : FirCallableSymbol<*>> getDirectOverriddenSymbols(symbol: S): Collection<MemberWithBaseScope<S>> {
val intersectionOverride = intersectionOverrides.getValueIfComputed(symbol)
val allDirectOverridden = overriddenSymbols[symbol].orEmpty() + intersectionOverride?.let {
overriddenSymbols[it.member]