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:
Dmitry Petrov
2018-02-01 16:18:01 +03:00
parent f4ed4ec9d9
commit 9137e68d4e
33 changed files with 1566 additions and 233 deletions
@@ -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