Generate equals/hashCode(): Support content equality from stdlib

#KT-22361 Fixed
This commit is contained in:
Alexey Sedunov
2018-07-06 21:17:39 +03:00
parent 5fcc6cfa0b
commit b441c76313
5 changed files with 76 additions and 30 deletions
@@ -24,6 +24,7 @@ import com.intellij.openapi.project.Project
import com.intellij.psi.PsiNameIdentifierOwner import com.intellij.psi.PsiNameIdentifierOwner
import com.intellij.util.IncorrectOperationException import com.intellij.util.IncorrectOperationException
import org.jetbrains.kotlin.builtins.KotlinBuiltIns import org.jetbrains.kotlin.builtins.KotlinBuiltIns
import org.jetbrains.kotlin.config.ApiVersion
import org.jetbrains.kotlin.config.LanguageFeature import org.jetbrains.kotlin.config.LanguageFeature
import org.jetbrains.kotlin.descriptors.ClassDescriptor import org.jetbrains.kotlin.descriptors.ClassDescriptor
import org.jetbrains.kotlin.descriptors.FunctionDescriptor import org.jetbrains.kotlin.descriptors.FunctionDescriptor
@@ -159,6 +160,42 @@ class KotlinGenerateEqualsAndHashcodeAction : KotlinGenerateMemberActionBase<Kot
} }
} }
private fun isNestedArray(variable: VariableDescriptor) =
KotlinBuiltIns.isArrayOrPrimitiveArray(variable.builtIns.getArrayElementType(variable.type))
private fun KtElement.canUseArrayContentFunctions() =
languageVersionSettings.apiVersion >= ApiVersion.KOTLIN_1_1
private fun generateArraysEqualsCall(
variable: VariableDescriptor,
canUseContentFunctions: Boolean,
arg1: String,
arg2: String
): String {
return if (canUseContentFunctions) {
val methodName = if (isNestedArray(variable)) "contentDeepEquals" else "contentEquals"
"$arg1.$methodName($arg2)"
} else {
val methodName = if (isNestedArray(variable)) "deepEquals" else "equals"
"java.util.Arrays.$methodName($arg1, $arg2)"
}
}
private fun generateArrayHashCodeCall(
variable: VariableDescriptor,
canUseContentFunctions: Boolean,
argument: String
): String {
return if (canUseContentFunctions) {
val methodName = if (isNestedArray(variable)) "contentDeepHashCode" else "contentHashCode"
val dot = if (TypeUtils.isNullableType(variable.type)) "?." else "."
"$argument$dot$methodName()"
} else {
val methodName = if (isNestedArray(variable)) "deepHashCode" else "hashCode"
"java.util.Arrays.$methodName($argument)"
}
}
private fun generateEquals(project: Project, info: Info, targetClass: KtClassOrObject): KtNamedFunction? { private fun generateEquals(project: Project, info: Info, targetClass: KtClassOrObject): KtNamedFunction? {
with(info) { with(info) {
if (!needEquals) return null if (!needEquals) return null
@@ -197,17 +234,26 @@ class KotlinGenerateEqualsAndHashcodeAction : KotlinGenerateMemberActionBase<Kot
append('\n') append('\n')
variablesForEquals.forEach { variablesForEquals.forEach {
val isNullable = TypeUtils.isNullableType(it.type)
val isArray = KotlinBuiltIns.isArrayOrPrimitiveArray(it.type)
val canUseArrayContentFunctions = targetClass.canUseArrayContentFunctions()
val propName = (DescriptorToSourceUtilsIde.getAnyDeclaration(project, it) as PsiNameIdentifierOwner).nameIdentifier!!.text val propName = (DescriptorToSourceUtilsIde.getAnyDeclaration(project, it) as PsiNameIdentifierOwner).nameIdentifier!!.text
val notEquals = when { val notEquals = when {
KotlinBuiltIns.isArrayOrPrimitiveArray(it.type) -> { isArray -> {
val isNestedArray = KotlinBuiltIns.isArrayOrPrimitiveArray(classDescriptor.builtIns.getArrayElementType(it.type)) "!${generateArraysEqualsCall(it, canUseArrayContentFunctions, propName, "$paramName.$propName")}"
val methodName = if (isNestedArray) "deepEquals" else "equals" } else -> {
"!java.util.Arrays.$methodName($propName, $paramName.$propName)"
}
else ->
"$propName != $paramName.$propName" "$propName != $paramName.$propName"
}
}
val equalsCheck = "if ($notEquals) return false\n"
if (isArray && isNullable && canUseArrayContentFunctions) {
append("if ($propName != null) {\n")
append("if ($paramName.$propName == null) return false\n")
append(equalsCheck)
append("} else if ($paramName.$propName != null) return false\n")
} else {
append(equalsCheck)
} }
append("if ($notEquals) return false\n")
} }
append('\n') append('\n')
@@ -234,9 +280,11 @@ class KotlinGenerateEqualsAndHashcodeAction : KotlinGenerateMemberActionBase<Kot
typeClass == builtIns.byte || typeClass == builtIns.short || typeClass == builtIns.int -> typeClass == builtIns.byte || typeClass == builtIns.short || typeClass == builtIns.int ->
ref ref
KotlinBuiltIns.isArrayOrPrimitiveArray(type) -> { KotlinBuiltIns.isArrayOrPrimitiveArray(type) -> {
val isNestedArray = KotlinBuiltIns.isArrayOrPrimitiveArray(builtIns.getArrayElementType(type)) val canUseArrayContentFunctions = targetClass.canUseArrayContentFunctions()
val methodName = if (isNestedArray) "deepHashCode" else "hashCode" val shouldWrapInLet = isNullable && !canUseArrayContentFunctions
if (isNullable) "$ref?.let { java.util.Arrays.$methodName(it) }" else "java.util.Arrays.$methodName($ref)" val hashCodeArg = if (shouldWrapInLet) "it" else ref
val hashCodeCall = generateArrayHashCodeCall(this, canUseArrayContentFunctions, hashCodeArg)
if (shouldWrapInLet) "$ref?.let { $hashCodeCall }" else hashCodeCall
} }
else -> else ->
if (isNullable) "$ref?.hashCode()" else "$ref.hashCode()" if (isNullable) "$ref?.hashCode()" else "$ref.hashCode()"
@@ -1,5 +1,3 @@
import java.util.Arrays
class A(val n: IntArray, val s: Array<String>) { class A(val n: IntArray, val s: Array<String>) {
val f: Float = 1.0f val f: Float = 1.0f
@@ -13,16 +11,16 @@ class A(val n: IntArray, val s: Array<String>) {
other as A other as A
if (!Arrays.equals(n, other.n)) return false if (!n.contentEquals(other.n)) return false
if (!Arrays.equals(s, other.s)) return false if (!s.contentEquals(other.s)) return false
if (f != other.f) return false if (f != other.f) return false
return true return true
} }
override fun hashCode(): Int { override fun hashCode(): Int {
var result = Arrays.hashCode(n) var result = n.contentHashCode()
result = 31 * result + Arrays.hashCode(s) result = 31 * result + s.contentHashCode()
result = 31 * result + f.hashCode() result = 31 * result + f.hashCode()
return result return result
} }
@@ -1,5 +1,3 @@
import java.util.Arrays
data class A(val a: IntArray) { data class A(val a: IntArray) {
<caret>override fun equals(other: Any?): Boolean { <caret>override fun equals(other: Any?): Boolean {
if (this === other) return true if (this === other) return true
@@ -7,12 +5,12 @@ data class A(val a: IntArray) {
other as A other as A
if (!Arrays.equals(a, other.a)) return false if (!a.contentEquals(other.a)) return false
return true return true
} }
override fun hashCode(): Int { override fun hashCode(): Int {
return Arrays.hashCode(a) return a.contentHashCode()
} }
} }
@@ -1,5 +1,3 @@
import java.util.Arrays
class EqKotlin(val a: Array<Array<String>>) { class EqKotlin(val a: Array<Array<String>>) {
<caret>override fun equals(other: Any?): Boolean { <caret>override fun equals(other: Any?): Boolean {
if (this === other) return true if (this === other) return true
@@ -7,12 +5,12 @@ class EqKotlin(val a: Array<Array<String>>) {
other as EqKotlin other as EqKotlin
if (!Arrays.deepEquals(a, other.a)) return false if (!a.contentDeepEquals(other.a)) return false
return true return true
} }
override fun hashCode(): Int { override fun hashCode(): Int {
return Arrays.deepHashCode(a) return a.contentDeepHashCode()
} }
} }
@@ -1,5 +1,3 @@
import java.util.Arrays
class A(val n: IntArray?, val s: Array<String>?) { class A(val n: IntArray?, val s: Array<String>?) {
val f: Float = 1.0f val f: Float = 1.0f
@@ -13,16 +11,22 @@ class A(val n: IntArray?, val s: Array<String>?) {
other as A other as A
if (!Arrays.equals(n, other.n)) return false if (n != null) {
if (!Arrays.equals(s, other.s)) return false if (other.n == null) return false
if (!n.contentEquals(other.n)) return false
} else if (other.n != null) return false
if (s != null) {
if (other.s == null) return false
if (!s.contentEquals(other.s)) return false
} else if (other.s != null) return false
if (f != other.f) return false if (f != other.f) return false
return true return true
} }
override fun hashCode(): Int { override fun hashCode(): Int {
var result = n?.let { Arrays.hashCode(it) } ?: 0 var result = n?.contentHashCode() ?: 0
result = 31 * result + (s?.let { Arrays.hashCode(it) } ?: 0) result = 31 * result + (s?.contentHashCode() ?: 0)
result = 31 * result + f.hashCode() result = 31 * result + f.hashCode()
return result return result
} }