diff --git a/java/java-psi-api/src/com/intellij/psi/util/PsiUtil.java b/java/java-psi-api/src/com/intellij/psi/util/PsiUtil.java index e16ab4cf3e7f..bdc9cb58c853 100644 --- a/java/java-psi-api/src/com/intellij/psi/util/PsiUtil.java +++ b/java/java-psi-api/src/com/intellij/psi/util/PsiUtil.java @@ -36,6 +36,7 @@ import org.jetbrains.annotations.NotNull; import org.jetbrains.annotations.Nullable; import java.util.*; +import java.util.function.Predicate; public final class PsiUtil extends PsiUtilCore { private static final Logger LOG = Logger.getInstance("#com.intellij.psi.util.PsiUtil"); @@ -359,20 +360,27 @@ public final class PsiUtil extends PsiUtilCore { result.add(((PsiExpressionStatement)ruleBody).getExpression()); } else if (ruleBody instanceof PsiBlockStatement) { - PsiCodeBlock codeBlock = ((PsiBlockStatement)ruleBody).getCodeBlock(); - ArrayList breaks = new ArrayList<>(); - addStatements(breaks, codeBlock, PsiBreakStatement.class); - for (PsiBreakStatement aBreak : breaks) { - ContainerUtil.addIfNotNull(result, aBreak.getExpression()); - } + collectSwitchResultExpressions(result, (PsiBlockStatement)ruleBody); } } + else if (statement instanceof PsiBlockStatement) { + collectSwitchResultExpressions(result, (PsiBlockStatement)statement); + } } return result; } return Collections.emptyList(); } + private static void collectSwitchResultExpressions(List result, PsiBlockStatement ruleBody) { + PsiCodeBlock codeBlock = ruleBody.getCodeBlock(); + ArrayList breaks = new ArrayList<>(); + addStatements(breaks, codeBlock, PsiBreakStatement.class, element -> element instanceof PsiSwitchBlock); + for (PsiBreakStatement aBreak : breaks) { + ContainerUtil.addIfNotNull(result, aBreak.getExpression()); + } + } + @MagicConstant(intValues = {ACCESS_LEVEL_PUBLIC, ACCESS_LEVEL_PROTECTED, ACCESS_LEVEL_PACKAGE_LOCAL, ACCESS_LEVEL_PRIVATE}) public @interface AccessLevel {} @@ -1326,20 +1334,20 @@ public final class PsiUtil extends PsiUtilCore { public static PsiReturnStatement[] findReturnStatements(@Nullable PsiCodeBlock body) { ArrayList vector = new ArrayList<>(); if (body != null) { - addStatements(vector, body, PsiReturnStatement.class); + addStatements(vector, body, PsiReturnStatement.class, statement -> false); } return vector.toArray(PsiReturnStatement.EMPTY_ARRAY); } - private static void addStatements(List vector, PsiElement element, Class clazz) { + private static void addStatements(List vector, PsiElement element, Class clazz, Predicate stopAt) { if (PsiTreeUtil.instanceOf(element, clazz)) { //noinspection unchecked vector.add((T)element); } - else if (!(element instanceof PsiClass) && !(element instanceof PsiLambdaExpression)) { + else if (!(element instanceof PsiClass) && !(element instanceof PsiLambdaExpression) && !stopAt.test(element)) { PsiElement[] children = element.getChildren(); for (PsiElement child : children) { - addStatements(vector, child, clazz); + addStatements(vector, child, clazz, stopAt); } } } diff --git a/java/java-psi-impl/src/com/intellij/psi/impl/source/resolve/graphInference/InferenceSession.java b/java/java-psi-impl/src/com/intellij/psi/impl/source/resolve/graphInference/InferenceSession.java index 9c96e34b2a59..a2bfb596df34 100644 --- a/java/java-psi-impl/src/com/intellij/psi/impl/source/resolve/graphInference/InferenceSession.java +++ b/java/java-psi-impl/src/com/intellij/psi/impl/source/resolve/graphInference/InferenceSession.java @@ -857,6 +857,10 @@ public class InferenceSession { return getTargetTypeFromParentLambda(PsiTreeUtil.getParentOfType(parent, PsiLambdaExpression.class, true, PsiMethod.class), errorMessage, inferParent); } + PsiSwitchExpression switchExpression = PsiTreeUtil.getParentOfType(parent, PsiSwitchExpression.class); + if (switchExpression != null && PsiUtil.getSwitchResultExpressions(switchExpression).contains(context)) { + return getTargetType(switchExpression); + } return null; } diff --git a/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/advHighlighting12/SimpleInferenceCases.java b/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/advHighlighting12/SimpleInferenceCases.java index 8ab10d4e536c..73f0f8cf992e 100644 --- a/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/advHighlighting12/SimpleInferenceCases.java +++ b/java/java-tests/testData/codeInsight/daemonCodeAnalyzer/advHighlighting12/SimpleInferenceCases.java @@ -37,5 +37,47 @@ no instance(s) of type variable(s) exist so that Integer conforms to String">() break () -> 1; } }; + String s7 = switch (i) { + case 1: { + break 1; + } + default: { + int i1 = switch (0) { + default -> { + break 1; + } + }; + break ""; + } + }; + } + + void switchChain(final int i) { + String s = switch (i) { + default -> switch (0) { + default -> { + break 1; + } + }; + }; + String s1 = switch (i) { + default -> { + break switch (0) { + default -> { + break 1; + } + }; + } + }; + + String s2 = switch (i) { + default: { + break switch (0) { + default -> { + break 1; + } + }; + } + }; } } \ No newline at end of file