[Kotlin] extract K2 equals and hashCode generation

KTIJ-28464

GitOrigin-RevId: 367f910260abb32cc7ab72a346061de263934b36
This commit is contained in:
Pavel Kirpichenkov
2024-07-11 11:41:35 +00:00
committed by intellij-monorepo-bot
parent 808bfe5e81
commit b3ff98c995
5 changed files with 348 additions and 325 deletions
@@ -56,5 +56,6 @@
<orderEntry type="module" module-name="intellij.java" />
<orderEntry type="module" module-name="kotlin.refactorings.common" />
<orderEntry type="module" module-name="intellij.platform.core.ui" />
<orderEntry type="module" module-name="kotlin.base.facet" />
</component>
</module>
@@ -0,0 +1,316 @@
// Copyright 2000-2024 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license.
package org.jetbrains.kotlin.idea.codeinsights.impl.base.intentions
import com.intellij.codeInsight.CodeInsightSettings
import org.jetbrains.kotlin.analysis.api.KaSession
import org.jetbrains.kotlin.analysis.api.analyze
import org.jetbrains.kotlin.analysis.api.symbols.*
import org.jetbrains.kotlin.analysis.api.types.KaType
import org.jetbrains.kotlin.builtins.StandardNames
import org.jetbrains.kotlin.config.ApiVersion
import org.jetbrains.kotlin.config.LanguageFeature
import org.jetbrains.kotlin.idea.base.codeInsight.KotlinNameSuggester
import org.jetbrains.kotlin.idea.base.facet.platform.platform
import org.jetbrains.kotlin.idea.base.projectStructure.languageVersionSettings
import org.jetbrains.kotlin.idea.base.psi.classIdIfNonLocal
import org.jetbrains.kotlin.idea.codeinsight.utils.isNonNullableBooleanType
import org.jetbrains.kotlin.idea.codeinsight.utils.isNullableAnyType
import org.jetbrains.kotlin.idea.core.CollectingNameValidator
import org.jetbrains.kotlin.name.Name
import org.jetbrains.kotlin.platform.isCommon
import org.jetbrains.kotlin.platform.isJs
import org.jetbrains.kotlin.psi.*
import org.jetbrains.kotlin.psi.psiUtil.hasExpectModifier
import org.jetbrains.kotlin.psi.psiUtil.isIdentifier
import org.jetbrains.kotlin.psi.psiUtil.quoteIfNeeded
import org.jetbrains.kotlin.util.OperatorNameConventions.EQUALS
import org.jetbrains.kotlin.util.OperatorNameConventions.HASH_CODE
context(KaSession)
fun matchesEqualsMethodSignature(function: KaNamedFunctionSymbol): Boolean {
if (function.modality == KaSymbolModality.ABSTRACT) return false
if (function.name != EQUALS) return false
if (function.typeParameters.isNotEmpty()) return false
val param = function.valueParameters.singleOrNull() ?: return false
val paramType = param.returnType
val returnType = function.returnType
return paramType.isNullableAnyType() && returnType.isNonNullableBooleanType()
}
context(KaSession)
fun matchesHashCodeMethodSignature(function: KaNamedFunctionSymbol): Boolean {
if (function.modality == KaSymbolModality.ABSTRACT) return false
if (function.name != HASH_CODE) return false
if (function.typeParameters.isNotEmpty()) return false
if (function.valueParameters.isNotEmpty()) return false
val returnType = function.returnType
return returnType.isInt && !returnType.isMarkedNullable
}
fun generateEqualsHeaderAndBodyTexts(targetClass: KtClass): Pair<String, String> {
val hasExpectModifier = targetClass.hasExpectModifier()
var functionText = buildString {
append("override fun equals(other: Any?): Boolean")
if (!hasExpectModifier) {
append(" {\n return super.equals(other)\n}")
}
}
var bodyText = ""
analyze(targetClass) {
val classSymbol = targetClass.getClassOrObjectSymbol() ?: return@analyze
val equalsMethod = findEqualsMethodForClass(classSymbol) as? KaNamedFunctionSymbol ?: return@analyze
val superContainingEqualsMethod = equalsMethod.containingDeclaration ?: return@analyze
val parameterName = equalsMethod.valueParameters.singleOrNull()?.name?.asString() ?: return@analyze
/**
* TODO: Correctly emits parameter type and return type when there are user-defined classes "Any" and "Boolean".
* The counterpart in FE1.0 uses [org.jetbrains.kotlin.idea.actions.generate.generateFunctionSkeleton] that automatically
* generates `equals(other: Any?)` by default and it generates `equals(other: kotlin.Any?)` if the project has a custom class
* named "Any". To preserve the same behavior for FIR, we tried
* [org.jetbrains.kotlin.analysis.api.components.KtSymbolDeclarationRendererMixIn.render] with
* [org.jetbrains.kotlin.analysis.api.components.KaBuiltinTypes.any],
* `builtinTypes.ANY.expandedClassSymbol.classId?.asSingleFqName()?.asString()`, and
* `org.jetbrains.kotlin.builtins.StandardNames.FqNames.any.render()`, but they all have the same rendering result
* regardless of the user-defined "Any" class. We specify the parameter and the return type as `kotlin.Any?` and
* `kotlin.Boolean` and rely on [org.jetbrains.kotlin.idea.base.analysis.api.utils.shortenReferences], but it always
* shortens the name regardless of the existence of the user-defined classes.
*/
functionText = buildString {
append("override fun equals(${parameterName}: kotlin.Any?): kotlin.Boolean")
if (!hasExpectModifier) {
append(" {\n return super.equals(${parameterName})\n}")
}
}
var typeForCast = targetClass.classIdIfNonLocal?.asSingleFqName()?.asString()
val typeParams = targetClass.typeParameters
if (typeParams.isNotEmpty()) {
typeForCast += typeParams.joinToString(prefix = "<", postfix = ">") { "*" }
}
val useIsCheck = CodeInsightSettings.getInstance().USE_INSTANCEOF_ON_EQUALS_PARAMETER
val isNotInstanceCondition = if (useIsCheck) {
"$parameterName !is $typeForCast"
} else {
generateClassLiteralsNotEqual(parameterName, targetClass)
}
bodyText = buildString {
append("if (this === $parameterName) return true\n")
append("if ($isNotInstanceCondition) return false\n")
if (superContainingEqualsMethod != builtinTypes.ANY.expandedSymbol) {
append("if (!super.equals($parameterName)) return false\n")
}
/**
* TODO: We have to add a wizard to select members used for equals. See how
* [org.jetbrains.kotlin.idea.actions.generate.KotlinGenerateEqualsAndHashcodeAction.prepareMembersInfo] uses
* [org.jetbrains.kotlin.idea.actions.generate.KotlinGenerateEqualsWizard].
*/
val variablesForEquals = getPropertiesToUseInGeneratedMember(targetClass)
if (variablesForEquals.isNotEmpty()) {
if (!useIsCheck) {
append("\n$parameterName as $typeForCast\n")
}
append('\n')
variablesForEquals.forEach {
val variableType = it.expressionType ?: return@forEach
val isNullableType = variableType.isMarkedNullable
val isArray = variableType.isArrayOrPrimitiveArray
val canUseArrayContentFunctions = targetClass.canUseArrayContentFunctions()
val propName = it.name ?: return@forEach
val notEquals = when {
isArray -> {
"!${
generateArraysEqualsCall(
variableType, canUseArrayContentFunctions, propName, "$parameterName.$propName"
)
}"
}
else -> {
"$propName != $parameterName.$propName"
}
}
val equalsCheck = "if ($notEquals) return false\n"
if (isArray && isNullableType && canUseArrayContentFunctions) {
append("if ($propName != null) {\n")
append("if ($parameterName.$propName == null) return false\n")
append(equalsCheck)
append("} else if ($parameterName.$propName != null) return false\n")
} else {
append(equalsCheck)
}
}
append('\n')
}
append("return true")
}
}
return Pair(functionText, bodyText)
}
fun generateHashCodeHeaderAndBodyTexts(targetClass: KtClass): Pair<String, String> {
val hasExpectModifier = targetClass.hasExpectModifier()
val functionText = buildString {
append("override fun hashCode(): kotlin.Int")
if (!hasExpectModifier) {
append(" {\n return super.hashCode()\n}")
}
}
var bodyText = ""
analyze(targetClass) {
fun KtNamedDeclaration.genVariableHashCode(parenthesesNeeded: Boolean): String {
val ref = name ?: return "0"
val type = returnType
val isNullable = type.isMarkedNullable
var text = when {
type.isEqualTo(builtinTypes.BYTE) || type.isEqualTo(builtinTypes.SHORT) || type.isEqualTo(builtinTypes.INT) -> ref
type.isArrayOrPrimitiveArray -> {
val canUseArrayContentFunctions = targetClass.canUseArrayContentFunctions()
val shouldWrapInLet = isNullable && !canUseArrayContentFunctions
val hashCodeArg = if (shouldWrapInLet) StandardNames.IMPLICIT_LAMBDA_PARAMETER_NAME.identifier else ref
val hashCodeCall = generateArrayHashCodeCall(type, canUseArrayContentFunctions, hashCodeArg)
if (shouldWrapInLet) "$ref?.let { $hashCodeCall }" else hashCodeCall
}
else -> if (isNullable) "$ref?.hashCode()" else "$ref.hashCode()"
}
if (isNullable) {
text += " ?: 0"
if (parenthesesNeeded) {
text = "($text)"
}
}
return text
}
val classSymbol = targetClass.getClassOrObjectSymbol() ?: return@analyze
val hashCodeMethod = findHashCodeMethodForClass(classSymbol) as? KaNamedFunctionSymbol ?: return@analyze
val superContainingHashCodeMethod = hashCodeMethod.containingDeclaration ?: return@analyze
/**
* TODO: We have to add a wizard to select members used for hashCode. See how
* [org.jetbrains.kotlin.idea.actions.generate.KotlinGenerateEqualsAndHashcodeAction.prepareMembersInfo] uses
* [org.jetbrains.kotlin.idea.actions.generate.KotlinGenerateEqualsWizard].
*/
val variablesForHashCode = getPropertiesToUseInGeneratedMember(targetClass)
val propertyIterator = variablesForHashCode.iterator()
val initialValue = when {
superContainingHashCodeMethod != builtinTypes.ANY.expandedSymbol -> "super.hashCode()"
propertyIterator.hasNext() -> propertyIterator.next().genVariableHashCode(false)
else -> generateClassLiteral(targetClass) + ".hashCode()"
}
bodyText =
if (propertyIterator.hasNext()) { // TODO: Confirm that `variablesForHashCode.map { it.name?.quoteIfNeeded()!! }` is safe here
val validator = CollectingNameValidator(variablesForHashCode.map { it.name?.quoteIfNeeded()!! })
val resultVarName = KotlinNameSuggester.suggestNameByName("result", validator)
buildString {
append("var $resultVarName = $initialValue\n")
propertyIterator.forEach { append("$resultVarName = 31 * $resultVarName + ${it.genVariableHashCode(true)}\n") }
append("return $resultVarName")
}
} else "return $initialValue"
}
return Pair(functionText, bodyText)
}
context(KaSession)
private fun findEqualsMethodForClass(classSymbol: KaClassSymbol): KaCallableSymbol? =
findMethod(classSymbol, EQUALS) { callableSymbol ->
(callableSymbol as? KaNamedFunctionSymbol)?.let { matchesEqualsMethodSignature(it) } == true
}
/**
* Finds methods whose name is [methodName] not only from the class [classSymbol] but also its parent classes,
* and returns method symbols after filtering them using [condition].
*/
context(KaSession)
private fun findMethod(
classSymbol: KaClassSymbol, methodName: Name, condition: (KaCallableSymbol) -> Boolean
): KaCallableSymbol? = classSymbol.memberScope.getCallableSymbols(methodName).filter(condition).singleOrNull()
context(KaSession)
private fun findHashCodeMethodForClass(classSymbol: KaClassSymbol): KaCallableSymbol? =
findMethod(classSymbol, HASH_CODE) { callableSymbol ->
(callableSymbol as? KaNamedFunctionSymbol)?.let { matchesHashCodeMethodSignature(it) } == true
}
/**
* A function to generate the "not equals" comparison between the class of `this` and the class of the parameter.
*/
private fun generateClassLiteralsNotEqual(paramName: String, targetClass: KtClassOrObject): String {
val defaultExpression = "javaClass != $paramName?.javaClass"
if (!targetClass.languageVersionSettings.supportsFeature(LanguageFeature.BoundCallableReferences)) return defaultExpression
return when {
targetClass.platform.isCommon() -> "other == null || this::class != $paramName::class"
targetClass.platform.isJs() -> "other == null || this::class.js != $paramName::class.js"
else -> defaultExpression
}
}
/**
* A function to generate the class i.e., `javaClass` or `this::class`.
*/
private fun generateClassLiteral(targetClass: KtClassOrObject): String {
val defaultExpression = "javaClass"
if (!targetClass.languageVersionSettings.supportsFeature(LanguageFeature.BoundCallableReferences)) return defaultExpression
return when {
targetClass.platform.isCommon() -> "this::class"
targetClass.platform.isJs() -> "this::class.js"
else -> defaultExpression
}
}
context(KaSession)
private fun getPropertiesToUseInGeneratedMember(classOrObject: KtClassOrObject): List<KtNamedDeclaration> =
buildList<KtNamedDeclaration> {
classOrObject.primaryConstructorParameters.filterTo(this) { it.hasValOrVar() }
classOrObject.declarations.asSequence().filterIsInstance<KtProperty>().filterTo(this) {
it.getVariableSymbol() is KaPropertySymbol
}
}.filter {
it.name?.quoteIfNeeded().isIdentifier()
}
private fun KtElement.canUseArrayContentFunctions() = languageVersionSettings.apiVersion >= ApiVersion.KOTLIN_1_1
context(KaSession)
private fun generateArraysEqualsCall(
type: KaType, canUseContentFunctions: Boolean, arg1: String, arg2: String
): String {
return if (canUseContentFunctions) {
val methodName = if (type.isNestedArray) "contentDeepEquals" else "contentEquals"
"$arg1.$methodName($arg2)"
} else {
val methodName = if (type.isNestedArray) "deepEquals" else "equals"
"java.util.Arrays.$methodName($arg1, $arg2)"
}
}
context(KaSession)
private fun generateArrayHashCodeCall(
variableType: KaType?, canUseContentFunctions: Boolean, argument: String
): String {
return if (canUseContentFunctions) {
val methodName = if (variableType?.isNestedArray == true) "contentDeepHashCode" else "contentHashCode"
val dot = if (variableType?.isMarkedNullable == true) "?." else "."
"$argument$dot$methodName()"
} else {
val methodName = if (variableType?.isNestedArray == true) "deepHashCode" else "hashCode"
"java.util.Arrays.$methodName($argument)"
}
}
@@ -0,0 +1,8 @@
// Copyright 2000-2024 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license.
package org.jetbrains.kotlin.idea.codeinsights.impl.base.quickFix
import org.jetbrains.kotlin.idea.base.resources.KotlinBundle
class GenerateEqualsFix(function: String, body: String) : GenerateFunctionFix(function, body) {
override fun getName() = KotlinBundle.message("equals.text")
}
@@ -0,0 +1,8 @@
// Copyright 2000-2024 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license.
package org.jetbrains.kotlin.idea.codeinsights.impl.base.quickFix
import org.jetbrains.kotlin.idea.base.resources.KotlinBundle
class GenerateHashCodeFix(function: String, body: String) : GenerateFunctionFix(function, body) {
override fun getName() = KotlinBundle.message("hash.code.text")
}
@@ -1,338 +1,28 @@
// Copyright 2000-2024 JetBrains s.r.o. and contributors. Use of this source code is governed by the Apache 2.0 license.
package org.jetbrains.kotlin.idea.k2.codeinsight.inspections
import com.intellij.codeInsight.CodeInsightSettings
import com.intellij.codeInspection.InspectionsBundle
import com.intellij.codeInspection.ProblemsHolder
import com.intellij.psi.PsiElementVisitor
import com.intellij.psi.createSmartPointer
import org.jetbrains.kotlin.analysis.api.KaSession
import org.jetbrains.kotlin.analysis.api.analyze
import org.jetbrains.kotlin.analysis.api.symbols.*
import org.jetbrains.kotlin.analysis.api.types.KaType
import org.jetbrains.kotlin.builtins.StandardNames
import org.jetbrains.kotlin.config.ApiVersion
import org.jetbrains.kotlin.config.LanguageFeature
import org.jetbrains.kotlin.descriptors.Modality
import org.jetbrains.kotlin.idea.base.codeInsight.KotlinNameSuggester
import org.jetbrains.kotlin.idea.base.facet.platform.platform
import org.jetbrains.kotlin.idea.base.projectStructure.languageVersionSettings
import org.jetbrains.kotlin.idea.base.psi.classIdIfNonLocal
import org.jetbrains.kotlin.analysis.api.symbols.KaNamedFunctionSymbol
import org.jetbrains.kotlin.idea.base.resources.KotlinBundle
import org.jetbrains.kotlin.idea.codeinsight.api.classic.inspections.AbstractKotlinInspection
import org.jetbrains.kotlin.idea.codeinsight.utils.DeletePsiElementsFix
import org.jetbrains.kotlin.idea.codeinsight.utils.isNonNullableBooleanType
import org.jetbrains.kotlin.idea.codeinsight.utils.isNullableAnyType
import org.jetbrains.kotlin.idea.codeinsights.impl.base.quickFix.GenerateFunctionFix
import org.jetbrains.kotlin.idea.core.CollectingNameValidator
import org.jetbrains.kotlin.name.Name
import org.jetbrains.kotlin.platform.isCommon
import org.jetbrains.kotlin.platform.isJs
import org.jetbrains.kotlin.psi.*
import org.jetbrains.kotlin.psi.psiUtil.hasExpectModifier
import org.jetbrains.kotlin.psi.psiUtil.isIdentifier
import org.jetbrains.kotlin.psi.psiUtil.quoteIfNeeded
import org.jetbrains.kotlin.util.OperatorNameConventions.EQUALS
import org.jetbrains.kotlin.util.OperatorNameConventions.HASH_CODE
import org.jetbrains.kotlin.idea.codeinsights.impl.base.intentions.generateEqualsHeaderAndBodyTexts
import org.jetbrains.kotlin.idea.codeinsights.impl.base.intentions.generateHashCodeHeaderAndBodyTexts
import org.jetbrains.kotlin.idea.codeinsights.impl.base.intentions.matchesEqualsMethodSignature
import org.jetbrains.kotlin.idea.codeinsights.impl.base.intentions.matchesHashCodeMethodSignature
import org.jetbrains.kotlin.idea.codeinsights.impl.base.quickFix.GenerateEqualsFix
import org.jetbrains.kotlin.idea.codeinsights.impl.base.quickFix.GenerateHashCodeFix
import org.jetbrains.kotlin.psi.KtClass
import org.jetbrains.kotlin.psi.KtNamedFunction
import org.jetbrains.kotlin.psi.KtObjectDeclaration
import org.jetbrains.kotlin.psi.classOrObjectVisitor
internal class EqualsOrHashCodeInspection : AbstractKotlinInspection() {
private class Equals(function: String, body: String) : GenerateFunctionFix(function, body) {
override fun getName() = KotlinBundle.message("equals.text")
}
private class HashCode(function: String, body: String) : GenerateFunctionFix(function, body) {
override fun getName() = KotlinBundle.message("hash.code.text")
}
context(KaSession)
private fun matchesEqualsMethodSignature(function: KaNamedFunctionSymbol): Boolean {
if (function.modality == KaSymbolModality.ABSTRACT) return false
if (function.name != EQUALS) return false
if (function.typeParameters.isNotEmpty()) return false
val param = function.valueParameters.singleOrNull() ?: return false
val paramType = param.returnType
val returnType = function.returnType
return with(this) {
paramType.isNullableAnyType() && returnType.isNonNullableBooleanType()
}
}
context(KaSession)
private fun matchesHashCodeMethodSignature(function: KaNamedFunctionSymbol): Boolean {
if (function.modality == KaSymbolModality.ABSTRACT) return false
if (function.name != HASH_CODE) return false
if (function.typeParameters.isNotEmpty()) return false
if (function.valueParameters.isNotEmpty()) return false
val returnType = function.returnType
return returnType.isInt && !returnType.isMarkedNullable
}
/**
* Finds methods whose name is [methodName] not only from the class [classSymbol] but also its parent classes,
* and returns method symbols after filtering them using [condition].
*/
context(KaSession)
private fun findMethod(
classSymbol: KaClassSymbol, methodName: Name, condition: (KaCallableSymbol) -> Boolean
): KaCallableSymbol? = classSymbol.memberScope.getCallableSymbols(methodName).filter(condition).singleOrNull()
context(KaSession)
private fun findEqualsMethodForClass(classSymbol: KaClassSymbol): KaCallableSymbol? =
findMethod(classSymbol, EQUALS) { callableSymbol ->
(callableSymbol as? KaNamedFunctionSymbol)?.let { matchesEqualsMethodSignature(it) } == true
}
context(KaSession)
private fun findHashCodeMethodForClass(classSymbol: KaClassSymbol): KaCallableSymbol? =
findMethod(classSymbol, HASH_CODE) { callableSymbol ->
(callableSymbol as? KaNamedFunctionSymbol)?.let { matchesHashCodeMethodSignature(it) } == true
}
/**
* A function to generate the "not equals" comparison between the class of `this` and the class of the parameter.
*/
private fun generateClassLiteralsNotEqual(paramName: String, targetClass: KtClassOrObject): String {
val defaultExpression = "javaClass != $paramName?.javaClass"
if (!targetClass.languageVersionSettings.supportsFeature(LanguageFeature.BoundCallableReferences)) return defaultExpression
return when {
targetClass.platform.isCommon() -> "other == null || this::class != $paramName::class"
targetClass.platform.isJs() -> "other == null || this::class.js != $paramName::class.js"
else -> defaultExpression
}
}
/**
* A function to generate the class i.e., `javaClass` or `this::class`.
*/
private fun generateClassLiteral(targetClass: KtClassOrObject): String {
val defaultExpression = "javaClass"
if (!targetClass.languageVersionSettings.supportsFeature(LanguageFeature.BoundCallableReferences)) return defaultExpression
return when {
targetClass.platform.isCommon() -> "this::class"
targetClass.platform.isJs() -> "this::class.js"
else -> defaultExpression
}
}
context(KaSession)
private fun getPropertiesToUseInGeneratedMember(classOrObject: KtClassOrObject): List<KtNamedDeclaration> =
buildList<KtNamedDeclaration> {
classOrObject.primaryConstructorParameters.filterTo(this) { it.hasValOrVar() }
classOrObject.declarations.asSequence().filterIsInstance<KtProperty>().filterTo(this) {
it.getVariableSymbol() is KaPropertySymbol
}
}.filter {
it.name?.quoteIfNeeded().isIdentifier()
}
private fun KtElement.canUseArrayContentFunctions() = languageVersionSettings.apiVersion >= ApiVersion.KOTLIN_1_1
context(KaSession)
private fun generateArraysEqualsCall(
type: KaType, canUseContentFunctions: Boolean, arg1: String, arg2: String
): String {
return if (canUseContentFunctions) {
val methodName = if (type.isNestedArray) "contentDeepEquals" else "contentEquals"
"$arg1.$methodName($arg2)"
} else {
val methodName = if (type.isNestedArray) "deepEquals" else "equals"
"java.util.Arrays.$methodName($arg1, $arg2)"
}
}
context(KaSession)
private fun generateArrayHashCodeCall(
variableType: KaType?, canUseContentFunctions: Boolean, argument: String
): String {
return if (canUseContentFunctions) {
val methodName = if (variableType?.isNestedArray == true) "contentDeepHashCode" else "contentHashCode"
val dot = if (variableType?.isMarkedNullable == true) "?." else "."
"$argument$dot$methodName()"
} else {
val methodName = if (variableType?.isNestedArray == true) "deepHashCode" else "hashCode"
"java.util.Arrays.$methodName($argument)"
}
}
private fun generateEqualsFunctionAndBodyTexts(targetClass: KtClass): Pair<String, String> {
val hasExpectModifier = targetClass.hasExpectModifier()
var functionText = buildString {
append("override fun equals(other: Any?): Boolean")
if (!hasExpectModifier) {
append(" {\n return super.equals(other)\n}")
}
}
var bodyText = ""
analyze(targetClass) {
val classSymbol = targetClass.getClassOrObjectSymbol() ?: return@analyze
val equalsMethod = findEqualsMethodForClass(classSymbol) as? KaNamedFunctionSymbol ?: return@analyze
val superContainingEqualsMethod = equalsMethod.containingDeclaration ?: return@analyze
val parameterName = equalsMethod.valueParameters.singleOrNull()?.name?.asString() ?: return@analyze
/**
* TODO: Correctly emits parameter type and return type when there are user-defined classes "Any" and "Boolean".
* The counterpart in FE1.0 uses [org.jetbrains.kotlin.idea.actions.generate.generateFunctionSkeleton] that automatically
* generates `equals(other: Any?)` by default and it generates `equals(other: kotlin.Any?)` if the project has a custom class
* named "Any". To preserve the same behavior for FIR, we tried
* [org.jetbrains.kotlin.analysis.api.components.KtSymbolDeclarationRendererMixIn.render] with
* [org.jetbrains.kotlin.analysis.api.components.KaBuiltinTypes.any],
* `builtinTypes.ANY.expandedClassSymbol.classId?.asSingleFqName()?.asString()`, and
* `org.jetbrains.kotlin.builtins.StandardNames.FqNames.any.render()`, but they all have the same rendering result
* regardless of the user-defined "Any" class. We specify the parameter and the return type as `kotlin.Any?` and
* `kotlin.Boolean` and rely on [org.jetbrains.kotlin.idea.base.analysis.api.utils.shortenReferences], but it always
* shortens the name regardless of the existence of the user-defined classes.
*/
functionText = buildString {
append("override fun equals(${parameterName}: kotlin.Any?): kotlin.Boolean")
if (!hasExpectModifier) {
append(" {\n return super.equals(${parameterName})\n}")
}
}
var typeForCast = targetClass.classIdIfNonLocal?.asSingleFqName()?.asString()
val typeParams = targetClass.typeParameters
if (typeParams.isNotEmpty()) {
typeForCast += typeParams.joinToString(prefix = "<", postfix = ">") { "*" }
}
val useIsCheck = CodeInsightSettings.getInstance().USE_INSTANCEOF_ON_EQUALS_PARAMETER
val isNotInstanceCondition = if (useIsCheck) {
"$parameterName !is $typeForCast"
} else {
generateClassLiteralsNotEqual(parameterName, targetClass)
}
bodyText = buildString {
append("if (this === $parameterName) return true\n")
append("if ($isNotInstanceCondition) return false\n")
if (superContainingEqualsMethod != builtinTypes.ANY.expandedSymbol) {
append("if (!super.equals($parameterName)) return false\n")
}
/**
* TODO: We have to add a wizard to select members used for equals. See how
* [org.jetbrains.kotlin.idea.actions.generate.KotlinGenerateEqualsAndHashcodeAction.prepareMembersInfo] uses
* [org.jetbrains.kotlin.idea.actions.generate.KotlinGenerateEqualsWizard].
*/
val variablesForEquals = getPropertiesToUseInGeneratedMember(targetClass)
if (variablesForEquals.isNotEmpty()) {
if (!useIsCheck) {
append("\n$parameterName as $typeForCast\n")
}
append('\n')
variablesForEquals.forEach {
val variableType = it.expressionType ?: return@forEach
val isNullableType = variableType.isMarkedNullable
val isArray = variableType.isArrayOrPrimitiveArray
val canUseArrayContentFunctions = targetClass.canUseArrayContentFunctions()
val propName = it.name ?: return@forEach
val notEquals = when {
isArray -> {
"!${
generateArraysEqualsCall(
variableType, canUseArrayContentFunctions, propName, "$parameterName.$propName"
)
}"
}
else -> {
"$propName != $parameterName.$propName"
}
}
val equalsCheck = "if ($notEquals) return false\n"
if (isArray && isNullableType && canUseArrayContentFunctions) {
append("if ($propName != null) {\n")
append("if ($parameterName.$propName == null) return false\n")
append(equalsCheck)
append("} else if ($parameterName.$propName != null) return false\n")
} else {
append(equalsCheck)
}
}
append('\n')
}
append("return true")
}
}
return Pair(functionText, bodyText)
}
private fun generateHashCodeFunctionAndBodyTexts(targetClass: KtClass): Pair<String, String> {
val hasExpectModifier = targetClass.hasExpectModifier()
val functionText = buildString {
append("override fun hashCode(): kotlin.Int")
if (!hasExpectModifier) {
append(" {\n return super.hashCode()\n}")
}
}
var bodyText = ""
analyze(targetClass) {
fun KtNamedDeclaration.genVariableHashCode(parenthesesNeeded: Boolean): String {
val ref = name ?: return "0"
val type = returnType
val isNullable = type.isMarkedNullable
var text = when {
type.isEqualTo(builtinTypes.BYTE) || type.isEqualTo(builtinTypes.SHORT) || type.isEqualTo(builtinTypes.INT) -> ref
type.isArrayOrPrimitiveArray -> {
val canUseArrayContentFunctions = targetClass.canUseArrayContentFunctions()
val shouldWrapInLet = isNullable && !canUseArrayContentFunctions
val hashCodeArg = if (shouldWrapInLet) StandardNames.IMPLICIT_LAMBDA_PARAMETER_NAME.identifier else ref
val hashCodeCall = generateArrayHashCodeCall(type, canUseArrayContentFunctions, hashCodeArg)
if (shouldWrapInLet) "$ref?.let { $hashCodeCall }" else hashCodeCall
}
else -> if (isNullable) "$ref?.hashCode()" else "$ref.hashCode()"
}
if (isNullable) {
text += " ?: 0"
if (parenthesesNeeded) {
text = "($text)"
}
}
return text
}
val classSymbol = targetClass.getClassOrObjectSymbol() ?: return@analyze
val hashCodeMethod = findHashCodeMethodForClass(classSymbol) as? KaNamedFunctionSymbol ?: return@analyze
val superContainingHashCodeMethod = hashCodeMethod.containingDeclaration ?: return@analyze
/**
* TODO: We have to add a wizard to select members used for hashCode. See how
* [org.jetbrains.kotlin.idea.actions.generate.KotlinGenerateEqualsAndHashcodeAction.prepareMembersInfo] uses
* [org.jetbrains.kotlin.idea.actions.generate.KotlinGenerateEqualsWizard].
*/
val variablesForHashCode = getPropertiesToUseInGeneratedMember(targetClass)
val propertyIterator = variablesForHashCode.iterator()
val initialValue = when {
superContainingHashCodeMethod != builtinTypes.ANY.expandedSymbol -> "super.hashCode()"
propertyIterator.hasNext() -> propertyIterator.next().genVariableHashCode(false)
else -> generateClassLiteral(targetClass) + ".hashCode()"
}
bodyText =
if (propertyIterator.hasNext()) { // TODO: Confirm that `variablesForHashCode.map { it.name?.quoteIfNeeded()!! }` is safe here
val validator = CollectingNameValidator(variablesForHashCode.map { it.name?.quoteIfNeeded()!! })
val resultVarName = KotlinNameSuggester.suggestNameByName("result", validator)
buildString {
append("var $resultVarName = $initialValue\n")
propertyIterator.forEach { append("$resultVarName = 31 * $resultVarName + ${it.genVariableHashCode(true)}\n") }
append("return $resultVarName")
}
} else "return $initialValue"
}
return Pair(functionText, bodyText)
}
override fun buildVisitor(holder: ProblemsHolder, isOnTheFly: Boolean): PsiElementVisitor {
return classOrObjectVisitor(fun(classOrObject) {
val nameIdentifier = classOrObject.nameIdentifier ?: return
@@ -373,11 +63,11 @@ internal class EqualsOrHashCodeInspection : AbstractKotlinInspection() {
)
val fix = if (equalsDeclaration != null) {
val (function, body) = generateHashCodeFunctionAndBodyTexts(classOrObject)
HashCode(function, body)
val (function, body) = generateHashCodeHeaderAndBodyTexts(classOrObject)
GenerateHashCodeFix(function, body)
} else {
val (function, body) = generateEqualsFunctionAndBodyTexts(classOrObject)
Equals(function, body)
val (function, body) = generateEqualsHeaderAndBodyTexts(classOrObject)
GenerateEqualsFix(function, body)
}
holder.registerProblem(