diff --git a/plugins/junit/src/com/intellij/execution/junit/JUnitConfigurationProducer.java b/plugins/junit/src/com/intellij/execution/junit/JUnitConfigurationProducer.java index 1924f5438041..91ea90de42a9 100644 --- a/plugins/junit/src/com/intellij/execution/junit/JUnitConfigurationProducer.java +++ b/plugins/junit/src/com/intellij/execution/junit/JUnitConfigurationProducer.java @@ -48,7 +48,7 @@ public abstract class JUnitConfigurationProducer extends JavaRuntimeConfiguratio @NotNull RunnerAndConfigurationSettings[] existingConfigurations, ConfigurationContext context) { final PsiElement[] elements = LangDataKeys.PSI_ELEMENT_ARRAY.getData(context.getDataContext()); - if (elements != null && elements.length > 1) { + if (elements != null && PatternConfigurationProducer.collectTestClasses(elements).size() > 1) { return null; } final Module predefinedModule = diff --git a/plugins/junit/src/com/intellij/execution/junit/PatternConfigurationProducer.java b/plugins/junit/src/com/intellij/execution/junit/PatternConfigurationProducer.java index 0648114b4d96..c8b12081a8c1 100644 --- a/plugins/junit/src/com/intellij/execution/junit/PatternConfigurationProducer.java +++ b/plugins/junit/src/com/intellij/execution/junit/PatternConfigurationProducer.java @@ -20,12 +20,11 @@ import com.intellij.execution.Location; import com.intellij.execution.RunConfigurationExtension; import com.intellij.execution.RunnerAndConfigurationSettings; import com.intellij.execution.actions.ConfigurationContext; +import com.intellij.openapi.actionSystem.DataContext; import com.intellij.openapi.actionSystem.LangDataKeys; import com.intellij.openapi.project.Project; import com.intellij.openapi.util.Comparing; -import com.intellij.psi.PsiClass; -import com.intellij.psi.PsiClassOwner; -import com.intellij.psi.PsiElement; +import com.intellij.psi.*; import org.jetbrains.annotations.NotNull; import java.util.LinkedHashSet; @@ -40,7 +39,7 @@ public class PatternConfigurationProducer extends JUnitConfigurationProducer { final Project project = location.getProject(); final LinkedHashSet classes = new LinkedHashSet(); myElements = collectPatternElements(context, classes); - if (classes.isEmpty()) return null; + if (classes.size() <= 1) return null; RunnerAndConfigurationSettings settings = cloneTemplateConfiguration(project, context); final JUnitConfiguration configuration = (JUnitConfiguration)settings.getConfiguration(); final JUnitConfiguration.Data data = configuration.getPersistentData(); @@ -72,12 +71,21 @@ public class PatternConfigurationProducer extends JUnitConfigurationProducer { } private static PsiElement[] collectPatternElements(ConfigurationContext context, LinkedHashSet classes) { - PsiElement[] elements = LangDataKeys.PSI_ELEMENT_ARRAY.getData(context.getDataContext()); - if (elements != null && elements.length > 1) { + final DataContext dataContext = context.getDataContext(); + PsiElement[] elements = LangDataKeys.PSI_ELEMENT_ARRAY.getData(dataContext); + if (elements != null) { for (PsiClass psiClass : collectTestClasses(elements)) { classes.add(psiClass.getQualifiedName()); } return elements; + } else { + final PsiFile file = LangDataKeys.PSI_FILE.getData(dataContext); + if (file instanceof PsiClassOwner) { + for (PsiClass psiClass : collectTestClasses(((PsiClassOwner)file).getClasses())) { + classes.add(psiClass.getQualifiedName()); + } + return new PsiElement[]{file}; + } } return null; } diff --git a/plugins/junit/src/com/intellij/execution/junit/TestClassConfigurationProducer.java b/plugins/junit/src/com/intellij/execution/junit/TestClassConfigurationProducer.java index fb9bc19a597a..d2edb10de567 100644 --- a/plugins/junit/src/com/intellij/execution/junit/TestClassConfigurationProducer.java +++ b/plugins/junit/src/com/intellij/execution/junit/TestClassConfigurationProducer.java @@ -35,7 +35,7 @@ public class TestClassConfigurationProducer extends JUnitConfigurationProducer { final Project project = location.getProject(); final PsiElement[] elements = LangDataKeys.PSI_ELEMENT_ARRAY.getData(context.getDataContext()); - if (elements != null && elements.length > 1) { + if (elements != null && PatternConfigurationProducer.collectTestClasses(elements).size() > 1) { return null; } myTestClass = JUnitUtil.getTestClass(location);