diff --git a/core/src/main/java/com/google/errorprone/bugpatterns/formatstring/FormatStringValidation.java b/core/src/main/java/com/google/errorprone/bugpatterns/formatstring/FormatStringValidation.java index 660bcb32750..787ccd6ab6b 100644 --- a/core/src/main/java/com/google/errorprone/bugpatterns/formatstring/FormatStringValidation.java +++ b/core/src/main/java/com/google/errorprone/bugpatterns/formatstring/FormatStringValidation.java @@ -26,9 +26,14 @@ import com.google.errorprone.VisitorState; import com.google.errorprone.suppliers.Supplier; import com.google.errorprone.util.ASTHelpers; +import com.sun.source.tree.BlockTree; +import com.sun.source.tree.CaseTree; import com.sun.source.tree.ConditionalExpressionTree; import com.sun.source.tree.ExpressionTree; +import com.sun.source.tree.StatementTree; +import com.sun.source.tree.SwitchExpressionTree; import com.sun.source.tree.Tree; +import com.sun.source.tree.YieldTree; import com.sun.source.util.SimpleTreeVisitor; import com.sun.tools.javac.code.Symbol.MethodSymbol; import com.sun.tools.javac.code.Type; @@ -150,14 +155,69 @@ protected Void defaultAction(Tree tree, Void unused) { * or {@link Integer}. */ private static @Nullable Object getInstance(Tree tree, VisitorState state) { + tree = ASTHelpers.stripParentheses(tree); Object value = ASTHelpers.constValue(tree); if (value != null) { return value; } - Type type = ASTHelpers.getType(tree); + Type type = getExpressionType(tree, state); return getInstance(type, state); } + private static @Nullable Type getExpressionType(Tree tree, VisitorState state) { + tree = ASTHelpers.stripParentheses(tree); + if (tree instanceof SwitchExpressionTree switchExpression) { + Type fromCases = switchExpressionResultType(switchExpression, state); + if (fromCases != null && fromCases.getKind() != TypeKind.ERROR) { + return fromCases; + } + } + return ASTHelpers.getType(tree); + } + + private static @Nullable Type switchExpressionResultType( + SwitchExpressionTree switchExpression, VisitorState state) { + Types types = state.getTypes(); + Type common = null; + for (CaseTree caseTree : switchExpression.getCases()) { + Type caseType = getSwitchCaseResultType(caseTree, state); + if (caseType == null || caseType.getKind() == TypeKind.ERROR) { + continue; + } + Type normalized = types.unboxedTypeOrType(types.erasure(caseType)); + if (common == null) { + common = normalized; + } else if (!types.isSameType(common, normalized)) { + return null; + } + } + return common; + } + + private static @Nullable Type getSwitchCaseResultType(CaseTree caseTree, VisitorState state) { + Tree body = caseTree.getBody(); + if (body == null) { + for (StatementTree statement : caseTree.getStatements()) { + if (statement instanceof YieldTree yieldTree) { + return getExpressionType(yieldTree.getValue(), state); + } + } + return null; + } + body = ASTHelpers.stripParentheses(body); + if (body instanceof ExpressionTree expressionTree) { + return getExpressionType(expressionTree, state); + } + if (body instanceof BlockTree blockTree) { + for (StatementTree statement : blockTree.getStatements()) { + if (statement instanceof YieldTree yieldTree) { + return getExpressionType(yieldTree.getValue(), state); + } + } + } + return getExpressionType(body, state); + } + private static @Nullable Object getInstance(Type type, VisitorState state) { Types types = state.getTypes(); if (type.getKind() == TypeKind.NULL) { diff --git a/core/src/test/java/com/google/errorprone/bugpatterns/formatstring/FormatStringTest.java b/core/src/test/java/com/google/errorprone/bugpatterns/formatstring/FormatStringTest.java index ca69f6e2c79..f4564a82b7e 100644 --- a/core/src/test/java/com/google/errorprone/bugpatterns/formatstring/FormatStringTest.java +++ b/core/src/test/java/com/google/errorprone/bugpatterns/formatstring/FormatStringTest.java @@ -466,4 +466,115 @@ public static void main() { """) .doTest(); } + + @Test + public void switchExpressionArgument() { + compilationHelper + .addSourceLines( + "Test.java", + """ + class Test { + static final int FRIDAY = 5; + + void f() { + int day = FRIDAY; + var viaVar = switch (day) { + case 1, 5, 7 -> 6; + case 2 -> 7; + default -> 9; + }; + System.out.printf("via var: %d%n", viaVar); + System.out.printf( + "inline: %d%n", + switch (day) { + case 1, 5, 7 -> 6; + case 2 -> 7; + default -> 9; + }); + } + } + """) + .doTest(); + } + + @Test + public void switchExpressionArgument_nestedSwitch() { + compilationHelper + .addSourceLines( + "Test.java", + """ + class Test { + void f(int day, int mode) { + System.out.printf( + "%d%n", + switch (day) { + case 1 -> switch (mode) { + case 0 -> 1; + default -> 2; + }; + default -> 3; + }); + } + } + """) + .doTest(); + } + + @Test + public void switchExpressionArgument_blockYield() { + compilationHelper + .addSourceLines( + "Test.java", + """ + class Test { + void f(int day) { + System.out.printf( + "%d%n", + switch (day) { + case 1 -> { + yield 6; + } + default -> 9; + }); + } + } + """) + .doTest(); + } + + @Test + public void switchExpressionArgument_colonYield() { + compilationHelper + .addSourceLines( + "Test.java", + """ + class Test { + void f(int day) { + System.out.printf( + "%d%n", + switch (day) { + case 1: + yield 6; + default: + yield 9; + }); + } + } + """) + .doTest(); + } + + @Test + public void switchExpressionArgument_stringCases() { + testFormat( + "illegal format conversion: 'java.lang.String' cannot be formatted using '%d'", + "System.out.printf(\"%d\", switch (1) { case 1 -> \"a\"; default -> \"b\"; });"); + } + + @Test + public void switchExpressionArgument_mixedNumericCases() { + testFormat( + "cannot be formatted using '%d'", + "System.out.printf(\"%d\", switch (1) { case 1 -> 1; default -> 2.0; });"); + } }