[FIR] Comparison to String implies smart cast to String
When the left-hand side of an equality comparison is known to be a String, and the equality condition resolves to true, then the right-hand side can be smart-cast to a String as well. This was working for String expressions on the left-hand side but not for String constants. ^KT-57513 Fixed
This commit is contained in:
+25
-2
@@ -34,6 +34,7 @@ import org.jetbrains.kotlin.fir.symbols.impl.FirPropertySymbol
|
|||||||
import org.jetbrains.kotlin.fir.symbols.impl.FirRegularClassSymbol
|
import org.jetbrains.kotlin.fir.symbols.impl.FirRegularClassSymbol
|
||||||
import org.jetbrains.kotlin.fir.symbols.impl.FirTypeParameterSymbol
|
import org.jetbrains.kotlin.fir.symbols.impl.FirTypeParameterSymbol
|
||||||
import org.jetbrains.kotlin.fir.types.*
|
import org.jetbrains.kotlin.fir.types.*
|
||||||
|
import org.jetbrains.kotlin.name.ClassId
|
||||||
import org.jetbrains.kotlin.name.StandardClassIds
|
import org.jetbrains.kotlin.name.StandardClassIds
|
||||||
import org.jetbrains.kotlin.types.ConstantValueKind
|
import org.jetbrains.kotlin.types.ConstantValueKind
|
||||||
import org.jetbrains.kotlin.util.OperatorNameConventions
|
import org.jetbrains.kotlin.util.OperatorNameConventions
|
||||||
@@ -439,7 +440,7 @@ abstract class FirDataFlowAnalyzer(
|
|||||||
|
|
||||||
private fun processEqConst(
|
private fun processEqConst(
|
||||||
flow: MutableFlow,
|
flow: MutableFlow,
|
||||||
expression: FirExpression,
|
expression: FirEqualityOperatorCall,
|
||||||
operand: FirExpression,
|
operand: FirExpression,
|
||||||
const: FirConstExpression<*>,
|
const: FirConstExpression<*>,
|
||||||
isEq: Boolean
|
isEq: Boolean
|
||||||
@@ -450,6 +451,7 @@ abstract class FirDataFlowAnalyzer(
|
|||||||
|
|
||||||
val operandVariable = variableStorage.getOrCreateIfReal(flow, operand) ?: return
|
val operandVariable = variableStorage.getOrCreateIfReal(flow, operand) ?: return
|
||||||
val expressionVariable = variableStorage.createSynthetic(expression)
|
val expressionVariable = variableStorage.createSynthetic(expression)
|
||||||
|
|
||||||
if (const.kind == ConstantValueKind.Boolean && operand.resolvedType.isBooleanOrNullableBoolean) {
|
if (const.kind == ConstantValueKind.Boolean && operand.resolvedType.isBooleanOrNullableBoolean) {
|
||||||
val expected = (const.value as Boolean)
|
val expected = (const.value as Boolean)
|
||||||
flow.addImplication((expressionVariable eq isEq) implies (operandVariable eq expected))
|
flow.addImplication((expressionVariable eq isEq) implies (operandVariable eq expected))
|
||||||
@@ -460,6 +462,11 @@ abstract class FirDataFlowAnalyzer(
|
|||||||
// expression == non-null const -> expression != null
|
// expression == non-null const -> expression != null
|
||||||
flow.addImplication((expressionVariable eq isEq) implies (operandVariable notEq null))
|
flow.addImplication((expressionVariable eq isEq) implies (operandVariable notEq null))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// We can imply type information if the constant is the left operand and is a supported primitive type.
|
||||||
|
if (operandVariable is RealVariable && const == expression.arguments[0] && isSmartcastPrimitive(const.resolvedType.classId)) {
|
||||||
|
flow.addImplication((expressionVariable eq isEq) implies (operandVariable typeEq const.resolvedType))
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
private fun processEqNull(flow: MutableFlow, expression: FirExpression, operand: FirExpression, isEq: Boolean) {
|
private fun processEqNull(flow: MutableFlow, expression: FirExpression, operand: FirExpression, isEq: Boolean) {
|
||||||
@@ -542,8 +549,9 @@ abstract class FirDataFlowAnalyzer(
|
|||||||
val status = resolvedStatus
|
val status = resolvedStatus
|
||||||
if (checkModality && status.modality != Modality.FINAL) return true
|
if (checkModality && status.modality != Modality.FINAL) return true
|
||||||
if (status.isExpect) return true
|
if (status.isExpect) return true
|
||||||
|
if (isSmartcastPrimitive(classId)) return false
|
||||||
when (classId) {
|
when (classId) {
|
||||||
StandardClassIds.Any, StandardClassIds.String -> return false
|
StandardClassIds.Any -> return false
|
||||||
// Float and Double effectively had non-trivial `equals` semantics while they don't have explicit overrides (see KT-50535)
|
// Float and Double effectively had non-trivial `equals` semantics while they don't have explicit overrides (see KT-50535)
|
||||||
StandardClassIds.Float, StandardClassIds.Double -> return true
|
StandardClassIds.Float, StandardClassIds.Double -> return true
|
||||||
}
|
}
|
||||||
@@ -563,6 +571,21 @@ abstract class FirDataFlowAnalyzer(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Determines if type smart-casting to the specified [ClassId] can be performed when values are
|
||||||
|
* compared via equality. Because this is determined using the ClassId, only standard built-in
|
||||||
|
* types are considered.
|
||||||
|
*/
|
||||||
|
private fun isSmartcastPrimitive(classId: ClassId?): Boolean {
|
||||||
|
return when (classId) {
|
||||||
|
// Support other primitives as well: KT-62246.
|
||||||
|
StandardClassIds.String,
|
||||||
|
-> true
|
||||||
|
|
||||||
|
else -> false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// ----------------------------------- Jump -----------------------------------
|
// ----------------------------------- Jump -----------------------------------
|
||||||
|
|
||||||
fun enterJump(jump: FirJump<*>) {
|
fun enterJump(jump: FirJump<*>) {
|
||||||
|
|||||||
@@ -8,8 +8,8 @@ import checkSubtype
|
|||||||
fun main(args : Array<String>) {
|
fun main(args : Array<String>) {
|
||||||
val x = checkSubtype<Any>(args[0])
|
val x = checkSubtype<Any>(args[0])
|
||||||
if(x is <!PLATFORM_CLASS_MAPPED_TO_KOTLIN!>java.lang.CharSequence<!>) {
|
if(x is <!PLATFORM_CLASS_MAPPED_TO_KOTLIN!>java.lang.CharSequence<!>) {
|
||||||
if (<!EQUALITY_NOT_APPLICABLE_WARNING!>"a" == x<!>) x.<!FUNCTION_CALL_EXPECTED!>length<!> else x.length() // OK
|
if (<!EQUALITY_NOT_APPLICABLE_WARNING!>"a" == x<!>) x.length else x.length() // OK
|
||||||
if (<!EQUALITY_NOT_APPLICABLE_WARNING!>"a" == x<!> || <!EQUALITY_NOT_APPLICABLE_WARNING!>"b" == x<!>) x.<!FUNCTION_CALL_EXPECTED!>length<!> else x.length() // <– THEN ERROR
|
if (<!EQUALITY_NOT_APPLICABLE_WARNING!>"a" == x<!> || <!EQUALITY_NOT_APPLICABLE_WARNING!>"b" == x<!>) x.length else x.length() // <– THEN ERROR
|
||||||
if (<!EQUALITY_NOT_APPLICABLE_WARNING!>"a" == x<!> && <!EQUALITY_NOT_APPLICABLE_WARNING!>"a" == x<!>) x.<!FUNCTION_CALL_EXPECTED!>length<!> else x.length() // <– ELSE ERROR
|
if (<!EQUALITY_NOT_APPLICABLE_WARNING!>"a" == x<!> && "a" == x) x.length else x.length() // <– ELSE ERROR
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -23,3 +23,31 @@ fun gav(i: TestWithEquals?, j: TestWithEquals?) {
|
|||||||
if (i == j) foo(i)
|
if (i == j) foo(i)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fun string(foo: Any) {
|
||||||
|
val string = ""
|
||||||
|
if ("" == foo) foo.length
|
||||||
|
if (string == foo) foo.length
|
||||||
|
if (foo == "") foo.<!UNRESOLVED_REFERENCE!>length<!>
|
||||||
|
}
|
||||||
|
|
||||||
|
fun int(foo: Any) {
|
||||||
|
val int = 1
|
||||||
|
if (1 == foo) foo.<!UNRESOLVED_REFERENCE!>plus<!>(1)
|
||||||
|
if (int == foo) foo.<!UNRESOLVED_REFERENCE!>plus<!>(1)
|
||||||
|
if (foo == 1) foo.<!UNRESOLVED_REFERENCE!>plus<!>(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
fun long(foo: Any) {
|
||||||
|
val long = 1L
|
||||||
|
if (1L == foo) foo.<!UNRESOLVED_REFERENCE!>plus<!>(1L)
|
||||||
|
if (long == foo) foo.<!UNRESOLVED_REFERENCE!>plus<!>(1L)
|
||||||
|
if (foo == 1L) foo.<!UNRESOLVED_REFERENCE!>plus<!>(1L)
|
||||||
|
}
|
||||||
|
|
||||||
|
fun char(foo: Any) {
|
||||||
|
val char = 'a'
|
||||||
|
if ('a' == foo) foo.<!UNRESOLVED_REFERENCE!>compareTo<!>('a')
|
||||||
|
if (char == foo) foo.<!UNRESOLVED_REFERENCE!>compareTo<!>('a')
|
||||||
|
if (foo == 'a') foo.<!UNRESOLVED_REFERENCE!>compareTo<!>('a')
|
||||||
|
}
|
||||||
|
|||||||
@@ -23,3 +23,31 @@ fun gav(i: TestWithEquals?, j: TestWithEquals?) {
|
|||||||
if (i == <!DEBUG_INFO_CONSTANT!>j<!>) foo(<!DEBUG_INFO_CONSTANT!>i<!>)
|
if (i == <!DEBUG_INFO_CONSTANT!>j<!>) foo(<!DEBUG_INFO_CONSTANT!>i<!>)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fun string(foo: Any) {
|
||||||
|
val string = ""
|
||||||
|
if ("" == foo) <!DEBUG_INFO_SMARTCAST!>foo<!>.length
|
||||||
|
if (string == foo) <!DEBUG_INFO_SMARTCAST!>foo<!>.length
|
||||||
|
if (foo == "") foo.<!UNRESOLVED_REFERENCE!>length<!>
|
||||||
|
}
|
||||||
|
|
||||||
|
fun int(foo: Any) {
|
||||||
|
val int = 1
|
||||||
|
if (1 == foo) foo.<!UNRESOLVED_REFERENCE_WRONG_RECEIVER!>plus<!>(1)
|
||||||
|
if (int == foo) foo.<!UNRESOLVED_REFERENCE_WRONG_RECEIVER!>plus<!>(1)
|
||||||
|
if (foo == 1) foo.<!UNRESOLVED_REFERENCE_WRONG_RECEIVER!>plus<!>(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
fun long(foo: Any) {
|
||||||
|
val long = 1L
|
||||||
|
if (1L == foo) foo.<!UNRESOLVED_REFERENCE_WRONG_RECEIVER!>plus<!>(1L)
|
||||||
|
if (long == foo) foo.<!UNRESOLVED_REFERENCE_WRONG_RECEIVER!>plus<!>(1L)
|
||||||
|
if (foo == 1L) foo.<!UNRESOLVED_REFERENCE_WRONG_RECEIVER!>plus<!>(1L)
|
||||||
|
}
|
||||||
|
|
||||||
|
fun char(foo: Any) {
|
||||||
|
val char = 'a'
|
||||||
|
if ('a' == foo) foo.<!UNRESOLVED_REFERENCE!>compareTo<!>('a')
|
||||||
|
if (char == foo) foo.<!UNRESOLVED_REFERENCE!>compareTo<!>('a')
|
||||||
|
if (foo == 'a') foo.<!UNRESOLVED_REFERENCE!>compareTo<!>('a')
|
||||||
|
}
|
||||||
|
|||||||
@@ -2,8 +2,12 @@ package
|
|||||||
|
|
||||||
public fun bar(/*0*/ i: Test?): kotlin.Unit
|
public fun bar(/*0*/ i: Test?): kotlin.Unit
|
||||||
public fun bar(/*0*/ i: TestWithEquals?): kotlin.Unit
|
public fun bar(/*0*/ i: TestWithEquals?): kotlin.Unit
|
||||||
|
public fun char(/*0*/ foo: kotlin.Any): kotlin.Unit
|
||||||
public fun foo(/*0*/ x: kotlin.String?): kotlin.String?
|
public fun foo(/*0*/ x: kotlin.String?): kotlin.String?
|
||||||
public fun gav(/*0*/ i: TestWithEquals?, /*1*/ j: TestWithEquals?): kotlin.Unit
|
public fun gav(/*0*/ i: TestWithEquals?, /*1*/ j: TestWithEquals?): kotlin.Unit
|
||||||
|
public fun int(/*0*/ foo: kotlin.Any): kotlin.Unit
|
||||||
|
public fun long(/*0*/ foo: kotlin.Any): kotlin.Unit
|
||||||
|
public fun string(/*0*/ foo: kotlin.Any): kotlin.Unit
|
||||||
|
|
||||||
public final class Test {
|
public final class Test {
|
||||||
public constructor Test()
|
public constructor Test()
|
||||||
|
|||||||
Reference in New Issue
Block a user