From f96c4918fbb9a63dac0df3d437f9edcf079bf70a Mon Sep 17 00:00:00 2001 From: Aleksey Dobrynin Date: Thu, 14 Aug 2025 11:23:56 +0200 Subject: [PATCH] [junit 5] collect only the methods used for @MethodSource IDEA-374530 IJ-CR-165719 GitOrigin-RevId: e1515d97a5f542e63ece9329340d682cdfba3e43 --- .../JvmReferenceContributorTestBase.kt | 17 +++ .../JavaJUnit5ReferenceContributorTest.kt | 100 +++++++++++++++++- .../JUnitMalformedDeclarationInspection.kt | 20 ++-- .../BaseJunitAnnotationReference.kt | 27 ++++- 4 files changed, 150 insertions(+), 14 deletions(-) diff --git a/jvm/jvm-analysis-testFramework/src/com/intellij/jvm/analysis/testFramework/JvmReferenceContributorTestBase.kt b/jvm/jvm-analysis-testFramework/src/com/intellij/jvm/analysis/testFramework/JvmReferenceContributorTestBase.kt index 6bd58c775df4..0515c42b3ae7 100644 --- a/jvm/jvm-analysis-testFramework/src/com/intellij/jvm/analysis/testFramework/JvmReferenceContributorTestBase.kt +++ b/jvm/jvm-analysis-testFramework/src/com/intellij/jvm/analysis/testFramework/JvmReferenceContributorTestBase.kt @@ -2,7 +2,9 @@ package com.intellij.jvm.analysis.testFramework import com.intellij.codeInsight.lookup.LookupElement import com.intellij.psi.PsiElement +import com.intellij.psi.PsiPolyVariantReference import com.intellij.psi.PsiReference +import com.intellij.psi.ResolveResult import com.intellij.testFramework.fixtures.JavaCodeInsightTestFixture abstract class JvmReferenceContributorTestBase : LightJvmCodeInsightFixtureTestCase() { @@ -20,6 +22,21 @@ abstract class JvmReferenceContributorTestBase : LightJvmCodeInsightFixtureTestC assertion(reference, resolved!!) } + protected fun JavaCodeInsightTestFixture.assertMultiresolveReference( + lang: JvmLanguage, + before: String, + fileName: String = generateFileName(), + assertion: (PsiReference, Array) -> Unit = { _, _ -> } + ) { + configureByText("$fileName${lang.ext}", before) + val offset = getCaretOffset() + val reference = myFixture.file.findReferenceAt(offset) ?: error("Could not find reference at caret offset") + assertTrue(reference is PsiPolyVariantReference) + if (reference !is PsiPolyVariantReference) return + + assertion(reference, reference.multiResolve(false)) + } + protected fun JavaCodeInsightTestFixture.assertUnResolvableReference( lang: JvmLanguage, before: String, diff --git a/plugins/junit/java-tests/test/com/intellij/execution/junit/JavaJUnit5ReferenceContributorTest.kt b/plugins/junit/java-tests/test/com/intellij/execution/junit/JavaJUnit5ReferenceContributorTest.kt index ac3c97bc1a2d..ccd7dd53fcca 100644 --- a/plugins/junit/java-tests/test/com/intellij/execution/junit/JavaJUnit5ReferenceContributorTest.kt +++ b/plugins/junit/java-tests/test/com/intellij/execution/junit/JavaJUnit5ReferenceContributorTest.kt @@ -3,6 +3,7 @@ package com.intellij.execution.junit import com.intellij.junit.testFramework.JUnit5ReferenceContributorTestBase import com.intellij.jvm.analysis.testFramework.JvmLanguage +import com.intellij.psi.PsiMethod class JavaJUnit5ReferenceContributorTest : JUnit5ReferenceContributorTestBase() { fun `test resolve to source method`() { @@ -18,8 +19,103 @@ class JavaJUnit5ReferenceContributorTest : JUnit5ReferenceContributorTestBase() private static void cde() {} } - """.trimIndent()) { reference, resolved -> - assertContainsElements(reference.lookupStringVariants(), "abc", "cde") + """.trimIndent()) { reference, _ -> + assertContainsElements(reference.lookupStringVariants(), "abc") + } + } + + fun `test filter resolved to source method`() { + myFixture.assertMultiresolveReference(JvmLanguage.JAVA, """ + import org.junit.jupiter.params.ParameterizedTest; + import org.junit.jupiter.params.provider.MethodSource; + + interface MyFirstInterface { + static void abc() {} + } + + interface MySecondInterface { + static void abc() {} + } + + abstract class MyAbstractClass implements MyFirstInterface, MySecondInterface {} + + class ParameterizedTestsDemo extends MyAbstractClass { + @MethodSource("abc") + void testWithProvider(String abc) {} + } + """.trimIndent()) { _, results -> + assertEquals(1, results.size) + val resolved = results.first().element + assertTrue(resolved is PsiMethod) + if (resolved !is PsiMethod) return@assertMultiresolveReference + assertEquals("MyFirstInterface", resolved.containingClass?.name) + assertEquals("abc", resolved.name) + } + } + + fun `test filter resolved to source method with inherited`() { + myFixture.assertMultiresolveReference(JvmLanguage.JAVA, """ + import org.junit.jupiter.params.ParameterizedTest; + import org.junit.jupiter.params.provider.MethodSource; + + interface MyFirstInterface { + static void abc() {} + } + + interface MySecondInterface { + static void abc() {} + } + + class ParameterizedTestsDemo implements MySecondInterface, MyFirstInterface { + @MethodSource("abc") + void testWithProvider(String abc) {} + } + + class ChildOfParameterizedTestsDemo extends ParameterizedTestsDemo { + static void abc() {} + } + """.trimIndent()) { _, results -> + val classes = results.map { (it.element as PsiMethod).containingClass?.name }.toSet() + assertEquals(setOf("MySecondInterface", "ChildOfParameterizedTestsDemo"), classes) + } + } + + fun `test filter resolved to source method with meta-annotation`() { + myFixture.assertMultiresolveReference(JvmLanguage.JAVA, """ + import org.junit.jupiter.params.ParameterizedTest; + import org.junit.jupiter.params.provider.MethodSource; + + interface MyFirstInterface { + static void abc() {} + } + + interface MySecondInterface { + static void abc() {} + } + + interface MyThirdInterface { + static void abc() {} + } + + @MethodSource("abc") + @interface MyAnnotation {} + + class ParameterizedTestsDemo1 implements MySecondInterface, MyFirstInterface { + @MyAnnotation + void testWithProvider(String abc) {} + } + + class ParameterizedTestsDemo2 implements MyThirdInterface, MyFirstInterface { + @MyAnnotation + void testWithProvider(String abc) {} + } + + class ChildOfParameterizedTestsDemo extends ParameterizedTestsDemo1 { + static void abc() {} + } + """.trimIndent()) { _, results -> + val classes = results.map { (it.element as PsiMethod).containingClass?.name }.toSet() + assertEquals(setOf("MySecondInterface", "MyThirdInterface", "ChildOfParameterizedTestsDemo"), classes) } } diff --git a/plugins/junit/src/com/intellij/execution/junit/codeInspection/JUnitMalformedDeclarationInspection.kt b/plugins/junit/src/com/intellij/execution/junit/codeInspection/JUnitMalformedDeclarationInspection.kt index 40a1247c66ec..9b4923bc57e7 100644 --- a/plugins/junit/src/com/intellij/execution/junit/codeInspection/JUnitMalformedDeclarationInspection.kt +++ b/plugins/junit/src/com/intellij/execution/junit/codeInspection/JUnitMalformedDeclarationInspection.kt @@ -21,11 +21,9 @@ import com.intellij.execution.junit.references.PsiSourceResolveResult import com.intellij.jvm.analysis.quickFix.CompositeModCommandQuickFix import com.intellij.jvm.analysis.quickFix.createModifierQuickfixes import com.intellij.lang.Language -import com.intellij.lang.java.request.CreateFieldFromJavaUsageRequest import com.intellij.lang.jvm.JvmMethod import com.intellij.lang.jvm.JvmModifier import com.intellij.lang.jvm.JvmModifiersOwner -import com.intellij.lang.jvm.JvmValue import com.intellij.lang.jvm.actions.* import com.intellij.lang.jvm.types.JvmPrimitiveTypeKind import com.intellij.lang.jvm.types.JvmType @@ -46,7 +44,6 @@ import com.intellij.psi.util.TypeConversionUtil import com.intellij.psi.util.parentOfType import com.intellij.uast.UastHintedVisitorAdapter import com.intellij.util.asSafely -import com.siyeh.ig.fixes.SerialVersionUIDBuilder import com.siyeh.ig.junit.JUnitCommonClassNames.* import com.siyeh.ig.psiutils.TestUtils import com.siyeh.ig.psiutils.TypeUtils @@ -391,8 +388,10 @@ private class JUnitMalformedSignatureVisitor( val javaClass = aClass.javaPsi if (aClass.isInterface || aClass.javaPsi.hasModifier(JvmModifier.ABSTRACT)) return val hasNestedAnnotation = javaClass.hasAnnotation(ORG_JUNIT_JUPITER_API_NESTED) - if (!hasNestedAnnotation && !aClass.methods.any { it.javaPsi.hasAnnotation(ORG_JUNIT_JUPITER_API_TEST) || - it.javaPsi.hasAnnotation(ORG_JUNIT_JUPITER_PARAMS_PARAMETERIZED_TEST)}) return + if (!hasNestedAnnotation && !aClass.methods.any { + it.javaPsi.hasAnnotation(ORG_JUNIT_JUPITER_API_TEST) || + it.javaPsi.hasAnnotation(ORG_JUNIT_JUPITER_PARAMS_PARAMETERIZED_TEST) + }) return if (!hasNestedAnnotation && aClass.isStatic) return if (hasNestedAnnotation && !aClass.isStatic && aClass.visibility != UastVisibility.PRIVATE) return val message = JUnitBundle.message("jvm.inspections.junit.malformed.missing.nested.annotation.descriptor") @@ -631,10 +630,11 @@ private class JUnitMalformedSignatureVisitor( } private fun PsiSourceResolveResult.getSourceForClass(owner: PsiClass): PsiElement? { - if(element is PsiMethod) { + if (element is PsiMethod) { if (owners.isEmpty()) return element // direct link return owner.findMethodBySignature(element as PsiMethod, true) - } else if (element is PsiField) { + } + else if (element is PsiField) { if (owners.isEmpty()) return element // direct link return owner.findFieldByName((element as PsiField).name, true) } @@ -642,7 +642,7 @@ private class JUnitMalformedSignatureVisitor( } private fun checkFieldSource(declaration: UDeclaration, methodSource: PsiAnnotation) { - if(declaration !is UMethod) return + if (declaration !is UMethod) return val psiMethod = declaration.javaPsi val containingClass = psiMethod.containingClass ?: return val annotationMemberValue = methodSource.flattenedAttributeValues(PsiAnnotation.DEFAULT_REFERENCED_METHOD_NAME) @@ -707,7 +707,7 @@ private class JUnitMalformedSignatureVisitor( } private fun checkAbsentFieldSourceProvider( - containingClass: PsiClass, anchor: PsiElement, sourceProviderName: String, method: UMethod + containingClass: PsiClass, anchor: PsiElement, sourceProviderName: String, method: UMethod, ) { val message = JUnitBundle.message( "jvm.inspections.junit.malformed.param.field.source.unresolved.descriptor", @@ -797,7 +797,7 @@ private class JUnitMalformedSignatureVisitor( ) { val actions = mutableListOf() val sameClass = sourceProvider.containingClass == containingClass - if(sameClass) { + if (sameClass) { val annotation = JavaPsiFacade.getElementFactory(containingClass.project).createAnnotationFromText( TEST_INSTANCE_PER_CLASS, containingClass ) diff --git a/plugins/junit/src/com/intellij/execution/junit/references/BaseJunitAnnotationReference.kt b/plugins/junit/src/com/intellij/execution/junit/references/BaseJunitAnnotationReference.kt index 13d4188b2344..a2610b98017e 100644 --- a/plugins/junit/src/com/intellij/execution/junit/references/BaseJunitAnnotationReference.kt +++ b/plugins/junit/src/com/intellij/execution/junit/references/BaseJunitAnnotationReference.kt @@ -108,15 +108,38 @@ abstract class BaseJunitAnnotationReference( return filteredElements(factoryMethods, scope.javaPsi, testMethod?.javaPsi).firstOrNull() } + /** + * Recursively finds a relevant `Psi` element from the provided list that is associated with the given class scope. + * priority: interface (by ordering) -> super class + * + * @param scope The class scope used to find the relevant `Psi` element. + * @param elements The list of `Psi` elements to be filtered. + * @return The relevant `Psi` element if found, or `null` if no matching element exists. + */ + private fun getRelevantElements(scope: PsiClass, elements: List): Psi? { + elements.find { it.containingClass == scope }?.let { return it } + for (iface in scope.interfaces) { + getRelevantElements(iface, elements)?.let { return it } + } + scope.superClass?.let { superClass -> + getRelevantElements(superClass, elements)?.let { return it } + } + return null + } + private fun fastResolveFor(literal: UExpression, scope: UClass, testMethod: UMethod?): Set { val name = literal.evaluate() as String? ?: return setOf() val currentTestClass = scope.javaPsi val clazzElements = filteredElements(getPsiElementsByName(currentTestClass, name, true), currentTestClass, testMethod?.javaPsi) + val relevantElement = getRelevantElements(scope.javaPsi, clazzElements) val elements = ClassInheritorsSearch.search(currentTestClass, currentTestClass.resolveScope, true) - .flatMap { inheritedTestClass -> filteredElements(getPsiElementsByName(inheritedTestClass, name, true), inheritedTestClass, testMethod?.javaPsi) } + .mapNotNull { inheritedTestClass -> + getRelevantElements(inheritedTestClass, + filteredElements(getPsiElementsByName(inheritedTestClass, name, true), + inheritedTestClass, testMethod?.javaPsi)) } .toMutableSet() - elements.addAll(clazzElements) + elements.addAll(if (relevantElement != null) listOf(relevantElement) else clazzElements) return elements }