[Kotlin] Implement K2 'array in data class' inspection

KTIJ-28464

GitOrigin-RevId: 63c220f1e93328f16912d5b64dd1686e46967d5a
This commit is contained in:
Pavel Kirpichenkov
2024-07-11 11:41:35 +00:00
committed by intellij-monorepo-bot
parent c7193e5a38
commit 63d99e08a4
26 changed files with 527 additions and 18 deletions
@@ -7,6 +7,7 @@ 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.analysis.api.types.symbol
import org.jetbrains.kotlin.builtins.StandardNames
import org.jetbrains.kotlin.config.ApiVersion
import org.jetbrains.kotlin.config.LanguageFeature
@@ -118,7 +119,7 @@ fun generateEqualsHeaderAndBodyTexts(targetClass: KtClass): Pair<String, String>
append('\n')
variablesForEquals.forEach {
val variableType = it.expressionType ?: return@forEach
val variableType = it.returnType
val isNullableType = variableType.isMarkedNullable
val isArray = variableType.isArrayOrPrimitiveArray
val canUseArrayContentFunctions = targetClass.canUseArrayContentFunctions()
@@ -229,9 +230,26 @@ fun generateHashCodeHeaderAndBodyTexts(targetClass: KtClass): Pair<String, Strin
context(KaSession)
private fun findEqualsMethodForClass(classSymbol: KaClassSymbol): KaCallableSymbol? =
findMethod(classSymbol, EQUALS) { callableSymbol ->
(callableSymbol as? KaNamedFunctionSymbol)?.let { matchesEqualsMethodSignature(it) } == true
}
findMethodInMemberScopeOrInAny(classSymbol, EQUALS) { matchesEqualsMethodSignature(it) }
context(KaSession)
private fun findHashCodeMethodForClass(classSymbol: KaClassSymbol): KaCallableSymbol? =
findMethodInMemberScopeOrInAny(classSymbol, HASH_CODE) { matchesHashCodeMethodSignature(it) }
context(KaSession)
private fun findMethodInMemberScopeOrInAny(
classSymbol: KaClassSymbol,
methodName: Name,
signatureFilter: (KaNamedFunctionSymbol) -> Boolean
): KaCallableSymbol? {
findMethod(classSymbol, methodName) { callableSymbol ->
if (callableSymbol !is KaNamedFunctionSymbol) return@findMethod false
signatureFilter(callableSymbol) && callableSymbol.origin != KaSymbolOrigin.SOURCE_MEMBER_GENERATED
}?.let { return it }
val anySuperClassSymbol = classSymbol.superTypes.find { it.isAnyType }?.symbol as? KaClassSymbol ?: return null
return anySuperClassSymbol.memberScope.callables(methodName).singleOrNull()
}
/**
* Finds methods whose name is [methodName] not only from the class [classSymbol] but also its parent classes,
@@ -242,12 +260,6 @@ private fun findMethod(
classSymbol: KaClassSymbol, methodName: Name, condition: (KaCallableSymbol) -> Boolean
): KaCallableSymbol? = classSymbol.memberScope.callables(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.
*/
@@ -415,5 +415,14 @@
language="kotlin" editorAttributes="NOT_USED_ELEMENT_ATTRIBUTES"
key="inspection.can.be.parameter.display.name" bundle="messages.KotlinBundle"/>
<localInspection implementationClass="org.jetbrains.kotlin.idea.k2.codeinsight.inspections.declarations.ArrayInDataClassInspection"
shortName="ArrayInDataClass"
groupPath="Kotlin"
groupBundle="messages.KotlinBundle" groupKey="group.names.probable.bugs"
enabledByDefault="true"
level="WARNING"
language="kotlin"
key="inspection.array.in.data.class.display.name" bundle="messages.KotlinBundle"/>
</extensions>
</idea-plugin>
@@ -0,0 +1,142 @@
// 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.declarations
import com.intellij.codeInspection.ProblemsHolder
import com.intellij.modcommand.ModPsiUpdater
import com.intellij.openapi.project.Project
import org.jetbrains.kotlin.analysis.api.KaSession
import org.jetbrains.kotlin.analysis.api.analyze
import org.jetbrains.kotlin.analysis.api.symbols.KaFunctionSymbol
import org.jetbrains.kotlin.idea.base.analysis.api.utils.shortenReferences
import org.jetbrains.kotlin.idea.base.resources.KotlinBundle
import org.jetbrains.kotlin.idea.codeinsight.api.applicable.inspections.KotlinApplicableInspectionBase
import org.jetbrains.kotlin.idea.codeinsight.api.applicable.inspections.KotlinModCommandQuickFix
import org.jetbrains.kotlin.idea.codeinsight.utils.isNullableAnyType
import org.jetbrains.kotlin.idea.codeinsights.impl.base.intentions.generateEqualsHeaderAndBodyTexts
import org.jetbrains.kotlin.idea.codeinsights.impl.base.intentions.generateHashCodeHeaderAndBodyTexts
import org.jetbrains.kotlin.lexer.KtTokens
import org.jetbrains.kotlin.psi.*
import org.jetbrains.kotlin.psi.psiUtil.containingClass
import org.jetbrains.kotlin.util.OperatorNameConventions
import org.jetbrains.kotlin.utils.addToStdlib.safeAs
class ArrayInDataClassInspection : KotlinApplicableInspectionBase.Simple<KtParameter, ArrayInDataClassInspection.Context>() {
class Context(
val equalsHeader: String?,
val equalsBody: String?,
val hashCodeHeader: String?,
val hashCodeBody: String?,
) {
init {
check((equalsHeader == null) == (equalsBody == null))
check((hashCodeHeader == null) == (hashCodeBody == null))
}
}
override fun getProblemDescription(element: KtParameter, context: Context): String {
return KotlinBundle.message("array.property.in.data.class.it.s.recommended.to.override.equals.hashcode")
}
override fun createQuickFix(element: KtParameter, context: Context): KotlinModCommandQuickFix<KtParameter> {
return object : KotlinModCommandQuickFix<KtParameter>() {
override fun getFamilyName(): String =
KotlinBundle.message("generate.equals.and.hashcode.fix.text")
override fun applyFix(project: Project, element: KtParameter, updater: ModPsiUpdater): Unit = with(context) {
val psiFactory = KtPsiFactory(project, markGenerated = true)
val containingClass = element.containingClass() ?: return
if (equalsHeader != null && equalsBody != null) {
generateFunctionDeclarationInClass(psiFactory, containingClass, equalsHeader, equalsBody)
}
if (hashCodeHeader != null && hashCodeBody != null) {
generateFunctionDeclarationInClass(psiFactory, containingClass, hashCodeHeader, hashCodeBody)
}
}
private fun generateFunctionDeclarationInClass(factory: KtPsiFactory, containingClass: KtClass, header: String, body: String) {
val function = factory.createFunction(header)
shortenReferences(function)
if (body.isNotEmpty()) function.bodyExpression?.replace(factory.createBlock(body))
containingClass.addDeclaration(function)
}
}
}
override fun buildVisitor(holder: ProblemsHolder, isOnTheFly: Boolean): KtVisitor<*, *> {
return classVisitor { klass ->
if (!klass.isData()) return@classVisitor
val constructor = klass.primaryConstructor ?: return@classVisitor
for (parameter in constructor.valueParameters) {
visitTargetElement(parameter, holder, isOnTheFly)
}
}
}
override fun isApplicableByPsi(element: KtParameter): Boolean {
return element.hasValOrVar()
}
context(KaSession)
override fun prepareContext(element: KtParameter): Context? {
val parameterType = element.symbol.returnType
if (!parameterType.isArrayOrPrimitiveArray) return null
val containingClass = element.containingClass() ?: return null
return when (checkOverriddenEqualsAndHashCode(containingClass)) {
EqualsHashCodeOverrides.HAS_EQUALS_AND_HASHCODE -> null
EqualsHashCodeOverrides.HAS_EQUALS -> {
val (hashCodeHeader, hashCodeBody) = generateHashCodeHeaderAndBodyTexts(containingClass)
Context(equalsHeader = null, equalsBody = null, hashCodeHeader, hashCodeBody)
}
EqualsHashCodeOverrides.HAS_HASHCODE -> {
val (equalsHeader, equalsBody) = generateEqualsHeaderAndBodyTexts(containingClass)
Context(equalsHeader, equalsBody, hashCodeHeader = null, hashCodeBody = null)
}
EqualsHashCodeOverrides.HAS_NONE -> {
val (equalsHeader, equalsBody) = generateEqualsHeaderAndBodyTexts(containingClass)
val (hashCodeHeader, hashCodeBody) = generateHashCodeHeaderAndBodyTexts(containingClass)
Context(equalsHeader, equalsBody, hashCodeHeader, hashCodeBody)
}
}
}
private fun checkOverriddenEqualsAndHashCode(klass: KtClass): EqualsHashCodeOverrides {
var overriddenEquals = false
var overriddenHashCode = false
for (declaration in klass.declarations) {
if (declaration !is KtFunction) continue
if (!declaration.hasModifier(KtTokens.OVERRIDE_KEYWORD)) continue
if (declaration.nameAsName == OperatorNameConventions.EQUALS && declaration.valueParameters.size == 1) {
analyze(declaration) {
val parameterType = declaration.symbol.safeAs<KaFunctionSymbol>()?.valueParameters?.singleOrNull()?.returnType
if (parameterType?.isNullableAnyType() == true) {
overriddenEquals = true
}
}
}
if (declaration.nameAsName == OperatorNameConventions.HASH_CODE && declaration.valueParameters.size == 0) {
overriddenHashCode = true
}
}
return EqualsHashCodeOverrides.of(overriddenEquals, overriddenHashCode)
}
private enum class EqualsHashCodeOverrides {
HAS_EQUALS_AND_HASHCODE,
HAS_EQUALS,
HAS_HASHCODE,
HAS_NONE;
companion object {
fun of(hasEquals: Boolean, hasHashCode: Boolean): EqualsHashCodeOverrides = when {
hasEquals && hasHashCode -> HAS_EQUALS_AND_HASHCODE
hasEquals -> HAS_EQUALS
hasHashCode -> HAS_HASHCODE
else -> HAS_NONE
}
}
}
}
@@ -373,4 +373,27 @@ public abstract class K2InspectionTestGenerated extends AbstractK2InspectionTest
}
}
}
@RunWith(JUnit3RunnerWithInners.class)
@TestMetadata("../../../idea/tests/testData/inspections/arrayInDataClass")
public abstract static class ArrayInDataClass extends AbstractK2InspectionTest {
@RunWith(JUnit3RunnerWithInners.class)
@TestMetadata("../../../idea/tests/testData/inspections/arrayInDataClass/inspectionData")
public static class InspectionData extends AbstractK2InspectionTest {
@java.lang.Override
@org.jetbrains.annotations.NotNull
public final KotlinPluginMode getPluginMode() {
return KotlinPluginMode.K2;
}
private void runTest(String testDataFilePath) throws Exception {
KotlinTestUtils.runTest(this::doTest, this, testDataFilePath);
}
@TestMetadata("inspections.test")
public void testInspections_test() throws Exception {
runTest("../../../idea/tests/testData/inspections/arrayInDataClass/inspectionData/inspections.test");
}
}
}
}
@@ -6648,6 +6648,65 @@ public abstract class K2LocalInspectionTestGenerated extends AbstractK2LocalInsp
}
}
@RunWith(JUnit3RunnerWithInners.class)
@TestMetadata("../../../idea/tests/testData/inspectionsLocal/arrayInDataClass")
public static class ArrayInDataClass extends AbstractK2LocalInspectionTest {
@java.lang.Override
@org.jetbrains.annotations.NotNull
public final KotlinPluginMode getPluginMode() {
return KotlinPluginMode.K2;
}
private void runTest(String testDataFilePath) throws Exception {
KotlinTestUtils.runTest(this::doTest, this, testDataFilePath);
}
@TestMetadata("genericArray.kt")
public void testGenericArray() throws Exception {
runTest("../../../idea/tests/testData/inspectionsLocal/arrayInDataClass/genericArray.kt");
}
@TestMetadata("intArray.kt")
public void testIntArray() throws Exception {
runTest("../../../idea/tests/testData/inspectionsLocal/arrayInDataClass/intArray.kt");
}
@TestMetadata("justEquals.kt")
public void testJustEquals() throws Exception {
runTest("../../../idea/tests/testData/inspectionsLocal/arrayInDataClass/justEquals.kt");
}
@TestMetadata("justHashCode.kt")
public void testJustHashCode() throws Exception {
runTest("../../../idea/tests/testData/inspectionsLocal/arrayInDataClass/justHashCode.kt");
}
@TestMetadata("mixedParameters.kt")
public void testMixedParameters() throws Exception {
runTest("../../../idea/tests/testData/inspectionsLocal/arrayInDataClass/mixedParameters.kt");
}
@TestMetadata("negativeEqualsHashCodeOverrides.kt")
public void testNegativeEqualsHashCodeOverrides() throws Exception {
runTest("../../../idea/tests/testData/inspectionsLocal/arrayInDataClass/negativeEqualsHashCodeOverrides.kt");
}
@TestMetadata("negativeNonArray.kt")
public void testNegativeNonArray() throws Exception {
runTest("../../../idea/tests/testData/inspectionsLocal/arrayInDataClass/negativeNonArray.kt");
}
@TestMetadata("nonOverrideEquals.kt")
public void testNonOverrideEquals() throws Exception {
runTest("../../../idea/tests/testData/inspectionsLocal/arrayInDataClass/nonOverrideEquals.kt");
}
@TestMetadata("nonOverrideHashCode.kt")
public void testNonOverrideHashCode() throws Exception {
runTest("../../../idea/tests/testData/inspectionsLocal/arrayInDataClass/nonOverrideHashCode.kt");
}
}
@RunWith(JUnit3RunnerWithInners.class)
@TestMetadata("testData/inspectionsLocal")
public abstract static class InspectionsLocal extends AbstractK2LocalInspectionTest {
@@ -82,9 +82,49 @@ public abstract class LocalInspectionTestGenerated extends AbstractLocalInspecti
KotlinTestUtils.runTest(this::doTest, this, testDataFilePath);
}
@TestMetadata("test.kt")
public void testTest() throws Exception {
runTest("testData/inspectionsLocal/arrayInDataClass/test.kt");
@TestMetadata("genericArray.kt")
public void testGenericArray() throws Exception {
runTest("testData/inspectionsLocal/arrayInDataClass/genericArray.kt");
}
@TestMetadata("intArray.kt")
public void testIntArray() throws Exception {
runTest("testData/inspectionsLocal/arrayInDataClass/intArray.kt");
}
@TestMetadata("justEquals.kt")
public void testJustEquals() throws Exception {
runTest("testData/inspectionsLocal/arrayInDataClass/justEquals.kt");
}
@TestMetadata("justHashCode.kt")
public void testJustHashCode() throws Exception {
runTest("testData/inspectionsLocal/arrayInDataClass/justHashCode.kt");
}
@TestMetadata("mixedParameters.kt")
public void testMixedParameters() throws Exception {
runTest("testData/inspectionsLocal/arrayInDataClass/mixedParameters.kt");
}
@TestMetadata("negativeEqualsHashCodeOverrides.kt")
public void testNegativeEqualsHashCodeOverrides() throws Exception {
runTest("testData/inspectionsLocal/arrayInDataClass/negativeEqualsHashCodeOverrides.kt");
}
@TestMetadata("negativeNonArray.kt")
public void testNegativeNonArray() throws Exception {
runTest("testData/inspectionsLocal/arrayInDataClass/negativeNonArray.kt");
}
@TestMetadata("nonOverrideEquals.kt")
public void testNonOverrideEquals() throws Exception {
runTest("testData/inspectionsLocal/arrayInDataClass/nonOverrideEquals.kt");
}
@TestMetadata("nonOverrideHashCode.kt")
public void testNonOverrideHashCode() throws Exception {
runTest("testData/inspectionsLocal/arrayInDataClass/nonOverrideHashCode.kt");
}
}
@@ -1 +1,2 @@
// INSPECTION_CLASS: org.jetbrains.kotlin.idea.inspections.ArrayInDataClassInspection
// INSPECTION_CLASS: org.jetbrains.kotlin.idea.inspections.ArrayInDataClassInspection
// K2_INSPECTION_CLASS: org.jetbrains.kotlin.idea.k2.codeinsight.inspections.declarations.ArrayInDataClassInspection
@@ -0,0 +1 @@
org.jetbrains.kotlin.idea.k2.codeinsight.inspections.declarations.ArrayInDataClassInspection
@@ -0,0 +1,3 @@
// WITH_STDLIB
data class A(<caret>val a: Array<String>)
@@ -0,0 +1,18 @@
// WITH_STDLIB
data class A(val a: Array<String>) {
override fun equals(other: Any?): Boolean {
if (this === other) return true
if (javaClass != other?.javaClass) return false
other as A
if (!a.contentEquals(other.a)) return false
return true
}
override fun hashCode(): Int {
return a.contentHashCode()
}
}
@@ -0,0 +1,3 @@
// WITH_STDLIB
data class A(<caret>val a: IntArray)
@@ -15,4 +15,4 @@ data class A(val a: IntArray) {
override fun hashCode(): Int {
return a.contentHashCode()
}
}
}
@@ -0,0 +1,14 @@
// WITH_STDLIB
data class A(val <caret>a: IntArray) {
override fun equals(other: Any?): Boolean {
if (this === other) return true
if (javaClass != other?.javaClass) return false
other as A
if (!a.contentEquals(other.a)) return false
return true
}
}
@@ -0,0 +1,18 @@
// WITH_STDLIB
data class A(val a: IntArray) {
override fun equals(other: Any?): Boolean {
if (this === other) return true
if (javaClass != other?.javaClass) return false
other as A
if (!a.contentEquals(other.a)) return false
return true
}
override fun hashCode(): Int {
return a.contentHashCode()
}
}
@@ -0,0 +1,7 @@
// WITH_STDLIB
data class A(val <caret>a: IntArray) {
override fun hashCode(): Int {
return a.contentHashCode()
}
}
@@ -0,0 +1,18 @@
// WITH_STDLIB
data class A(val a: IntArray) {
override fun hashCode(): Int {
return a.contentHashCode()
}
override fun equals(other: Any?): Boolean {
if (this === other) return true
if (javaClass != other?.javaClass) return false
other as A
if (!a.contentEquals(other.a)) return false
return true
}
}
@@ -0,0 +1,11 @@
// WITH_STDLIB
class MyClass
data class A(
val a: <caret>IntArray,
val b: Array<String>,
val c: String,
val d: Int,
val e: MyClass,
)
@@ -0,0 +1,35 @@
// WITH_STDLIB
class MyClass
data class A(
val a: IntArray,
val b: Array<String>,
val c: String,
val d: Int,
val e: MyClass,
) {
override fun equals(other: Any?): Boolean {
if (this === other) return true
if (javaClass != other?.javaClass) return false
other as A
if (!a.contentEquals(other.a)) return false
if (!b.contentEquals(other.b)) return false
if (c != other.c) return false
if (d != other.d) return false
if (e != other.e) return false
return true
}
override fun hashCode(): Int {
var result = a.contentHashCode()
result = 31 * result + b.contentHashCode()
result = 31 * result + c.hashCode()
result = 31 * result + d
result = 31 * result + e.hashCode()
return result
}
}
@@ -0,0 +1,19 @@
// WITH_STDLIB
// PROBLEM: none
data class A(val <caret>a: IntArray) {
override fun equals(other: Any?): Boolean {
if (this === other) return true
if (javaClass != other?.javaClass) return false
other as A
if (!a.contentEquals(other.a)) return false
return true
}
override fun hashCode(): Int {
return a.contentHashCode()
}
}
@@ -0,0 +1,4 @@
// WITH_STDLIB
// PROBLEM: none
data class A(val <caret>str: String)
@@ -0,0 +1,11 @@
// WITH_STDLIB
data class A(val <caret>a: IntArray) {
fun equals(other: Any?, excessive: Any?): Boolean {
return true
}
override fun hashCode(): Int {
return a.contentHashCode()
}
}
@@ -0,0 +1,22 @@
// WITH_STDLIB
data class A(val a: IntArray) {
fun equals(other: Any?, excessive: Any?): Boolean {
return true
}
override fun hashCode(): Int {
return a.contentHashCode()
}
override fun equals(other: Any?): Boolean {
if (this === other) return true
if (javaClass != other?.javaClass) return false
other as A
if (!a.contentEquals(other.a)) return false
return true
}
}
@@ -0,0 +1,18 @@
// WITH_STDLIB
data class A(val <caret>a: IntArray) {
override fun equals(other: Any?): Boolean {
if (this === other) return true
if (javaClass != other?.javaClass) return false
other as A
if (!a.contentEquals(other.a)) return false
return true
}
fun hashCode(seed: Int): Int {
return 42
}
}
@@ -0,0 +1,22 @@
// WITH_STDLIB
data class A(val a: IntArray) {
override fun equals(other: Any?): Boolean {
if (this === other) return true
if (javaClass != other?.javaClass) return false
other as A
if (!a.contentEquals(other.a)) return false
return true
}
fun hashCode(seed: Int): Int {
return 42
}
override fun hashCode(): Int {
return a.contentHashCode()
}
}
@@ -1,3 +0,0 @@
// WITH_STDLIB
data class A(<caret>val a: IntArray)
@@ -58,6 +58,7 @@ internal fun MutableTWorkspace.generateK2InspectionTests() {
model("${idea}/inspectionsLocal/usePropertyAccessSyntax")
model("${idea}/inspectionsLocal/redundantUnitReturnType")
model("${idea}/inspectionsLocal/canBeParameter")
model("${idea}/inspectionsLocal/arrayInDataClass")
model("code-insight/inspections-k2/tests/testData/inspectionsLocal", pattern = pattern)
}
/**
@@ -79,6 +80,7 @@ internal fun MutableTWorkspace.generateK2InspectionTests() {
model("${idea}/inspections/protectedInFinal", pattern = pattern)
model("${idea}/intentions/convertToStringTemplate", pattern = pattern)
model("${idea}/inspections/unusedSymbol", pattern = pattern)
model("${idea}/inspections/arrayInDataClass", pattern = pattern)
}
testClass<AbstractK2MultiFileInspectionTest> {