Generate equals/hashCode(): Support content equality from stdlib
#KT-22361 Fixed
This commit is contained in:
+58
-10
@@ -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
|
||||||
}
|
}
|
||||||
|
|||||||
+2
-4
@@ -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()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
+10
-6
@@ -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
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user