switch expression type: initial

This commit is contained in:
Anna.Kozlova
2018-11-19 11:25:18 +01:00
parent 41b25091e7
commit c3bc7e3461
16 changed files with 332 additions and 14 deletions
@@ -18,6 +18,7 @@ public class PsiPolyExpressionUtil {
return !(expression instanceof PsiFunctionalExpression) &&
!(expression instanceof PsiParenthesizedExpression) &&
!(expression instanceof PsiConditionalExpression) &&
!(expression instanceof PsiSwitchExpression) &&
!(expression instanceof PsiCallExpression);
}
@@ -44,6 +45,9 @@ public class PsiPolyExpressionUtil {
return isInAssignmentOrInvocationContext(expression);
}
}
if (expression instanceof PsiSwitchExpression) {
return isInAssignmentOrInvocationContext(expression);
}
return false;
}
@@ -159,6 +163,9 @@ public class PsiPolyExpressionUtil {
final PsiMethod method = ((PsiMethodCallExpression)arg).resolveMethod();
return method != null && method.getReturnType() instanceof PsiPrimitiveType;
}
else if (arg instanceof PsiSwitchExpression) {
return isBooleanOrNumeric(arg) != null;
}
else {
assert false : arg;
return false;
@@ -201,6 +208,25 @@ public class PsiPolyExpressionUtil {
if (thenKind == elseKind || elseKind == ConditionalKind.NULL) return thenKind;
if (thenKind == ConditionalKind.NULL) return elseKind;
}
if (expr instanceof PsiSwitchExpression) {
ConditionalKind switchKind = null;
for (PsiExpression resultExpression : PsiUtil.getSwitchResultExpressions((PsiSwitchExpression)expr)) {
ConditionalKind resultKind = isBooleanOrNumeric(resultExpression);
if (resultKind == null) return null;
if (switchKind == null) {
switchKind = resultKind;
}
else if (switchKind != resultKind) {
if (switchKind == ConditionalKind.NULL) {
switchKind = resultKind;
}
else if (resultKind != ConditionalKind.NULL) {
return null;
}
}
}
}
return null;
}
@@ -33,6 +33,8 @@ import org.jetbrains.annotations.NotNull;
import org.jetbrains.annotations.Nullable;
import java.util.Map;
import java.util.Set;
import java.util.stream.Collectors;
/**
* @author ik, dsl
@@ -256,6 +258,12 @@ public class MethodCandidateInfo extends CandidateInfo{
return ThreeState.UNSURE;
}
}
else if (expression instanceof PsiSwitchExpression) {
Set<ThreeState> states =
PsiUtil.getSwitchResultExpressions((PsiSwitchExpression)expression).stream().map(expr -> isPotentialCompatible(expr, formalType, method)).collect(Collectors.toSet());
if (states.contains(ThreeState.NO)) return ThreeState.NO;
if (states.contains(ThreeState.UNSURE)) return ThreeState.UNSURE;
}
return ThreeState.YES;
}
@@ -347,6 +347,25 @@ public final class PsiUtil extends PsiUtilCore {
return false;
}
public static List<PsiExpression> getSwitchResultExpressions(PsiSwitchExpression switchExpression) {
PsiCodeBlock body = switchExpression.getBody();
if (body != null) {
List<PsiExpression> result = new ArrayList<>();
PsiStatement[] statements = body.getStatements();
for (PsiStatement statement : statements) {
if (statement instanceof PsiSwitchLabeledRuleStatement) {
PsiStatement ruleBody = ((PsiSwitchLabeledRuleStatement)statement).getBody();
if (ruleBody instanceof PsiExpressionStatement) {//todo break statements
result.add(((PsiExpressionStatement)ruleBody).getExpression());
}
}
}
return result;
}
return Collections.emptyList();
}
@MagicConstant(intValues = {ACCESS_LEVEL_PUBLIC, ACCESS_LEVEL_PROTECTED, ACCESS_LEVEL_PACKAGE_LOCAL, ACCESS_LEVEL_PRIVATE})
public @interface AccessLevel {}
@@ -40,8 +40,8 @@ public class TypeConversionUtil {
public static final int CHAR_RANK = 3;
public static final int INT_RANK = 4;
public static final int LONG_RANK = 5;
private static final int FLOAT_RANK = 6;
private static final int DOUBLE_RANK = 7;
public static final int FLOAT_RANK = 6;
public static final int DOUBLE_RANK = 7;
private static final int BOOL_RANK = 10;
private static final int STRING_RANK = 100;
private static final int MAX_NUMERIC_RANK = DOUBLE_RANK;
@@ -528,6 +528,11 @@ public class InferenceSession {
processReturnExpression(additionalConstraints, ignoredConstraints, ((PsiConditionalExpression)returnExpression).getThenExpression(), functionalType, addConstraint, initialSubstitutor);
processReturnExpression(additionalConstraints, ignoredConstraints, ((PsiConditionalExpression)returnExpression).getElseExpression(), functionalType, addConstraint, initialSubstitutor);
}
else if (returnExpression instanceof PsiSwitchExpression) {
for (PsiExpression resultExpression : PsiUtil.getSwitchResultExpressions((PsiSwitchExpression)returnExpression)) {
processReturnExpression(additionalConstraints, ignoredConstraints, resultExpression, functionalType, addConstraint, initialSubstitutor);
}
}
else if (returnExpression instanceof PsiLambdaExpression) {
collectLambdaReturnExpression(additionalConstraints, ignoredConstraints, (PsiLambdaExpression)returnExpression, functionalType, myErased, initialSubstitutor);
}
@@ -796,9 +801,13 @@ public class InferenceSession {
}
return false;
}
public static PsiType getTargetType(final PsiElement context) {
return getTargetTypeFromParent(context, new Ref<>(), true);
PsiType targetType = getTargetTypeFromParent(context, new Ref<>(), true);
if (targetType instanceof PsiClassType) {
return ((PsiClassType)targetType).setLanguageLevel(PsiUtil.getLanguageLevel(context));
}
return targetType;
}
/**
@@ -1887,6 +1896,11 @@ public class InferenceSession {
return argConstraints(thenExpression, session, sInterfaceMethod, sSubstitutor, tInterfaceMethod, tSubstitutor) &&
argConstraints(elseExpression, session, sInterfaceMethod, sSubstitutor, tInterfaceMethod, tSubstitutor);
}
if (arg instanceof PsiSwitchExpression) {
return PsiUtil.getSwitchResultExpressions((PsiSwitchExpression)arg).stream()
.allMatch(resultExpression -> argConstraints(resultExpression, session, sInterfaceMethod, sSubstitutor, tInterfaceMethod, tSubstitutor));
}
return false;
}
@@ -8,6 +8,7 @@ import com.intellij.psi.impl.source.resolve.graphInference.InferenceSession;
import com.intellij.psi.impl.source.resolve.graphInference.InferenceVariable;
import com.intellij.psi.impl.source.resolve.graphInference.PsiPolyExpressionUtil;
import com.intellij.psi.infos.MethodCandidateInfo;
import com.intellij.psi.util.PsiUtil;
import com.intellij.psi.util.TypeConversionUtil;
import com.intellij.util.ArrayUtil;
import org.jetbrains.annotations.NotNull;
@@ -82,6 +83,11 @@ public class ExpressionCompatibilityConstraint extends InputOutputConstraintForm
return true;
}
if (myExpression instanceof PsiSwitchExpression) {
PsiUtil.getSwitchResultExpressions((PsiSwitchExpression)myExpression).forEach(expression -> constraints.add(new ExpressionCompatibilityConstraint(expression,myT)));
return true;
}
if (myExpression instanceof PsiCall) {
final InferenceSession callSession = reduceExpressionCompatibilityConstraint(session, myExpression, myT, true);
if (callSession == null) {
@@ -24,6 +24,8 @@ import org.jetbrains.annotations.Nullable;
import java.util.HashSet;
import java.util.Set;
import java.util.stream.Collectors;
import java.util.stream.Stream;
public abstract class InputOutputConstraintFormula implements ConstraintFormula {
@@ -93,6 +95,15 @@ public abstract class InputOutputConstraintFormula implements ConstraintFormula
return thenResult;
}
}
if (psiExpression instanceof PsiSwitchExpression) {
Set<InferenceVariable> variables =
PsiUtil.getSwitchResultExpressions((PsiSwitchExpression)psiExpression).stream().flatMap(expression -> {
Set<InferenceVariable> inputVariables = createSelfConstraint(type, expression).getInputVariables(session);
return inputVariables != null ? inputVariables.stream() : Stream.empty();
}).collect(Collectors.toSet());
return variables.isEmpty() ? null : variables;
}
return null;
}
@@ -71,11 +71,7 @@ public class PsiConditionalExpressionImpl extends ExpressionPsiElement implement
!MethodCandidateInfo.ourOverloadGuard.currentStack().contains(PsiUtil.skipParenthesizedExprUp(this.getParent()))) {
//15.25.3 Reference Conditional Expressions
// The type of a poly reference conditional expression is the same as its target type.
final PsiType targetType = InferenceSession.getTargetType(this);
if (targetType instanceof PsiClassType) {
return ((PsiClassType)targetType).setLanguageLevel(PsiUtil.getLanguageLevel(this));
}
return targetType;
return InferenceSession.getTargetType(this);
}
final int typeRank1 = TypeConversionUtil.getTypeRank(type1);
@@ -3,13 +3,23 @@ package com.intellij.psi.impl.source.tree.java;
import com.intellij.lang.ASTNode;
import com.intellij.psi.*;
import com.intellij.psi.impl.source.PsiImmediateClassType;
import com.intellij.psi.impl.source.resolve.graphInference.InferenceSession;
import com.intellij.psi.impl.source.resolve.graphInference.PsiPolyExpressionUtil;
import com.intellij.psi.impl.source.tree.ElementType;
import com.intellij.psi.impl.source.tree.JavaElementType;
import com.intellij.psi.impl.source.tree.JavaSourceUtil;
import com.intellij.psi.impl.source.tree.TreeElement;
import com.intellij.psi.infos.MethodCandidateInfo;
import com.intellij.psi.util.PsiUtil;
import com.intellij.psi.util.TypeConversionUtil;
import com.intellij.util.ArrayUtil;
import com.intellij.util.containers.ContainerUtil;
import org.jetbrains.annotations.NotNull;
import java.util.HashSet;
import java.util.List;
import java.util.Set;
public class PsiSwitchExpressionImpl extends PsiSwitchBlockImpl implements PsiSwitchExpression {
public PsiSwitchExpressionImpl() {
super(JavaElementType.SWITCH_EXPRESSION);
@@ -22,9 +32,80 @@ public class PsiSwitchExpressionImpl extends PsiSwitchBlockImpl implements PsiSw
@Override
public PsiType getType() {
//todo[ann] http://cr.openjdk.java.net/~gbierman/switch-expressions.html#jep325-15.29.1
PsiClass objClass = JavaPsiFacade.getInstance(getProject()).findClass(CommonClassNames.JAVA_LANG_OBJECT, getResolveScope());
return objClass != null ? new PsiImmediateClassType(objClass, PsiSubstitutor.EMPTY) : null;
if (PsiPolyExpressionUtil.isPolyExpression(this) &&
!MethodCandidateInfo.ourOverloadGuard.currentStack().contains(PsiUtil.skipParenthesizedExprUp(getParent()))) {
return InferenceSession.getTargetType(this);
}
List<PsiExpression> resultExpressions = PsiUtil.getSwitchResultExpressions(this);
Set<PsiType> resultTypes = new HashSet<>();
for (PsiExpression expression : resultExpressions) {
PsiType resultExpressionType = expression.getType();
if (resultExpressionType == null) return null;
resultTypes.add(resultExpressionType);
}
//If the result expressions all have the same type (which may be the null type), then that is the type of the switch expression.
if (resultTypes.size() == 1) {
return ContainerUtil.getFirstItem(resultTypes);
}
//Otherwise, if the type of each result expression is boolean or Boolean,
//an unboxing conversion (5.1.8) is applied to each result expression of type Boolean, and the switch expression has type boolean.
if (resultTypes.stream().allMatch(type -> PsiType.BOOLEAN.isAssignableFrom(type))) {
return PsiType.BOOLEAN;
}
//Otherwise, if the type of each result expression is convertible to a numeric type (5.1.8),
// the type of the switch expression is given by numeric promotion (5.6.3) applied to the result expressions.
int[] ranks = resultTypes.stream().mapToInt(type -> TypeConversionUtil.getTypeRank(type)).toArray();
int maxRank = ArrayUtil.max(ranks);
if (TypeConversionUtil.isNumericType(maxRank)) {
if (maxRank == TypeConversionUtil.DOUBLE_RANK) {
return PsiType.DOUBLE;
}
if (maxRank == TypeConversionUtil.FLOAT_RANK) {
return PsiType.FLOAT;
}
if (maxRank == TypeConversionUtil.LONG_RANK) {
return PsiType.LONG;
}
if (isNumericPromotion(resultExpressions, ranks, PsiType.CHAR)) {
return PsiType.CHAR;
}
if (isNumericPromotion(resultExpressions, ranks, PsiType.SHORT)) {
return PsiType.SHORT;
}
if (isNumericPromotion(resultExpressions, ranks, PsiType.BYTE)) {
return PsiType.BYTE;
}
return PsiType.INT;
}
//Otherwise, boxing conversion (5.1.7) is applied to each result expression that has a primitive type, after which the type of the switch expression is the result of applying capture conversion (5.1.10)
// to the least upper bound (4.10.4) of the types of the result expressions.
PsiType leastUpperBound = PsiType.NULL;
for (PsiType type : resultTypes) {
if (TypeConversionUtil.isPrimitiveAndNotNull(type)) {
type = ((PsiPrimitiveType)type).getBoxedType(this);
}
if (leastUpperBound == PsiType.NULL) {
leastUpperBound = type;
}
else {
leastUpperBound = GenericsUtil.getLeastUpperBound(leastUpperBound, leastUpperBound, getManager());
}
}
return leastUpperBound != null ? PsiUtil.captureToplevelWildcards(leastUpperBound, this) : null;
}
private static boolean isNumericPromotion(List<PsiExpression> resultExpressions, int[] ranks, final PsiPrimitiveType type) {
return ArrayUtil.find(ranks, TypeConversionUtil.getTypeRank(type)) > -1 &&
resultExpressions.stream().allMatch(expression -> TypeConversionUtil.areTypesAssignmentCompatible(type, expression));
}
@Override
@@ -749,6 +749,10 @@ public class JavaMethodsConflictResolver implements PsiConflictResolver{
isFunctionalTypeMoreSpecific(((PsiConditionalExpression)expr).getElseExpression(), sType, tType);
}
if (expr instanceof PsiSwitchExpression) {
return PsiUtil.getSwitchResultExpressions((PsiSwitchExpression)expr).stream().allMatch(resultExpr -> isFunctionalTypeMoreSpecific(resultExpr, sType, tType));
}
if (expr instanceof PsiFunctionalExpression) {
if (expr instanceof PsiLambdaExpression && !((PsiLambdaExpression)expr).hasFormalParameterTypes()) {
@@ -0,0 +1,11 @@
class MyTest {
<T> T foo(T t) {
return t;
}
void m(int i) {
String s = foo(switch (i) {default -> "str";});
String s1 = <error descr="Incompatible types. Required String but 'foo' was inferred to T:
no instance(s) of type variable(s) exist so that Object conforms to String">foo(switch (i) {case 1 -> new Object(); default -> "str";});</error>
}
}
@@ -0,0 +1,49 @@
class SwitchExpressions {
byte B = 1;
short S = 1;
char C = 1;
final int I = 1;
void m(int i) {
var v1 = switch(i) {
case 1 -> 1;
default -> 1.0;
};
double d = v1;
<error descr="Incompatible types. Found: 'double', required: 'float'">float f = v1;</error>
<error descr="Incompatible types. Found: 'double', required: 'int'">int in = v1;</error>
var v2 = switch (i) {
case 1 -> C;
default -> I;
};
in = v2;
<error descr="Incompatible types. Found: 'char', required: 'byte'">byte b = v2;</error>
<error descr="Incompatible types. Found: 'char', required: 'short'">short s = v2;</error>
char c = v2;
var v3 = switch (i) {
case 1 -> B;
default -> I;
};
in = v3;
b = v3;
s = v3;
<error descr="Incompatible types. Found: 'byte', required: 'char'">c = v3</error>;
var v4 = switch (i) {
case 1 -> B;
default -> Integer.MAX_VALUE;
};
in = v4;
<error descr="Incompatible types. Found: 'int', required: 'byte'">b = v4</error>;
<error descr="Incompatible types. Found: 'int', required: 'short'">s = v4</error>;
<error descr="Incompatible types. Found: 'int', required: 'char'">c = v4</error>;
}
}
@@ -10,6 +10,8 @@ class LightJava12HighlightingTest : LightCodeInsightFixtureTestCase() {
fun testEnhancedSwitchStatements() = doTest()
fun testSwitchExpressions() = doTest()
fun testSwitchNumericPromotion() = doTest()
fun testSimpleInferenceCases() = doTest()
private fun doTest() {
myFixture.configureByFile(getTestName(false) + ".java")
@@ -0,0 +1,82 @@
// Copyright 2000-2018 JetBrains s.r.o. Use of this source code is governed by the Apache 2.0 license that can be found in the LICENSE file.
package com.intellij.java.propertyBased;
import com.intellij.openapi.application.PathManager;
import com.intellij.openapi.projectRoots.impl.JavaAwareProjectJdkTableImpl;
import com.intellij.openapi.util.io.FileUtil;
import com.intellij.psi.PsiFile;
import com.intellij.psi.PsiSwitchBlock;
import com.intellij.psi.util.PsiTreeUtil;
import com.intellij.testFramework.LightProjectDescriptor;
import com.intellij.testFramework.SkipSlowTestLocally;
import com.intellij.testFramework.fixtures.LightCodeInsightFixtureTestCase;
import com.intellij.testFramework.propertyBased.InvokeIntention;
import com.intellij.testFramework.propertyBased.MadTestingAction;
import com.intellij.testFramework.propertyBased.MadTestingUtil;
import com.intellij.testFramework.propertyBased.StripTestDataMarkup;
import com.intellij.util.containers.ContainerUtil;
import org.jetbrains.annotations.NotNull;
import org.jetbrains.annotations.Nullable;
import org.jetbrains.jetCheck.Generator;
import org.jetbrains.jetCheck.PropertyChecker;
import java.io.IOException;
import java.util.Collection;
import java.util.List;
import java.util.function.Function;
import java.util.function.Supplier;
@SkipSlowTestLocally
public class Java12SwitchExpressionSanityTest extends LightCodeInsightFixtureTestCase {
@Override
protected void tearDown() throws Exception {
// remove jdk if it was created during highlighting to avoid leaks
try {
JavaAwareProjectJdkTableImpl.removeInternalJdkInTests();
}
finally {
super.tearDown();
}
}
@NotNull
@Override
protected LightProjectDescriptor getProjectDescriptor() {
return JAVA_12;
}
public void testIntentionsAroundSwitch() {
MadTestingUtil.enableAllInspections(getProject(), getTestRootDisposable(),
"BoundedWildcard" // IDEA-194460
);
Function<PsiFile, Generator<? extends MadTestingAction>> fileActions =
file -> Generator.sampledFrom(new InvokeIntention(file, new JavaIntentionPolicy()) {
@Override
protected int generateDocOffset(@NotNull Environment env, @Nullable String logMessage) {
Collection<PsiSwitchBlock> children = PsiTreeUtil.findChildrenOfType(getFile(), PsiSwitchBlock.class);
if (children.isEmpty()) {
return super.generateDocOffset(env, logMessage);
}
List<Generator<Integer>> generators =
ContainerUtil.map(children, stmt ->
Generator.integers(stmt.getTextRange().getStartOffset(),
stmt.getTextRange().getEndOffset()).noShrink());
return env.generateValue(Generator.anyOf(generators), logMessage);
}
}, new StripTestDataMarkup(file));
Supplier<MadTestingAction> fileChooser = MadTestingUtil.actionsOnFileContents(myFixture, PathManager.getHomePath(), f -> {
try {
return f.getName().endsWith(".java") && FileUtil.loadFile(f).contains(" switch");
}
catch (IOException e) {
return false;
}
}, fileActions);
PropertyChecker.checkScenarios(fileChooser);
}
}
@@ -80,7 +80,7 @@ public class InvokeIntention extends ActionOnFile {
return result;
}
private void doInvokeIntention(int offset, Environment env) {
protected void doInvokeIntention(int offset, Environment env) {
Project project = getProject();
Editor editor = FileEditorManager.getInstance(project).openTextEditor(new OpenFileDescriptor(project, getVirtualFile(), offset), true);
assert editor != null;
@@ -932,6 +932,15 @@ public class ArrayUtil extends ArrayUtilRt {
return min;
}
@Contract(pure = true)
public static int max(int[] values) {
int max = Integer.MIN_VALUE;
for (int value : values) {
if (value > max) max = value;
}
return max;
}
@Contract(pure = true)
public static int[] mergeSortedArrays(int[] a1, int[] a2, boolean mergeEqualItems) {
int newSize = a1.length + a2.length;