Handle equality checks for 'when' and data classes

This commit is contained in:
Dmitry Petrov
2018-02-06 17:51:12 +03:00
parent 299eb24ca9
commit 00325ae539
19 changed files with 631 additions and 89 deletions
@@ -37,7 +37,7 @@ import org.jetbrains.kotlin.types.KotlinType
class AssignmentGenerator(statementGenerator: StatementGenerator) : StatementGeneratorExtension(statementGenerator) {
fun generateAssignment(expression: KtBinaryExpression): IrExpression {
val ktLeft = expression.left!!
val irRhs = statementGenerator.generateExpression(expression.right!!)
val irRhs = expression.right!!.genExpr()
val irAssignmentReceiver = generateAssignmentReceiver(ktLeft, IrStatementOrigin.EQ)
return irAssignmentReceiver.assign(irRhs)
}
@@ -52,7 +52,7 @@ class AssignmentGenerator(statementGenerator: StatementGenerator) : StatementGen
return irAssignmentReceiver.assign { irLValue ->
val opCall = statementGenerator.pregenerateCallReceivers(opResolvedCall)
opCall.setExplicitReceiverValue(irLValue)
opCall.irValueArgumentsByIndex[0] = statementGenerator.generateExpression(ktRight)
opCall.irValueArgumentsByIndex[0] = ktRight.genExpr()
val irOpCall = CallGenerator(statementGenerator).generateCall(expression, opCall, origin)
if (isSimpleAssignment) {
@@ -137,7 +137,7 @@ class AssignmentGenerator(statementGenerator: StatementGenerator) : StatementGen
origin
)
else ->
OnceExpressionValue(statementGenerator.generateExpression(ktLeft))
OnceExpressionValue(ktLeft.genExpr())
}
}
@@ -237,8 +237,8 @@ class AssignmentGenerator(statementGenerator: StatementGenerator) : StatementGen
ktLeft: KtArrayAccessExpression,
origin: IrStatementOrigin
): ArrayAccessAssignmentReceiver {
val irArray = statementGenerator.generateExpression(ktLeft.arrayExpression!!)
val irIndexExpressions = ktLeft.indexExpressions.map { statementGenerator.generateExpression(it) }
val irArray = ktLeft.arrayExpression!!.genExpr()
val irIndexExpressions = ktLeft.indexExpressions.map { it.genExpr() }
val indexedGetResolvedCall = get(BindingContext.INDEXED_LVALUE_GET, ktLeft)
val indexedGetCall = indexedGetResolvedCall?.let { statementGenerator.pregenerateCallReceivers(it) }
@@ -40,8 +40,8 @@ class BranchingExpressionGenerator(statementGenerator: StatementGenerator) : Sta
var irElseBranch: IrExpression? = null
whenBranches@ while (true) {
val irCondition = statementGenerator.generateExpression(ktLastIf.condition!!)
val irThenBranch = statementGenerator.generateExpression(ktLastIf.then!!)
val irCondition = ktLastIf.condition!!.genExpr()
val irThenBranch = ktLastIf.then!!.genExpr()
irBranches.add(IrBranchImpl(irCondition, irThenBranch))
val ktElse = ktLastIf.`else`?.deparenthesize()
@@ -49,7 +49,7 @@ class BranchingExpressionGenerator(statementGenerator: StatementGenerator) : Sta
null -> break@whenBranches
is KtIfExpression -> ktLastIf = ktElse
is KtExpression -> {
irElseBranch = statementGenerator.generateExpression(ktElse)
irElseBranch = ktElse.genExpr()
break@whenBranches
}
else -> throw AssertionError("Unexpected else expression: ${ktElse.text}")
@@ -85,7 +85,7 @@ class BranchingExpressionGenerator(statementGenerator: StatementGenerator) : Sta
fun generateWhenExpression(expression: KtWhenExpression): IrExpression {
val irSubject = expression.subjectExpression?.let {
scope.createTemporaryVariable(statementGenerator.generateExpression(it), "subject")
scope.createTemporaryVariable(it.genExpr(), "subject")
}
@@ -104,7 +104,7 @@ class BranchingExpressionGenerator(statementGenerator: StatementGenerator) : Sta
for (ktEntry in expression.entries) {
if (ktEntry.isElse) {
val irElseResult = statementGenerator.generateExpression(ktEntry.expression!!)
val irElseResult = ktEntry.expression!!.genExpr()
irWhen.branches.add(IrBranchImpl.elseBranch(irElseResult))
break
}
@@ -117,10 +117,9 @@ class BranchingExpressionGenerator(statementGenerator: StatementGenerator) : Sta
else
generateWhenConditionNoSubject(ktCondition)
irBranchCondition = irBranchCondition?.let { context.whenComma(it, irCondition) } ?: irCondition
}
val irBranchResult = statementGenerator.generateExpression(ktEntry.expression!!)
val irBranchResult = ktEntry.expression!!.genExpr()
irWhen.branches.add(IrBranchImpl(irBranchCondition!!, irBranchResult))
}
addElseBranchForExhaustiveWhenIfNeeded(irWhen, expression)
@@ -163,7 +162,7 @@ class BranchingExpressionGenerator(statementGenerator: StatementGenerator) : Sta
}
private fun generateWhenConditionNoSubject(ktCondition: KtWhenCondition): IrExpression =
statementGenerator.generateExpression((ktCondition as KtWhenConditionWithExpression).expression!!)
(ktCondition as KtWhenConditionWithExpression).expression!!.genExpr()
private fun generateWhenConditionWithSubject(ktCondition: KtWhenCondition, irSubject: IrVariable): IrExpression {
return when (ktCondition) {
@@ -204,10 +203,13 @@ class BranchingExpressionGenerator(statementGenerator: StatementGenerator) : Sta
}
}
private fun generateEqualsCondition(irSubject: IrVariable, ktCondition: KtWhenConditionWithExpression): IrBinaryPrimitiveImpl =
IrBinaryPrimitiveImpl(
ktCondition.startOffset, ktCondition.endOffset,
IrStatementOrigin.EQEQ, context.irBuiltIns.eqeqSymbol,
irSubject.defaultLoad(), statementGenerator.generateExpression(ktCondition.expression!!)
private fun generateEqualsCondition(irSubject: IrVariable, ktCondition: KtWhenConditionWithExpression): IrExpression {
val ktExpression = ktCondition.expression
val irExpression = ktExpression!!.genExpr()
return OperatorExpressionGenerator(statementGenerator).generateEquality(
ktCondition.startOffset, ktCondition.endOffset, IrStatementOrigin.EQEQ,
irSubject.defaultLoad(), irExpression,
context.bindingContext[BindingContext.PRIMITIVE_NUMERIC_COMPARISON_INFO, ktExpression]
)
}
}
@@ -137,12 +137,9 @@ class DataClassMembersGenerator(declarationGenerator: DeclarationGenerator) : De
+irIfThenReturnFalse(irNotIs(irOther(), classDescriptor.defaultType))
val otherWithCast = irTemporary(irAs(irOther(), classDescriptor.defaultType), "other_with_cast")
for (property in properties) {
+irIfThenReturnFalse(
irNotEquals(
irGet(irThis(), getPropertyGetterSymbol(property)),
irGet(irGet(otherWithCast.symbol), getPropertyGetterSymbol(property))
)
)
val arg1 = irGet(irThis(), getPropertyGetterSymbol(property))
val arg2 = irGet(irGet(otherWithCast.symbol), getPropertyGetterSymbol(property))
+irIfThenReturnFalse(irNotEquals(arg1, arg2))
}
+irReturnTrue()
}
@@ -229,7 +226,8 @@ class DataClassMembersGenerator(declarationGenerator: DeclarationGenerator) : De
val typeConstructorDescriptor = property.type.constructor.declarationDescriptor
val irPropertyStringValue =
if (typeConstructorDescriptor is ClassDescriptor &&
KotlinBuiltIns.isArrayOrPrimitiveArray(typeConstructorDescriptor))
KotlinBuiltIns.isArrayOrPrimitiveArray(typeConstructorDescriptor)
)
irCall(context.irBuiltIns.dataClassArrayMemberToStringSymbol).apply {
putValueArgument(0, irPropertyValue)
}
@@ -46,19 +46,19 @@ class ErrorExpressionGenerator(statementGenerator: StatementGenerator) : Stateme
val type = getErrorExpressionType(ktCall)
val irErrorCall = IrErrorCallExpressionImpl(ktCall.startOffset, ktCall.endOffset, type, "") // TODO problem description?
irErrorCall.explicitReceiver = (ktCall.parent as? KtDotQualifiedExpression)?.let {
statementGenerator.generateExpression(it.receiverExpression)
irErrorCall.explicitReceiver = (ktCall.parent as? KtDotQualifiedExpression)?.run {
receiverExpression.genExpr()
}
ktCall.valueArguments.forEach {
val ktArgument = it.getArgumentExpression()
if (ktArgument != null) {
irErrorCall.addArgument(statementGenerator.generateExpression(ktArgument))
irErrorCall.addArgument(ktArgument.genExpr())
}
}
ktCall.lambdaArguments.forEach {
irErrorCall.addArgument(statementGenerator.generateExpression(it.getArgumentExpression()))
irErrorCall.addArgument(it.getArgumentExpression().genExpr())
}
irErrorCall
@@ -73,7 +73,7 @@ class ErrorExpressionGenerator(statementGenerator: StatementGenerator) : Stateme
val irErrorCall = IrErrorCallExpressionImpl(ktName.startOffset, ktName.endOffset, type, "") // TODO problem description?
irErrorCall.explicitReceiver = (ktName.parent as? KtDotQualifiedExpression)?.let { ktParent ->
if (ktParent.receiverExpression == ktName) null
else statementGenerator.generateExpression(ktParent.receiverExpression)
else ktParent.receiverExpression.genExpr()
}
irErrorCall
@@ -36,7 +36,7 @@ class LoopExpressionGenerator(statementGenerator: StatementGenerator) : Statemen
context.builtIns.unitType, IrStatementOrigin.WHILE_LOOP
)
irLoop.condition = statementGenerator.generateExpression(ktWhile.condition!!)
irLoop.condition = ktWhile.condition!!.genExpr()
statementGenerator.bodyGenerator.putLoop(ktWhile, irLoop)
@@ -44,7 +44,7 @@ class LoopExpressionGenerator(statementGenerator: StatementGenerator) : Statemen
if (ktLoopBody is KtBlockExpression)
generateWhileLoopBody(ktLoopBody)
else
statementGenerator.generateExpression(ktLoopBody)
ktLoopBody.genExpr()
}
irLoop.label = getLoopLabel(ktWhile)
@@ -64,10 +64,10 @@ class LoopExpressionGenerator(statementGenerator: StatementGenerator) : Statemen
if (ktLoopBody is KtBlockExpression)
generateDoWhileLoopBody(ktLoopBody)
else
statementGenerator.generateExpression(ktLoopBody)
ktLoopBody.genExpr()
}
irLoop.condition = statementGenerator.generateExpression(ktDoWhile.condition!!)
irLoop.condition = ktDoWhile.condition!!.genExpr()
irLoop.label = getLoopLabel(ktDoWhile)
@@ -79,14 +79,14 @@ class LoopExpressionGenerator(statementGenerator: StatementGenerator) : Statemen
private fun generateWhileLoopBody(ktLoopBody: KtBlockExpression): IrExpression =
IrBlockImpl(
ktLoopBody.startOffset, ktLoopBody.endOffset, context.builtIns.unitType, null,
ktLoopBody.statements.map { statementGenerator.generateStatement(it) }
ktLoopBody.statements.map { it.genStmt() }
)
private fun generateDoWhileLoopBody(ktLoopBody: KtBlockExpression): IrExpression =
IrCompositeImpl(
ktLoopBody.startOffset, ktLoopBody.endOffset, context.builtIns.unitType, null,
ktLoopBody.statements.map { statementGenerator.generateStatement(it) }
ktLoopBody.statements.map { it.genStmt() }
)
fun generateBreak(ktBreak: KtBreakExpression): IrExpression {
@@ -199,7 +199,7 @@ class LoopExpressionGenerator(statementGenerator: StatementGenerator) : Statemen
}
if (ktForBody != null) {
irInnerBody.statements.add(statementGenerator.generateExpression(ktForBody))
irInnerBody.statements.add(ktForBody.genExpr())
}
return irForBlock
@@ -33,6 +33,7 @@ import org.jetbrains.kotlin.psi.psiUtil.startOffset
import org.jetbrains.kotlin.psi2ir.findSingleFunction
import org.jetbrains.kotlin.resolve.BindingContext
import org.jetbrains.kotlin.resolve.calls.model.ResolvedCall
import org.jetbrains.kotlin.resolve.checkers.PrimitiveNumericComparisonInfo
import org.jetbrains.kotlin.resolve.constants.evaluate.ConstantExpressionEvaluator
import org.jetbrains.kotlin.types.KotlinType
import org.jetbrains.kotlin.types.typeUtil.isPrimitiveNumberType
@@ -89,7 +90,7 @@ class OperatorExpressionGenerator(statementGenerator: StatementGenerator) : Stat
return IrTypeOperatorCallImpl(
expression.startOffset, expression.endOffset, resultType, irOperator, rhsType,
statementGenerator.generateExpression(expression.left)
expression.left.genExpr()
)
}
@@ -100,7 +101,7 @@ class OperatorExpressionGenerator(statementGenerator: StatementGenerator) : Stat
return IrTypeOperatorCallImpl(
expression.startOffset, expression.endOffset, context.builtIns.booleanType, irOperator,
againstType, statementGenerator.generateExpression(expression.leftHandSide)
againstType, expression.leftHandSide.genExpr()
)
}
@@ -130,8 +131,8 @@ class OperatorExpressionGenerator(statementGenerator: StatementGenerator) : Stat
private fun generateElvis(expression: KtBinaryExpression): IrExpression {
val specialCallForElvis = getResolvedCall(expression)!!
val resultType = specialCallForElvis.resultingDescriptor.returnType!!
val irArgument0 = statementGenerator.generateExpression(expression.left!!)
val irArgument1 = statementGenerator.generateExpression(expression.right!!)
val irArgument0 = expression.left!!.genExpr()
val irArgument1 = expression.right!!.genExpr()
return irBlock(expression, IrStatementOrigin.ELVIS, resultType) {
val temporary = irTemporary(irArgument0, "elvis_lhs")
@@ -140,8 +141,8 @@ class OperatorExpressionGenerator(statementGenerator: StatementGenerator) : Stat
}
private fun generateBinaryBooleanOperator(expression: KtBinaryExpression, irOperator: IrStatementOrigin): IrExpression {
val irArgument0 = statementGenerator.generateExpression(expression.left!!)
val irArgument1 = statementGenerator.generateExpression(expression.right!!)
val irArgument0 = expression.left!!.genExpr()
val irArgument1 = expression.right!!.genExpr()
return when (irOperator) {
IrStatementOrigin.OROR ->
context.oror(expression.startOffset, expression.endOffset, irArgument0, irArgument1)
@@ -173,8 +174,8 @@ class OperatorExpressionGenerator(statementGenerator: StatementGenerator) : Stat
}
private fun generateIdentityOperator(expression: KtBinaryExpression, irOperator: IrStatementOrigin): IrExpression {
val irArgument0 = statementGenerator.generateExpression(expression.left!!)
val irArgument1 = statementGenerator.generateExpression(expression.right!!)
val irArgument0 = expression.left!!.genExpr()
val irArgument1 = expression.right!!.genExpr()
val irIdentityEquals = IrBinaryPrimitiveImpl(
expression.startOffset, expression.endOffset, irOperator,
@@ -196,31 +197,27 @@ class OperatorExpressionGenerator(statementGenerator: StatementGenerator) : Stat
}
}
private fun KtExpression.generateAsPrimitiveNumericComparisonOperand(primitiveNumericComparisonType: KotlinType?) =
statementGenerator.generateExpression(this)
.promoteToPrimitiveNumericType(
getPrimitiveNumericComparisonOperandType(this),
primitiveNumericComparisonType
)
private fun KtExpression.generateAsPrimitiveNumericComparisonOperand(
expressionType: KotlinType?,
comparisonType: KotlinType?
) = genExpr().promoteToPrimitiveNumericType(expressionType, comparisonType)
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 getPrimitiveNumericComparisonInfo(ktExpression: KtBinaryExpression) =
context.bindingContext[BindingContext.PRIMITIVE_NUMERIC_COMPARISON_INFO, ktExpression]
private fun generateEqualityOperator(expression: KtBinaryExpression, irOperator: IrStatementOrigin): IrExpression {
val primitiveNumericComparisonType = getPrimitiveNumericComparisonType(expression)
val comparisonInfo = getPrimitiveNumericComparisonInfo(expression)
val comparisonType = comparisonInfo?.comparisonType
val eqeqSymbol = context.irBuiltIns.ieee754equalsFunByOperandType[primitiveNumericComparisonType]?.symbol
val eqeqSymbol = context.irBuiltIns.ieee754equalsFunByOperandType[comparisonType]?.symbol
?: context.irBuiltIns.eqeqSymbol
val irEquals = IrBinaryPrimitiveImpl(
expression.startOffset, expression.endOffset,
irOperator,
eqeqSymbol,
expression.left!!.generateAsPrimitiveNumericComparisonOperand(primitiveNumericComparisonType),
expression.right!!.generateAsPrimitiveNumericComparisonOperand(primitiveNumericComparisonType)
expression.left!!.generateAsPrimitiveNumericComparisonOperand(comparisonInfo?.leftType, comparisonType),
expression.right!!.generateAsPrimitiveNumericComparisonOperand(comparisonInfo?.rightType, comparisonType)
)
return when (irOperator) {
@@ -238,6 +235,33 @@ class OperatorExpressionGenerator(statementGenerator: StatementGenerator) : Stat
}
}
fun generateEquality(
startOffset: Int,
endOffset: Int,
irOperator: IrStatementOrigin,
arg1: IrExpression,
arg2: IrExpression,
comparisonInfo: PrimitiveNumericComparisonInfo?
): IrExpression =
if (comparisonInfo != null) {
val comparisonType = comparisonInfo.comparisonType
val eqeqSymbol =
context.irBuiltIns.ieee754equalsFunByOperandType[comparisonType]?.symbol
?: context.irBuiltIns.eqeqSymbol
IrBinaryPrimitiveImpl(
startOffset, endOffset, irOperator,
eqeqSymbol,
arg1.promoteToPrimitiveNumericType(comparisonInfo.leftType, comparisonType),
arg2.promoteToPrimitiveNumericType(comparisonInfo.rightType, comparisonType)
)
} else {
IrBinaryPrimitiveImpl(
startOffset, endOffset, irOperator,
context.irBuiltIns.eqeqSymbol,
arg1, arg2
)
}
private fun IrExpression.promoteToPrimitiveNumericType(operandType: KotlinType?, targetType: KotlinType?): IrExpression {
if (targetType == null) return this
if (operandType == null) throw AssertionError("operandType should be non-null")
@@ -290,14 +314,14 @@ class OperatorExpressionGenerator(statementGenerator: StatementGenerator) : Stat
val startOffset = expression.startOffset
val endOffset = expression.endOffset
val primitiveNumberComparisonType = getPrimitiveNumericComparisonType(expression)
val comparisonInfo = getPrimitiveNumericComparisonInfo(expression)
return if (primitiveNumberComparisonType != null) {
return if (comparisonInfo != null) {
IrBinaryPrimitiveImpl(
startOffset, endOffset, origin,
getComparisonOperatorSymbol(origin, primitiveNumberComparisonType),
expression.left!!.generateAsPrimitiveNumericComparisonOperand(primitiveNumberComparisonType),
expression.right!!.generateAsPrimitiveNumericComparisonOperand(primitiveNumberComparisonType)
getComparisonOperatorSymbol(origin, comparisonInfo.comparisonType),
expression.left!!.generateAsPrimitiveNumericComparisonOperand(comparisonInfo.leftType, comparisonInfo.comparisonType),
expression.right!!.generateAsPrimitiveNumericComparisonOperand(comparisonInfo.rightType, comparisonInfo.comparisonType)
)
} else {
IrBinaryPrimitiveImpl(
@@ -327,7 +351,7 @@ class OperatorExpressionGenerator(statementGenerator: StatementGenerator) : Stat
private fun generateExclExclOperator(expression: KtPostfixExpression, origin: IrStatementOrigin): IrExpression {
val ktArgument = expression.baseExpression!!
val irArgument = statementGenerator.generateExpression(ktArgument)
val irArgument = ktArgument.genExpr()
val ktOperator = expression.operationReference
val resultType = irArgument.type.makeNotNullable()
@@ -39,7 +39,7 @@ class ReflectionReferencesGenerator(statementGenerator: StatementGenerator) : St
return if (lhs is DoubleColonLHS.Expression && !lhs.isObjectQualifier) {
IrGetClassImpl(
ktClassLiteral.startOffset, ktClassLiteral.endOffset, resultType,
statementGenerator.generateExpression(ktArgument)
ktArgument.genExpr()
)
} else {
val typeConstructorDeclaration = lhs.type.constructor.declarationDescriptor
@@ -407,4 +407,7 @@ class StatementGenerator(
abstract class StatementGeneratorExtension(val statementGenerator: StatementGenerator) : GeneratorWithScope {
override val scope: Scope get() = statementGenerator.scope
override val context: GeneratorContext get() = statementGenerator.context
fun KtExpression.genExpr() = statementGenerator.generateExpression(this)
fun KtExpression.genStmt() = statementGenerator.generateStatement(this)
}
@@ -30,7 +30,7 @@ class TryCatchExpressionGenerator(statementGenerator: StatementGenerator) : Stat
val resultType = getInferredTypeWithImplicitCastsOrFail(ktTry)
val irTryCatch = IrTryImpl(ktTry.startOffset, ktTry.endOffset, resultType)
irTryCatch.tryResult = statementGenerator.generateExpression(ktTry.tryBlock)
irTryCatch.tryResult = ktTry.tryBlock.genExpr()
for (ktCatchClause in ktTry.catchClauses) {
val ktCatchParameter = ktCatchClause.catchParameter!!
@@ -45,13 +45,13 @@ class TryCatchExpressionGenerator(statementGenerator: StatementGenerator) : Stat
catchParameterDescriptor
)
).apply {
result = statementGenerator.generateExpression(ktCatchBody)
result = ktCatchBody.genExpr()
}
irTryCatch.catches.add(irCatch)
}
irTryCatch.finallyExpression = ktTry.finallyBlock?.let { statementGenerator.generateExpression(it.finalExpression) }
irTryCatch.finallyExpression = ktTry.finallyBlock?.run { finalExpression.genExpr() }
return irTryCatch
}