diff --git a/src/main/java/org/openrewrite/staticanalysis/RemoveRedundantNullCheckBeforeInstanceof.java b/src/main/java/org/openrewrite/staticanalysis/RemoveRedundantNullCheckBeforeInstanceof.java index 6b2714c68..b7c8d35ce 100644 --- a/src/main/java/org/openrewrite/staticanalysis/RemoveRedundantNullCheckBeforeInstanceof.java +++ b/src/main/java/org/openrewrite/staticanalysis/RemoveRedundantNullCheckBeforeInstanceof.java @@ -31,6 +31,7 @@ import java.util.Set; import static java.util.Collections.singleton; +import static org.openrewrite.staticanalysis.SideEffects.mayHaveSideEffects; @EqualsAndHashCode(callSuper = false) @Value @@ -90,15 +91,17 @@ public J visitBinary(J.Binary binary, ExecutionContext ctx) { } private boolean isRedundantNullCheck(J.Binary nullCheck, J.InstanceOf instanceOf) { - if (nullCheck.getOperator() == J.Binary.Type.NotEqual) { - if (J.Literal.isLiteralValue(nullCheck.getLeft(), null)) { - return SemanticallyEqual.areEqual(nullCheck.getRight(), instanceOf.getExpression()); - } - if (J.Literal.isLiteralValue(nullCheck.getRight(), null)) { - return SemanticallyEqual.areEqual(nullCheck.getLeft(), instanceOf.getExpression()); - } + if (nullCheck.getOperator() != J.Binary.Type.NotEqual) { + return false; + } + Expression checked = J.Literal.isLiteralValue(nullCheck.getLeft(), null) ? nullCheck.getRight() : + J.Literal.isLiteralValue(nullCheck.getRight(), null) ? nullCheck.getLeft() : null; + if (checked == null || !SemanticallyEqual.areEqual(checked, instanceOf.getExpression())) { + return false; } - return false; + // The rewrite evaluates once what was evaluated twice, so both occurrences must be side-effect + // free; they can differ, as `SemanticallyEqual` matches a static field access against its qualified form + return !mayHaveSideEffects(checked) && !mayHaveSideEffects(instanceOf.getExpression()); } }); } diff --git a/src/test/java/org/openrewrite/staticanalysis/RemoveRedundantNullCheckBeforeInstanceofTest.java b/src/test/java/org/openrewrite/staticanalysis/RemoveRedundantNullCheckBeforeInstanceofTest.java index a1ad654a6..b369b3abe 100644 --- a/src/test/java/org/openrewrite/staticanalysis/RemoveRedundantNullCheckBeforeInstanceofTest.java +++ b/src/test/java/org/openrewrite/staticanalysis/RemoveRedundantNullCheckBeforeInstanceofTest.java @@ -87,35 +87,124 @@ void foo(Object obj) { ); } - + @Issue("https://github.com/openrewrite/rewrite-static-analysis/issues/953") @Test - void removeRedundantNullCheckWithMethodInvocation() { + void doNotChangeWhenNullCheckedExpressionIsMethodInvocation() { rewriteRun( //language=java java( """ class A { - void foo() { + void direct() { if (getValue() != null && getValue() instanceof String) { System.out.println("String value"); } + if (null != getValue() && getValue() instanceof String) { + System.out.println("String value"); + } + } + + void chained(boolean enabled) { + if (enabled && getValue() != null && getValue() instanceof String) { + System.out.println("String value"); + } } String getValue() { return "test"; } } + """ + ) + ); + } + + @Test + void removeRedundantNullCheckInChainedCondition() { + rewriteRun( + //language=java + java( + """ + class A { + void foo(boolean enabled, Object obj) { + if (enabled && obj != null && obj instanceof String) { + System.out.println("String value"); + } + } + } """, """ class A { + void foo(boolean enabled, Object obj) { + if (enabled && obj instanceof String) { + System.out.println("String value"); + } + } + } + """ + ) + ); + } + + @Issue("https://github.com/openrewrite/rewrite-static-analysis/issues/953") + @Test + void doNotChangeWhenConstructorCall() { + rewriteRun( + //language=java + java( + """ + class A { + void foo() { + if (new StringBuilder() != null && new StringBuilder() instanceof CharSequence) { + System.out.println("CharSequence value"); + } + } + } + """ + ) + ); + } + + @Issue("https://github.com/openrewrite/rewrite-static-analysis/issues/953") + @Test + void doNotChangeWhenArrayIndexHasSideEffect() { + rewriteRun( + //language=java + java( + """ + class A { + Object[] values = new Object[2]; + int i; + void foo() { - if (getValue() instanceof String) { + if (values[i++] != null && values[i++] instanceof String) { System.out.println("String value"); } } + } + """ + ) + ); + } - String getValue() { - return "test"; + @Issue("https://github.com/openrewrite/rewrite-static-analysis/issues/953") + @Test + void doNotChangeWhenOnlyTheNullCheckedOperandHasSideEffects() { + rewriteRun( + //language=java + java( + """ + class A { + static Integer count = 1; + + A getInstance() { + return this; + } + + void foo() { + if (getInstance().count != null && count instanceof Integer) { + System.out.println("Integer value"); + } } } """