[junit 5] collect only the methods used for @MethodSource IDEA-374530 IJ-CR-165719

GitOrigin-RevId: e1515d97a5f542e63ece9329340d682cdfba3e43
This commit is contained in:
Aleksey Dobrynin
2025-08-14 10:49:55 +00:00
committed by intellij-monorepo-bot
parent 2577a04dec
commit f96c4918fb
4 changed files with 150 additions and 14 deletions
@@ -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<ResolveResult>) -> 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,
@@ -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("ab<caret>c")
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("ab<caret>c")
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("ab<caret>c")
@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)
}
}
@@ -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<IntentionAction>()
val sameClass = sourceProvider.containingClass == containingClass
if(sameClass) {
if (sameClass) {
val annotation = JavaPsiFacade.getElementFactory(containingClass.project).createAnnotationFromText(
TEST_INSTANCE_PER_CLASS, containingClass
)
@@ -108,15 +108,38 @@ abstract class BaseJunitAnnotationReference<Psi : PsiMember, U : UDeclaration>(
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>): 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<PsiElement> {
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
}