Sealed class hierarchies are now correctly processed in when (by checking possible smart casts to nested sealed classes) #KT-10648 Fixed
(cherry picked from commit 2eaaf9c)
This commit is contained in:
committed by
Mikhail Glukhikh
parent
01430b4b99
commit
49fb9ff424
@@ -27,6 +27,7 @@ import org.jetbrains.kotlin.lexer.KtTokens
|
|||||||
import org.jetbrains.kotlin.psi.*
|
import org.jetbrains.kotlin.psi.*
|
||||||
import org.jetbrains.kotlin.psi.psiUtil.checkReservedPrefixWord
|
import org.jetbrains.kotlin.psi.psiUtil.checkReservedPrefixWord
|
||||||
import org.jetbrains.kotlin.resolve.BindingContext
|
import org.jetbrains.kotlin.resolve.BindingContext
|
||||||
|
import org.jetbrains.kotlin.resolve.BindingContext.SMARTCAST
|
||||||
import org.jetbrains.kotlin.resolve.BindingTrace
|
import org.jetbrains.kotlin.resolve.BindingTrace
|
||||||
import org.jetbrains.kotlin.resolve.CompileTimeConstantUtils
|
import org.jetbrains.kotlin.resolve.CompileTimeConstantUtils
|
||||||
import org.jetbrains.kotlin.resolve.DescriptorUtils
|
import org.jetbrains.kotlin.resolve.DescriptorUtils
|
||||||
@@ -139,7 +140,7 @@ private class ClassMissingCase(val descriptor: ClassDescriptor): WhenMissingCase
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
private abstract class WhenOnClassExhaustivenessChecker : WhenExhaustivenessChecker {
|
internal abstract class WhenOnClassExhaustivenessChecker : WhenExhaustivenessChecker {
|
||||||
private fun getReference(expression: KtExpression?): KtSimpleNameExpression? =
|
private fun getReference(expression: KtExpression?): KtSimpleNameExpression? =
|
||||||
when (expression) {
|
when (expression) {
|
||||||
is KtSimpleNameExpression -> expression
|
is KtSimpleNameExpression -> expression
|
||||||
@@ -223,7 +224,7 @@ private object WhenOnEnumExhaustivenessChecker : WhenOnClassExhaustivenessChecke
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
private object WhenOnSealedExhaustivenessChecker : WhenOnClassExhaustivenessChecker() {
|
internal object WhenOnSealedExhaustivenessChecker : WhenOnClassExhaustivenessChecker() {
|
||||||
override fun getMissingCases(
|
override fun getMissingCases(
|
||||||
expression: KtWhenExpression,
|
expression: KtWhenExpression,
|
||||||
context: BindingContext,
|
context: BindingContext,
|
||||||
@@ -233,8 +234,7 @@ private object WhenOnSealedExhaustivenessChecker : WhenOnClassExhaustivenessChec
|
|||||||
assert(DescriptorUtils.isSealedClass(subjectDescriptor)) {
|
assert(DescriptorUtils.isSealedClass(subjectDescriptor)) {
|
||||||
"isWhenOnSealedClassExhaustive should be called with a sealed class descriptor: $subjectDescriptor"
|
"isWhenOnSealedClassExhaustive should be called with a sealed class descriptor: $subjectDescriptor"
|
||||||
}
|
}
|
||||||
val memberClassDescriptors = LinkedHashSet<ClassDescriptor>()
|
val memberClassDescriptors = getNestedSubclasses(subjectDescriptor!!)
|
||||||
collectNestedSubclasses(subjectDescriptor!!, subjectDescriptor, memberClassDescriptors)
|
|
||||||
// When on a sealed class without derived members is considered non-exhaustive (see test WhenOnEmptySealed)
|
// When on a sealed class without derived members is considered non-exhaustive (see test WhenOnEmptySealed)
|
||||||
return getMissingClassCases(expression, memberClassDescriptors, context) +
|
return getMissingClassCases(expression, memberClassDescriptors, context) +
|
||||||
WhenOnNullableExhaustivenessChecker.getMissingCases(expression, context, nullable)
|
WhenOnNullableExhaustivenessChecker.getMissingCases(expression, context, nullable)
|
||||||
@@ -244,6 +244,12 @@ private object WhenOnSealedExhaustivenessChecker : WhenOnClassExhaustivenessChec
|
|||||||
return DescriptorUtils.isSealedClass(TypeUtils.getClassDescriptor(subjectType))
|
return DescriptorUtils.isSealedClass(TypeUtils.getClassDescriptor(subjectType))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
internal fun getNestedSubclasses(baseDescriptor: ClassDescriptor): Set<ClassDescriptor> {
|
||||||
|
val memberClassDescriptors = LinkedHashSet<ClassDescriptor>()
|
||||||
|
collectNestedSubclasses(baseDescriptor, baseDescriptor, memberClassDescriptors)
|
||||||
|
return memberClassDescriptors
|
||||||
|
}
|
||||||
|
|
||||||
private fun collectNestedSubclasses(
|
private fun collectNestedSubclasses(
|
||||||
baseDescriptor: ClassDescriptor,
|
baseDescriptor: ClassDescriptor,
|
||||||
currentDescriptor: ClassDescriptor,
|
currentDescriptor: ClassDescriptor,
|
||||||
@@ -291,7 +297,7 @@ object WhenChecker {
|
|||||||
}
|
}
|
||||||
|
|
||||||
private fun whenSubjectType(expression: KtWhenExpression, context: BindingContext) =
|
private fun whenSubjectType(expression: KtWhenExpression, context: BindingContext) =
|
||||||
expression.subjectExpression?.let { context.getType(it) } ?: null
|
expression.subjectExpression?.let { context.get(SMARTCAST, it) ?: context.getType(it) } ?: null
|
||||||
|
|
||||||
@JvmStatic
|
@JvmStatic
|
||||||
fun getEnumMissingCases(
|
fun getEnumMissingCases(
|
||||||
@@ -300,30 +306,9 @@ object WhenChecker {
|
|||||||
enumClassDescriptor: ClassDescriptor
|
enumClassDescriptor: ClassDescriptor
|
||||||
) = WhenOnEnumExhaustivenessChecker.getMissingCases(expression, context, enumClassDescriptor, false)
|
) = WhenOnEnumExhaustivenessChecker.getMissingCases(expression, context, enumClassDescriptor, false)
|
||||||
|
|
||||||
/**
|
|
||||||
* It's assumed that function is called for a final type. In this case the only possible smart cast is to not nullable type.
|
|
||||||
* @return true if type is nullable, and cannot be smart casted
|
|
||||||
*/
|
|
||||||
private fun isNullableTypeWithoutPossibleSmartCast(
|
|
||||||
expression: KtExpression?,
|
|
||||||
type: KotlinType,
|
|
||||||
context: BindingContext
|
|
||||||
): Boolean {
|
|
||||||
if (expression == null) return false // Normally should not happen
|
|
||||||
if (!TypeUtils.isNullableType(type)) return false
|
|
||||||
// We cannot read data flow information here due to lack of inputs (module descriptor is necessary)
|
|
||||||
if (context.get(BindingContext.SMARTCAST, expression) != null) {
|
|
||||||
// We have smart cast from enum or boolean to something
|
|
||||||
// Not very nice but we *can* decide it was smart cast to not-null
|
|
||||||
// because both enum and boolean are final
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
fun getMissingCases(expression: KtWhenExpression, context: BindingContext): List<WhenMissingCase> {
|
fun getMissingCases(expression: KtWhenExpression, context: BindingContext): List<WhenMissingCase> {
|
||||||
val type = whenSubjectType(expression, context) ?: return listOf(UnknownMissingCase)
|
val type = whenSubjectType(expression, context) ?: return listOf(UnknownMissingCase)
|
||||||
val nullable = !type.isFlexible() && isNullableTypeWithoutPossibleSmartCast(expression.subjectExpression, type, context)
|
val nullable = type.isMarkedNullable
|
||||||
val checkers = exhaustivenessCheckers.filter { it.isApplicable(type) }
|
val checkers = exhaustivenessCheckers.filter { it.isApplicable(type) }
|
||||||
if (checkers.isEmpty()) return listOf(UnknownMissingCase)
|
if (checkers.isEmpty()) return listOf(UnknownMissingCase)
|
||||||
return checkers.map { it.getMissingCases(expression, context, TypeUtils.getClassDescriptor(type), nullable) }.flatten()
|
return checkers.map { it.getMissingCases(expression, context, TypeUtils.getClassDescriptor(type), nullable) }.flatten()
|
||||||
|
|||||||
+37
-12
@@ -19,6 +19,9 @@ package org.jetbrains.kotlin.types.expressions
|
|||||||
import org.jetbrains.kotlin.builtins.KotlinBuiltIns
|
import org.jetbrains.kotlin.builtins.KotlinBuiltIns
|
||||||
import org.jetbrains.kotlin.builtins.KotlinBuiltIns.isBoolean
|
import org.jetbrains.kotlin.builtins.KotlinBuiltIns.isBoolean
|
||||||
import org.jetbrains.kotlin.cfg.WhenChecker
|
import org.jetbrains.kotlin.cfg.WhenChecker
|
||||||
|
import org.jetbrains.kotlin.cfg.WhenOnSealedExhaustivenessChecker
|
||||||
|
import org.jetbrains.kotlin.descriptors.ClassDescriptor
|
||||||
|
import org.jetbrains.kotlin.descriptors.Modality
|
||||||
import org.jetbrains.kotlin.diagnostics.Errors
|
import org.jetbrains.kotlin.diagnostics.Errors
|
||||||
import org.jetbrains.kotlin.diagnostics.Errors.*
|
import org.jetbrains.kotlin.diagnostics.Errors.*
|
||||||
import org.jetbrains.kotlin.psi.*
|
import org.jetbrains.kotlin.psi.*
|
||||||
@@ -206,22 +209,44 @@ class PatternMatchingTypingVisitor internal constructor(facade: ExpressionTyping
|
|||||||
return Pair(currentDataFlowInfo ?: contextAfterSubject.dataFlowInfo, jumpOutPossible)
|
return Pair(currentDataFlowInfo ?: contextAfterSubject.dataFlowInfo, jumpOutPossible)
|
||||||
}
|
}
|
||||||
|
|
||||||
private fun checkSmartCastsInSubjectIfRequired(expression: KtWhenExpression, contextBeforeSubject: ExpressionTypingContext, subjectType: KotlinType) {
|
private fun checkSmartCastsInSubjectIfRequired(
|
||||||
val subjectExpression = expression.subjectExpression
|
expression: KtWhenExpression,
|
||||||
if (subjectExpression != null &&
|
contextBeforeSubject: ExpressionTypingContext,
|
||||||
TypeUtils.isNullableType(subjectType) &&
|
subjectType: KotlinType
|
||||||
!WhenChecker.containsNullCase(expression, contextBeforeSubject.trace.bindingContext)
|
) {
|
||||||
) {
|
val subjectExpression = expression.subjectExpression ?: return
|
||||||
val trace = TemporaryBindingTrace.create(contextBeforeSubject.trace, "Temporary trace for when subject nullability")
|
val nullableType = TypeUtils.isNullableType(subjectType)
|
||||||
val subjectContext = contextBeforeSubject.replaceExpectedType(TypeUtils.makeNotNullable(subjectType)).replaceBindingTrace(trace)
|
val bindingContext = contextBeforeSubject.trace.bindingContext
|
||||||
val castResult = DataFlowAnalyzer.checkPossibleCast(
|
if (nullableType && !WhenChecker.containsNullCase(expression, bindingContext)) {
|
||||||
subjectType, KtPsiUtil.safeDeparenthesize(subjectExpression), subjectContext)
|
val notNullableType = TypeUtils.makeNotNullable(subjectType)
|
||||||
if (castResult != null && castResult.isCorrect) {
|
checkSmartCastToExpectedTypeInSubject(contextBeforeSubject, subjectExpression, subjectType, notNullableType)
|
||||||
trace.commit()
|
}
|
||||||
|
val subjectClass = subjectType.constructor.declarationDescriptor as? ClassDescriptor ?: return
|
||||||
|
if (subjectClass.modality == Modality.SEALED &&
|
||||||
|
WhenOnSealedExhaustivenessChecker.getMissingCases(expression, bindingContext, subjectClass, false).isNotEmpty()) {
|
||||||
|
for (descriptor in WhenOnSealedExhaustivenessChecker.getNestedSubclasses(subjectClass)) {
|
||||||
|
if (descriptor.modality == Modality.SEALED && DescriptorUtils.isDirectSubclass(descriptor, subjectClass)) {
|
||||||
|
checkSmartCastToExpectedTypeInSubject(contextBeforeSubject, subjectExpression, subjectType, descriptor.defaultType)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
private fun checkSmartCastToExpectedTypeInSubject(
|
||||||
|
contextBeforeSubject: ExpressionTypingContext,
|
||||||
|
subjectExpression: KtExpression,
|
||||||
|
subjectType: KotlinType,
|
||||||
|
expectedType: KotlinType
|
||||||
|
) {
|
||||||
|
val trace = TemporaryBindingTrace.create(contextBeforeSubject.trace, "Temporary trace for when subject nullability")
|
||||||
|
val subjectContext = contextBeforeSubject.replaceExpectedType(expectedType).replaceBindingTrace(trace)
|
||||||
|
val castResult = DataFlowAnalyzer.checkPossibleCast(
|
||||||
|
subjectType, KtPsiUtil.safeDeparenthesize(subjectExpression), subjectContext)
|
||||||
|
if (castResult != null && castResult.isCorrect) {
|
||||||
|
trace.commit()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
private fun analyzeWhenEntryConditions(
|
private fun analyzeWhenEntryConditions(
|
||||||
whenEntry: KtWhenEntry,
|
whenEntry: KtWhenEntry,
|
||||||
context: ExpressionTypingContext,
|
context: ExpressionTypingContext,
|
||||||
|
|||||||
@@ -0,0 +1,45 @@
|
|||||||
|
// See KT-10648: Exhaustiveness check does not work with nested sealed hierarchy
|
||||||
|
sealed class Base {
|
||||||
|
sealed class A : Base() {
|
||||||
|
class A1 : A()
|
||||||
|
class A2 : A()
|
||||||
|
}
|
||||||
|
sealed class B : Base() {
|
||||||
|
class B1 : B()
|
||||||
|
class B2 : B()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fun foo(b: Base) = when (b) {
|
||||||
|
is Base.A -> when(<!DEBUG_INFO_SMARTCAST!>b<!>) {
|
||||||
|
is Base.A.A1 -> 1
|
||||||
|
is Base.A.A2 -> 2
|
||||||
|
}
|
||||||
|
is Base.B -> when(<!DEBUG_INFO_SMARTCAST!>b<!>) {
|
||||||
|
is Base.B.B1 -> 3
|
||||||
|
is Base.B.B2 -> 4
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fun bar(b: Base?) = if (b == null) 0 else when (<!DEBUG_INFO_SMARTCAST!>b<!>) {
|
||||||
|
is Base.A -> when(<!DEBUG_INFO_SMARTCAST!>b<!>) {
|
||||||
|
is Base.A.A1 -> 1
|
||||||
|
is Base.A.A2 -> 2
|
||||||
|
}
|
||||||
|
is Base.B -> when(<!DEBUG_INFO_SMARTCAST!>b<!>) {
|
||||||
|
is Base.B.B1 -> 3
|
||||||
|
is Base.B.B2 -> 4
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fun gav(b: Base?) = when (b) {
|
||||||
|
null -> 0
|
||||||
|
is Base.A -> when(<!DEBUG_INFO_SMARTCAST!>b<!>) {
|
||||||
|
is Base.A.A1 -> 1
|
||||||
|
is Base.A.A2 -> 2
|
||||||
|
}
|
||||||
|
is Base.B -> when(<!DEBUG_INFO_SMARTCAST!>b<!>) {
|
||||||
|
is Base.B.B1 -> 3
|
||||||
|
is Base.B.B2 -> 4
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,54 @@
|
|||||||
|
package
|
||||||
|
|
||||||
|
public fun bar(/*0*/ b: Base?): kotlin.Int
|
||||||
|
public fun foo(/*0*/ b: Base): kotlin.Int
|
||||||
|
public fun gav(/*0*/ b: Base?): kotlin.Int
|
||||||
|
|
||||||
|
public sealed class Base {
|
||||||
|
private constructor Base()
|
||||||
|
public open override /*1*/ /*fake_override*/ fun equals(/*0*/ other: kotlin.Any?): kotlin.Boolean
|
||||||
|
public open override /*1*/ /*fake_override*/ fun hashCode(): kotlin.Int
|
||||||
|
public open override /*1*/ /*fake_override*/ fun toString(): kotlin.String
|
||||||
|
|
||||||
|
public sealed class A : Base {
|
||||||
|
private constructor A()
|
||||||
|
public open override /*1*/ /*fake_override*/ fun equals(/*0*/ other: kotlin.Any?): kotlin.Boolean
|
||||||
|
public open override /*1*/ /*fake_override*/ fun hashCode(): kotlin.Int
|
||||||
|
public open override /*1*/ /*fake_override*/ fun toString(): kotlin.String
|
||||||
|
|
||||||
|
public final class A1 : Base.A {
|
||||||
|
public constructor A1()
|
||||||
|
public open override /*1*/ /*fake_override*/ fun equals(/*0*/ other: kotlin.Any?): kotlin.Boolean
|
||||||
|
public open override /*1*/ /*fake_override*/ fun hashCode(): kotlin.Int
|
||||||
|
public open override /*1*/ /*fake_override*/ fun toString(): kotlin.String
|
||||||
|
}
|
||||||
|
|
||||||
|
public final class A2 : Base.A {
|
||||||
|
public constructor A2()
|
||||||
|
public open override /*1*/ /*fake_override*/ fun equals(/*0*/ other: kotlin.Any?): kotlin.Boolean
|
||||||
|
public open override /*1*/ /*fake_override*/ fun hashCode(): kotlin.Int
|
||||||
|
public open override /*1*/ /*fake_override*/ fun toString(): kotlin.String
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
public sealed class B : Base {
|
||||||
|
private constructor B()
|
||||||
|
public open override /*1*/ /*fake_override*/ fun equals(/*0*/ other: kotlin.Any?): kotlin.Boolean
|
||||||
|
public open override /*1*/ /*fake_override*/ fun hashCode(): kotlin.Int
|
||||||
|
public open override /*1*/ /*fake_override*/ fun toString(): kotlin.String
|
||||||
|
|
||||||
|
public final class B1 : Base.B {
|
||||||
|
public constructor B1()
|
||||||
|
public open override /*1*/ /*fake_override*/ fun equals(/*0*/ other: kotlin.Any?): kotlin.Boolean
|
||||||
|
public open override /*1*/ /*fake_override*/ fun hashCode(): kotlin.Int
|
||||||
|
public open override /*1*/ /*fake_override*/ fun toString(): kotlin.String
|
||||||
|
}
|
||||||
|
|
||||||
|
public final class B2 : Base.B {
|
||||||
|
public constructor B2()
|
||||||
|
public open override /*1*/ /*fake_override*/ fun equals(/*0*/ other: kotlin.Any?): kotlin.Boolean
|
||||||
|
public open override /*1*/ /*fake_override*/ fun hashCode(): kotlin.Int
|
||||||
|
public open override /*1*/ /*fake_override*/ fun toString(): kotlin.String
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -16143,6 +16143,12 @@ public class DiagnosticsTestGenerated extends AbstractDiagnosticsTest {
|
|||||||
doTest(fileName);
|
doTest(fileName);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
@TestMetadata("NestedSealed.kt")
|
||||||
|
public void testNestedSealed() throws Exception {
|
||||||
|
String fileName = KotlinTestUtils.navigationMetadata("compiler/testData/diagnostics/tests/sealed/NestedSealed.kt");
|
||||||
|
doTest(fileName);
|
||||||
|
}
|
||||||
|
|
||||||
@TestMetadata("NeverConstructed.kt")
|
@TestMetadata("NeverConstructed.kt")
|
||||||
public void testNeverConstructed() throws Exception {
|
public void testNeverConstructed() throws Exception {
|
||||||
String fileName = KotlinTestUtils.navigationMetadata("compiler/testData/diagnostics/tests/sealed/NeverConstructed.kt");
|
String fileName = KotlinTestUtils.navigationMetadata("compiler/testData/diagnostics/tests/sealed/NeverConstructed.kt");
|
||||||
|
|||||||
Reference in New Issue
Block a user