Implement special desugaring for numeric comparisons in PSI2IR
This introduces the following IR built-in functions required for proper
implementation of the number comparisons:
- ieee754Equals(T, T): Boolean,
for each T in {Float?, Double?}
- less(T, T): Boolean
lessOrEqual(T, T): Boolean
greater(T, T): Boolean
greaterOrEqual(T, T): Boolean
for each T in {Int, Long, Float, Double}
This commit is contained in:
+118
-28
@@ -17,20 +17,28 @@
|
||||
package org.jetbrains.kotlin.psi2ir.generators
|
||||
|
||||
import org.jetbrains.kotlin.builtins.KotlinBuiltIns
|
||||
import org.jetbrains.kotlin.descriptors.FunctionDescriptor
|
||||
import org.jetbrains.kotlin.ir.IrStatement
|
||||
import org.jetbrains.kotlin.ir.builders.*
|
||||
import org.jetbrains.kotlin.ir.expressions.IrExpression
|
||||
import org.jetbrains.kotlin.ir.expressions.IrStatementOrigin
|
||||
import org.jetbrains.kotlin.ir.expressions.IrTypeOperator
|
||||
import org.jetbrains.kotlin.ir.expressions.impl.IrBinaryPrimitiveImpl
|
||||
import org.jetbrains.kotlin.ir.expressions.impl.IrTypeOperatorCallImpl
|
||||
import org.jetbrains.kotlin.ir.expressions.impl.IrUnaryPrimitiveImpl
|
||||
import org.jetbrains.kotlin.ir.expressions.impl.*
|
||||
import org.jetbrains.kotlin.ir.symbols.IrSimpleFunctionSymbol
|
||||
import org.jetbrains.kotlin.lexer.KtTokens
|
||||
import org.jetbrains.kotlin.name.Name
|
||||
import org.jetbrains.kotlin.psi.*
|
||||
import org.jetbrains.kotlin.psi.psiUtil.endOffset
|
||||
import org.jetbrains.kotlin.psi.psiUtil.startOffset
|
||||
import org.jetbrains.kotlin.psi2ir.findSingleFunction
|
||||
import org.jetbrains.kotlin.psi2ir.intermediate.CallReceiver
|
||||
import org.jetbrains.kotlin.psi2ir.intermediate.OnceExpressionValue
|
||||
import org.jetbrains.kotlin.psi2ir.intermediate.SimpleCallReceiver
|
||||
import org.jetbrains.kotlin.resolve.BindingContext
|
||||
import org.jetbrains.kotlin.resolve.calls.model.ResolvedCall
|
||||
import org.jetbrains.kotlin.resolve.constants.evaluate.ConstantExpressionEvaluator
|
||||
import org.jetbrains.kotlin.types.KotlinType
|
||||
import org.jetbrains.kotlin.types.typeUtil.isPrimitiveNumberType
|
||||
import org.jetbrains.kotlin.types.typeUtil.makeNotNullable
|
||||
import org.jetbrains.kotlin.types.typeUtil.makeNullable
|
||||
import java.lang.AssertionError
|
||||
@@ -150,8 +158,7 @@ class OperatorExpressionGenerator(statementGenerator: StatementGenerator) : Stat
|
||||
private fun generateInOperator(expression: KtBinaryExpression, irOperator: IrStatementOrigin): IrExpression {
|
||||
val containsCall = getResolvedCall(expression)!!
|
||||
|
||||
val irContainsCall =
|
||||
CallGenerator(statementGenerator).generateCall(expression, statementGenerator.pregenerateCall(containsCall), irOperator)
|
||||
val irContainsCall = generateCall(containsCall, expression, irOperator)
|
||||
|
||||
return when (irOperator) {
|
||||
IrStatementOrigin.IN ->
|
||||
@@ -172,7 +179,6 @@ class OperatorExpressionGenerator(statementGenerator: StatementGenerator) : Stat
|
||||
val irArgument0 = statementGenerator.generateExpression(expression.left!!)
|
||||
val irArgument1 = statementGenerator.generateExpression(expression.right!!)
|
||||
|
||||
|
||||
val irIdentityEquals = IrBinaryPrimitiveImpl(
|
||||
expression.startOffset, expression.endOffset, irOperator,
|
||||
context.irBuiltIns.eqeqeqSymbol,
|
||||
@@ -191,18 +197,33 @@ class OperatorExpressionGenerator(statementGenerator: StatementGenerator) : Stat
|
||||
else ->
|
||||
throw AssertionError("Unexpected identity operator $irOperator")
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
private fun KtExpression.generateAsPrimitiveNumericComparisonOperand(primitiveNumericComparisonType: KotlinType?) =
|
||||
statementGenerator.generateExpression(this)
|
||||
.promoteToPrimitiveNumericType(
|
||||
getPrimitiveNumericComparisonOperandType(this),
|
||||
primitiveNumericComparisonType
|
||||
)
|
||||
|
||||
private fun getPrimitiveNumericComparisonType(ktExpression: KtBinaryExpression) =
|
||||
context.bindingContext[BindingContext.PRIMITIVE_NUMERIC_COMPARISON_TYPE, ktExpression]
|
||||
|
||||
private fun getPrimitiveNumericComparisonOperandType(ktExpression: KtExpression) =
|
||||
context.bindingContext[BindingContext.PRIMITIVE_NUMERIC_COMPARISON_OPERAND_TYPE, ktExpression]
|
||||
|
||||
private fun generateEqualityOperator(expression: KtBinaryExpression, irOperator: IrStatementOrigin): IrExpression {
|
||||
val irArgument0 = statementGenerator.generateExpression(expression.left!!)
|
||||
val irArgument1 = statementGenerator.generateExpression(expression.right!!)
|
||||
val primitiveNumericComparisonType = getPrimitiveNumericComparisonType(expression)
|
||||
|
||||
val eqeqSymbol = context.irBuiltIns.ieee754equalsFunByOperandType[primitiveNumericComparisonType]?.symbol
|
||||
?: context.irBuiltIns.eqeqSymbol
|
||||
|
||||
val irEquals = IrBinaryPrimitiveImpl(
|
||||
expression.startOffset, expression.endOffset,
|
||||
irOperator,
|
||||
context.irBuiltIns.eqeqSymbol,
|
||||
irArgument0, irArgument1
|
||||
eqeqSymbol,
|
||||
expression.left!!.generateAsPrimitiveNumericComparisonOperand(primitiveNumericComparisonType),
|
||||
expression.right!!.generateAsPrimitiveNumericComparisonOperand(primitiveNumericComparisonType)
|
||||
)
|
||||
|
||||
return when (irOperator) {
|
||||
@@ -210,33 +231,104 @@ class OperatorExpressionGenerator(statementGenerator: StatementGenerator) : Stat
|
||||
irEquals
|
||||
IrStatementOrigin.EXCLEQ ->
|
||||
IrUnaryPrimitiveImpl(
|
||||
expression.startOffset, expression.endOffset, IrStatementOrigin.EXCLEQ,
|
||||
expression.startOffset, expression.endOffset,
|
||||
IrStatementOrigin.EXCLEQ,
|
||||
context.irBuiltIns.booleanNotSymbol,
|
||||
irEquals
|
||||
)
|
||||
else ->
|
||||
throw AssertionError("Unexpected equality operator $irOperator")
|
||||
}
|
||||
}
|
||||
|
||||
private fun IrExpression.promoteToPrimitiveNumericType(operandType: KotlinType?, targetType: KotlinType?): IrExpression {
|
||||
if (targetType == null) return this
|
||||
if (operandType == null) throw AssertionError("operandType should be non-null")
|
||||
|
||||
val operandNNType = operandType.makeNotNullable()
|
||||
|
||||
val conversionFunction = operandNNType.findConversionFunctionTo(targetType)
|
||||
|
||||
return when {
|
||||
!operandNNType.isPrimitiveNumberType() ->
|
||||
throw AssertionError("Primitive number type or nullable primitive number type expected: $type")
|
||||
|
||||
operandType == targetType || operandNNType == targetType ->
|
||||
this
|
||||
|
||||
else ->
|
||||
SimpleCallReceiver(OnceExpressionValue(this), null)
|
||||
.invokeConversionFunction(
|
||||
startOffset, endOffset,
|
||||
conversionFunction ?: throw AssertionError("No conversion function for $type ~> $targetType")
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
private fun CallReceiver.invokeConversionFunction(
|
||||
startOffset: Int,
|
||||
endOffset: Int,
|
||||
functionDescriptor: FunctionDescriptor
|
||||
): IrExpression =
|
||||
call { dispatchReceiverValue, _ ->
|
||||
IrCallImpl(
|
||||
startOffset,
|
||||
endOffset,
|
||||
functionDescriptor.returnType!!,
|
||||
context.symbolTable.referenceFunction(functionDescriptor.original),
|
||||
functionDescriptor,
|
||||
typeArguments = null,
|
||||
origin = null, // TODO origin for widening conversions?
|
||||
superQualifierSymbol = null
|
||||
).apply {
|
||||
dispatchReceiver = dispatchReceiverValue!!.load()
|
||||
}
|
||||
}
|
||||
|
||||
private fun KotlinType.findConversionFunctionTo(targetType: KotlinType): FunctionDescriptor? {
|
||||
val targetTypeName = targetType.constructor.declarationDescriptor?.name?.asString() ?: return null
|
||||
return memberScope.findSingleFunction(Name.identifier("to$targetTypeName"))
|
||||
}
|
||||
|
||||
private fun generateComparisonOperator(expression: KtBinaryExpression, origin: IrStatementOrigin): IrExpression {
|
||||
val compareToCall = getResolvedCall(expression)!!
|
||||
val startOffset = expression.startOffset
|
||||
val endOffset = expression.endOffset
|
||||
|
||||
val irCompareToCall =
|
||||
CallGenerator(statementGenerator).generateCall(expression, statementGenerator.pregenerateCall(compareToCall), origin)
|
||||
val primitiveNumberComparisonType = getPrimitiveNumericComparisonType(expression)
|
||||
|
||||
val compareToZeroSymbol = when (origin) {
|
||||
IrStatementOrigin.LT -> context.irBuiltIns.lt0Symbol
|
||||
IrStatementOrigin.LTEQ -> context.irBuiltIns.lteq0Symbol
|
||||
IrStatementOrigin.GT -> context.irBuiltIns.gt0Symbol
|
||||
IrStatementOrigin.GTEQ -> context.irBuiltIns.gteq0Symbol
|
||||
else -> throw AssertionError("Unexpected comparison operator: $origin")
|
||||
return if (primitiveNumberComparisonType != null) {
|
||||
IrBinaryPrimitiveImpl(
|
||||
startOffset, endOffset, origin,
|
||||
getComparisonOperatorSymbol(origin, primitiveNumberComparisonType),
|
||||
expression.left!!.generateAsPrimitiveNumericComparisonOperand(primitiveNumberComparisonType),
|
||||
expression.right!!.generateAsPrimitiveNumericComparisonOperand(primitiveNumberComparisonType)
|
||||
)
|
||||
} else {
|
||||
IrBinaryPrimitiveImpl(
|
||||
startOffset, endOffset, origin,
|
||||
getComparisonOperatorSymbol(origin, context.irBuiltIns.int),
|
||||
generateCall(getResolvedCall(expression)!!, expression, origin),
|
||||
IrConstImpl.int(startOffset, endOffset, context.builtIns.intType, 0)
|
||||
)
|
||||
}
|
||||
|
||||
return IrUnaryPrimitiveImpl(expression.startOffset, expression.endOffset, origin, compareToZeroSymbol, irCompareToCall)
|
||||
}
|
||||
|
||||
private fun generateCall(
|
||||
resolvedCall: ResolvedCall<*>,
|
||||
ktExpression: KtExpression,
|
||||
origin: IrStatementOrigin?
|
||||
) =
|
||||
CallGenerator(statementGenerator).generateCall(ktExpression, statementGenerator.pregenerateCall(resolvedCall), origin)
|
||||
|
||||
private fun getComparisonOperatorSymbol(origin: IrStatementOrigin, primitiveNumericType: KotlinType): IrSimpleFunctionSymbol =
|
||||
when (origin) {
|
||||
IrStatementOrigin.LT -> context.irBuiltIns.lessFunByOperandType
|
||||
IrStatementOrigin.LTEQ -> context.irBuiltIns.lessOrEqualFunByOperandType
|
||||
IrStatementOrigin.GT -> context.irBuiltIns.greaterFunByOperandType
|
||||
IrStatementOrigin.GTEQ -> context.irBuiltIns.greaterOrEqualFunByOperandType
|
||||
else -> throw AssertionError("Unexpected comparison operator: $origin")
|
||||
}[primitiveNumericType]!!.symbol
|
||||
|
||||
private fun generateExclExclOperator(expression: KtPostfixExpression, origin: IrStatementOrigin): IrExpression {
|
||||
val ktArgument = expression.baseExpression!!
|
||||
val irArgument = statementGenerator.generateExpression(ktArgument)
|
||||
@@ -250,10 +342,8 @@ class OperatorExpressionGenerator(statementGenerator: StatementGenerator) : Stat
|
||||
}
|
||||
}
|
||||
|
||||
private fun generateBinaryOperatorAsCall(expression: KtBinaryExpression, origin: IrStatementOrigin?): IrExpression {
|
||||
val operatorCall = getResolvedCall(expression)!!
|
||||
return CallGenerator(statementGenerator).generateCall(expression, statementGenerator.pregenerateCall(operatorCall), origin)
|
||||
}
|
||||
private fun generateBinaryOperatorAsCall(expression: KtBinaryExpression, origin: IrStatementOrigin?): IrExpression =
|
||||
generateCall(getResolvedCall(expression)!!, expression, origin)
|
||||
|
||||
private fun generatePrefixOperatorAsCall(expression: KtPrefixExpression, origin: IrStatementOrigin): IrExpression {
|
||||
val resolvedCall = getResolvedCall(expression)!!
|
||||
@@ -267,6 +357,6 @@ class OperatorExpressionGenerator(statementGenerator: StatementGenerator) : Stat
|
||||
}
|
||||
}
|
||||
|
||||
return CallGenerator(statementGenerator).generateCall(expression, statementGenerator.pregenerateCall(resolvedCall), origin)
|
||||
return generateCall(resolvedCall, expression, origin)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -26,7 +26,7 @@ import org.jetbrains.kotlin.resolve.calls.model.ResolvedCall
|
||||
import org.jetbrains.kotlin.types.KotlinType
|
||||
|
||||
class CallBuilder(
|
||||
val original: ResolvedCall<*>,
|
||||
val original: ResolvedCall<*>, // TODO get rid of "original", sometimes we want to generate a call without ResolvedCall
|
||||
val descriptor: CallableDescriptor,
|
||||
val isExtensionInvokeCall: Boolean = false
|
||||
) {
|
||||
|
||||
@@ -32,7 +32,9 @@ import org.jetbrains.kotlin.name.FqName
|
||||
import org.jetbrains.kotlin.name.Name
|
||||
import org.jetbrains.kotlin.types.KotlinType
|
||||
import org.jetbrains.kotlin.types.KotlinTypeFactory
|
||||
import org.jetbrains.kotlin.types.SimpleType
|
||||
import org.jetbrains.kotlin.types.Variance
|
||||
import org.jetbrains.kotlin.types.typeUtil.makeNullable
|
||||
|
||||
class IrBuiltIns(val builtIns: KotlinBuiltIns) {
|
||||
private val packageFragment = IrBuiltinsPackageFragmentDescriptorImpl(builtIns.builtInsModule, KOTLIN_INTERNAL_IR_FQN)
|
||||
@@ -59,40 +61,54 @@ class IrBuiltIns(val builtIns: KotlinBuiltIns) {
|
||||
private fun <T : SimpleFunctionDescriptor> T.addStub(): IrSimpleFunction =
|
||||
addStubToPackageFragment(this)
|
||||
|
||||
private fun defineComparisonOperator(name: String, operandType: KotlinType) =
|
||||
defineOperator(name, bool, listOf(operandType, operandType))
|
||||
|
||||
private fun List<SimpleType>.defineComparisonOperatorForEachType(name: String) =
|
||||
associate { it to defineComparisonOperator(name, it) }
|
||||
|
||||
val bool = builtIns.booleanType
|
||||
val any = builtIns.anyType
|
||||
val anyN = builtIns.nullableAnyType
|
||||
val char = builtIns.charType
|
||||
val byte = builtIns.byteType
|
||||
val short = builtIns.shortType
|
||||
val int = builtIns.intType
|
||||
val long = builtIns.longType
|
||||
val float = builtIns.floatType
|
||||
val double = builtIns.doubleType
|
||||
val nothing = builtIns.nothingType
|
||||
val unit = builtIns.unitType
|
||||
val string = builtIns.stringType
|
||||
|
||||
val primitiveTypes = listOf(bool, char, byte, short, int, long, float, double)
|
||||
val primitiveTypesWithComparisons = listOf(int, long, float, double)
|
||||
val primitiveFloatingPointTypes = listOf(float, double)
|
||||
|
||||
val lessFunByOperandType = primitiveTypesWithComparisons.defineComparisonOperatorForEachType("less")
|
||||
val lessOrEqualFunByOperandType = primitiveTypesWithComparisons.defineComparisonOperatorForEachType("lessOrEqual")
|
||||
val greaterOrEqualFunByOperandType = primitiveTypesWithComparisons.defineComparisonOperatorForEachType("greaterOrEqual")
|
||||
val greaterFunByOperandType = primitiveTypesWithComparisons.defineComparisonOperatorForEachType("greater")
|
||||
|
||||
val ieee754equalsFunByOperandType =
|
||||
primitiveFloatingPointTypes.associate {
|
||||
it to defineOperator("ieee754equals", bool, listOf(it.makeNullable(), it.makeNullable()))
|
||||
}
|
||||
|
||||
val eqeqeqFun = defineOperator("EQEQEQ", bool, listOf(anyN, anyN))
|
||||
val eqeqFun = defineOperator("EQEQ", bool, listOf(anyN, anyN))
|
||||
val lt0Fun = defineOperator("LT0", bool, listOf(int))
|
||||
val lteq0Fun = defineOperator("LTEQ0", bool, listOf(int))
|
||||
val gt0Fun = defineOperator("GT0", bool, listOf(int))
|
||||
val gteq0Fun = defineOperator("GTEQ0", bool, listOf(int))
|
||||
val throwNpeFun = defineOperator("THROW_NPE", nothing, listOf())
|
||||
val booleanNotFun = defineOperator("NOT", bool, listOf(bool))
|
||||
val noWhenBranchMatchedExceptionFun = defineOperator("noWhenBranchMatchedException", unit, listOf())
|
||||
|
||||
val eqeqeq = eqeqeqFun.descriptor
|
||||
val eqeq = eqeqFun.descriptor
|
||||
val lt0 = lt0Fun.descriptor
|
||||
val lteq0 = lteq0Fun.descriptor
|
||||
val gt0 = gt0Fun.descriptor
|
||||
val gteq0 = gteq0Fun.descriptor
|
||||
val throwNpe = throwNpeFun.descriptor
|
||||
val booleanNot = booleanNotFun.descriptor
|
||||
val noWhenBranchMatchedException = noWhenBranchMatchedExceptionFun.descriptor
|
||||
|
||||
val eqeqeqSymbol = eqeqeqFun.symbol
|
||||
val eqeqSymbol = eqeqFun.symbol
|
||||
val lt0Symbol = lt0Fun.symbol
|
||||
val lteq0Symbol = lteq0Fun.symbol
|
||||
val gt0Symbol = gt0Fun.symbol
|
||||
val gteq0Symbol = gteq0Fun.symbol
|
||||
val throwNpeSymbol = throwNpeFun.symbol
|
||||
val booleanNotSymbol = booleanNotFun.symbol
|
||||
val noWhenBranchMatchedExceptionSymbol = noWhenBranchMatchedExceptionFun.symbol
|
||||
|
||||
Reference in New Issue
Block a user