diff --git a/platform/lang-api/src/com/intellij/testIntegration/TestFinder.java b/platform/lang-api/src/com/intellij/testIntegration/TestFinder.java index d956d465398e..f62c477a5ab9 100644 --- a/platform/lang-api/src/com/intellij/testIntegration/TestFinder.java +++ b/platform/lang-api/src/com/intellij/testIntegration/TestFinder.java @@ -35,4 +35,9 @@ public interface TestFinder { Collection findClassesForTest(@NotNull PsiElement element); boolean isTest(@NotNull PsiElement element); + + @Nullable + default PsiElement findSelectedElement(@NotNull final PsiElement element) { + return element; + } } diff --git a/platform/lang-impl/src/com/intellij/testIntegration/TestFinderHelper.java b/platform/lang-impl/src/com/intellij/testIntegration/TestFinderHelper.java index 67318ceb1238..675cce74e813 100644 --- a/platform/lang-impl/src/com/intellij/testIntegration/TestFinderHelper.java +++ b/platform/lang-impl/src/com/intellij/testIntegration/TestFinderHelper.java @@ -39,7 +39,8 @@ public class TestFinderHelper { public static Collection findTestsForClass(PsiElement element) { Collection result = new LinkedHashSet<>(); for (TestFinder each : getFinders()) { - result.addAll(each.findTestsForClass(element)); + final PsiElement selectedElement = each.findSelectedElement(element); + if (selectedElement != null) result.addAll(each.findTestsForClass(selectedElement)); } return result; } @@ -47,7 +48,8 @@ public class TestFinderHelper { public static Collection findClassesForTest(PsiElement element) { Collection result = new LinkedHashSet<>(); for (TestFinder each : getFinders()) { - result.addAll(each.findClassesForTest(element)); + final PsiElement selectedElement = each.findSelectedElement(element); + if (selectedElement != null) result.addAll(each.findClassesForTest(selectedElement)); } return result; }