diff --git a/nullaway/src/main/java/com/uber/nullaway/Config.java b/nullaway/src/main/java/com/uber/nullaway/Config.java
index 3efefe47c2..b35abe69c3 100644
--- a/nullaway/src/main/java/com/uber/nullaway/Config.java
+++ b/nullaway/src/main/java/com/uber/nullaway/Config.java
@@ -240,6 +240,14 @@ public interface Config {
*/
@Nullable String getCastToNonNullMethod();
+ /**
+ * Checks whether the castToNonNull() method should be treated as failing on null.
+ *
+ * @return true if castToNonNull() methods are guaranteed to throw when passed null,
+ * false otherwise
+ */
+ boolean castToNonNullMethodFailsOnNull();
+
/**
* Gets the suppression name aliases.
*
diff --git a/nullaway/src/main/java/com/uber/nullaway/DummyOptionsConfig.java b/nullaway/src/main/java/com/uber/nullaway/DummyOptionsConfig.java
index f564175788..ab16ec1d6b 100644
--- a/nullaway/src/main/java/com/uber/nullaway/DummyOptionsConfig.java
+++ b/nullaway/src/main/java/com/uber/nullaway/DummyOptionsConfig.java
@@ -179,6 +179,11 @@ public Set getOptionalClassPaths() {
throw new IllegalStateException(ERROR_MESSAGE);
}
+ @Override
+ public boolean castToNonNullMethodFailsOnNull() {
+ throw new IllegalStateException(ERROR_MESSAGE);
+ }
+
@Override
public Set getSuppressionNameAliases() {
throw new IllegalStateException(ERROR_MESSAGE);
diff --git a/nullaway/src/main/java/com/uber/nullaway/ErrorProneCLIFlagsConfig.java b/nullaway/src/main/java/com/uber/nullaway/ErrorProneCLIFlagsConfig.java
index 73f5119d7a..109bcc8ddb 100644
--- a/nullaway/src/main/java/com/uber/nullaway/ErrorProneCLIFlagsConfig.java
+++ b/nullaway/src/main/java/com/uber/nullaway/ErrorProneCLIFlagsConfig.java
@@ -66,6 +66,8 @@ final class ErrorProneCLIFlagsConfig implements Config {
static final String FL_NULLABLE_ANNOT = EP_FL_NAMESPACE + ":CustomNullableAnnotations";
static final String FL_NONNULL_ANNOT = EP_FL_NAMESPACE + ":CustomNonnullAnnotations";
static final String FL_CTNN_METHOD = EP_FL_NAMESPACE + ":CastToNonNullMethod";
+ static final String FL_CTNN_METHOD_FAILS_ON_NULL =
+ EP_FL_NAMESPACE + ":CastToNonNullMethodFailsOnNull";
static final String FL_EXTERNAL_INIT_ANNOT = EP_FL_NAMESPACE + ":ExternalInitAnnotations";
static final String FL_CONTRACT_ANNOT = EP_FL_NAMESPACE + ":CustomContractAnnotations";
static final String FL_UNANNOTATED_CLASSES = EP_FL_NAMESPACE + ":UnannotatedClasses";
@@ -243,6 +245,7 @@ final class ErrorProneCLIFlagsConfig implements Config {
private final ImmutableSet externalInitAnnotations;
private final ImmutableSet contractAnnotations;
private final @Nullable String castToNonNullMethod;
+ private final boolean castToNonNullMethodFailsOnNull;
private final String autofixSuppressionComment;
private final ImmutableSet suppressionNameAliases;
private final ImmutableSet skippedLibraryModels;
@@ -312,6 +315,7 @@ final class ErrorProneCLIFlagsConfig implements Config {
getPackagePattern(
getFlagStringSet(flags, FL_EXCLUDED_FIELD_ANNOT, DEFAULT_EXCLUDED_FIELD_ANNOT));
castToNonNullMethod = flags.get(FL_CTNN_METHOD).orElse(null);
+ castToNonNullMethodFailsOnNull = flags.getBoolean(FL_CTNN_METHOD_FAILS_ON_NULL).orElse(false);
legacyAnnotationLocation = flags.getBoolean(FL_LEGACY_ANNOTATION_LOCATION).orElse(false);
if (legacyAnnotationLocation && jspecifyMode) {
throw new IllegalStateException(
@@ -554,6 +558,11 @@ public boolean assertsEnabled() {
return castToNonNullMethod;
}
+ @Override
+ public boolean castToNonNullMethodFailsOnNull() {
+ return castToNonNullMethodFailsOnNull;
+ }
+
@Override
public ImmutableSet getSuppressionNameAliases() {
return suppressionNameAliases;
diff --git a/nullaway/src/main/java/com/uber/nullaway/handlers/LibraryModelsHandler.java b/nullaway/src/main/java/com/uber/nullaway/handlers/LibraryModelsHandler.java
index 676923b926..3d2cbaa181 100644
--- a/nullaway/src/main/java/com/uber/nullaway/handlers/LibraryModelsHandler.java
+++ b/nullaway/src/main/java/com/uber/nullaway/handlers/LibraryModelsHandler.java
@@ -401,8 +401,32 @@ private void setUnconditionalArgumentNullness(
AccessPath.AccessPathContext apContext) {
ImmutableSet requiredNonNullParameters =
getOptLibraryModels(state.context).failIfNullParameters(callee);
+ Set allNonNullParams;
+ if (config.castToNonNullMethodFailsOnNull()) {
+ ImmutableSet castToNonNullParameters =
+ getOptLibraryModels(state.context).castToNonNullMethod(callee);
+ String cliCastToNonNull = config.getCastToNonNullMethod();
+ boolean isCliCastToNonNull =
+ cliCastToNonNull != null
+ && callee.getParameters().size() == 1
+ && cliCastToNonNull.equals(
+ ASTHelpers.enclosingClass(callee) + "." + callee.getSimpleName());
+
+ if (castToNonNullParameters.isEmpty() && !isCliCastToNonNull) {
+ allNonNullParams = requiredNonNullParameters;
+ } else {
+ allNonNullParams = new HashSet<>(requiredNonNullParameters);
+ allNonNullParams.addAll(castToNonNullParameters);
+ if (isCliCastToNonNull) {
+ allNonNullParams.add(0);
+ }
+ }
+ } else {
+ allNonNullParams = requiredNonNullParameters;
+ }
+
for (AccessPath accessPath :
- accessPathsAtIndexes(requiredNonNullParameters, arguments, state, apContext)) {
+ accessPathsAtIndexes(allNonNullParams, arguments, state, apContext)) {
bothUpdates.set(accessPath, NONNULL);
}
}
diff --git a/nullaway/src/test/java/com/uber/nullaway/CoreTests.java b/nullaway/src/test/java/com/uber/nullaway/CoreTests.java
index cc09f0b041..43cd3e8144 100644
--- a/nullaway/src/test/java/com/uber/nullaway/CoreTests.java
+++ b/nullaway/src/test/java/com/uber/nullaway/CoreTests.java
@@ -23,6 +23,7 @@
package com.uber.nullaway;
import java.util.Arrays;
+import java.util.List;
import org.junit.Test;
import org.junit.runner.RunWith;
import org.junit.runners.JUnit4;
@@ -397,6 +398,123 @@ Object test4(@Nullable Object o) {
.doTest();
}
+ @SuppressWarnings("deprecation")
+ @Test
+ public void testCastToNonNullPropagationDefault() {
+ defaultCompilationHelper
+ .addSourceFile("testdata/Util.java")
+ .addSourceLines(
+ "Test.java",
+ """
+ package com.uber;
+ import javax.annotation.Nullable;
+ import static com.uber.nullaway.testdata.Util.castToNonNull;
+ class Test {
+ static class Foo {
+ @Nullable String getToken() { return ""; }
+ }
+ void test(Foo value) {
+ if (castToNonNull(value.getToken()).contains("abc")) {
+ // BUG: Diagnostic contains: dereferenced expression 'value.getToken()' is @Nullable
+ value.getToken().length();
+ }
+ }
+ }
+ """)
+ .doTest();
+ }
+
+ @SuppressWarnings("deprecation")
+ @Test
+ public void testCastToNonNullPropagationExplicitFalse() {
+ makeTestHelperWithArgs(
+ List.of(
+ "-d",
+ temporaryFolder.getRoot().getAbsolutePath(),
+ "-XepOpt:NullAway:AnnotatedPackages=com.uber",
+ "-XepOpt:NullAway:CastToNonNullMethod=com.uber.nullaway.testdata.Util.castToNonNull",
+ "-XepOpt:NullAway:CastToNonNullMethodFailsOnNull=false"))
+ .addSourceFile("testdata/Util.java")
+ .addSourceLines(
+ "Test.java",
+ """
+ package com.uber;
+ import javax.annotation.Nullable;
+ import static com.uber.nullaway.testdata.Util.castToNonNull;
+ class Test {
+ static class Foo {
+ @Nullable String getToken() { return ""; }
+ }
+ void test(Foo value) {
+ if (castToNonNull(value.getToken()).contains("abc")) {
+ // BUG: Diagnostic contains: dereferenced expression 'value.getToken()' is @Nullable
+ value.getToken().length();
+ }
+ }
+ }
+ """)
+ .doTest();
+ }
+
+ @SuppressWarnings("deprecation")
+ @Test
+ public void testCastToNonNullPropagationFailsOnNullEnabled() {
+ makeTestHelperWithArgs(
+ List.of(
+ "-d",
+ temporaryFolder.getRoot().getAbsolutePath(),
+ "-XepOpt:NullAway:AnnotatedPackages=com.uber",
+ "-XepOpt:NullAway:CastToNonNullMethod=com.uber.nullaway.testdata.Util.castToNonNull",
+ "-XepOpt:NullAway:CastToNonNullMethodFailsOnNull=true"))
+ .addSourceFile("testdata/Util.java")
+ .addSourceLines(
+ "Test.java",
+ """
+ package com.uber;
+ import javax.annotation.Nullable;
+ import static com.uber.nullaway.testdata.Util.castToNonNull;
+ class Test {
+ static class Foo {
+ @Nullable String getToken() { return ""; }
+ }
+ void test(Foo value) {
+ if (castToNonNull(value.getToken()).contains("abc")) {
+ value.getToken().length();
+ }
+ }
+ }
+ """)
+ .doTest();
+ }
+
+ @Test
+ public void testObjectsRequireNonNullPropagationUnchanged() {
+ makeTestHelperWithArgs(
+ List.of(
+ "-d",
+ temporaryFolder.getRoot().getAbsolutePath(),
+ "-XepOpt:NullAway:AnnotatedPackages=com.uber",
+ "-XepOpt:NullAway:CastToNonNullMethodFailsOnNull=false"))
+ .addSourceLines(
+ "Test.java",
+ """
+ package com.uber;
+ import java.util.Objects;
+ import javax.annotation.Nullable;
+ class Test {
+ static class Foo {
+ @Nullable String getToken() { return ""; }
+ }
+ void test(Foo value) {
+ if (Objects.requireNonNull(value.getToken()).contains("abc")) {
+ value.getToken().length();
+ }
+ }
+ }
+ """)
+ .doTest();
+ }
+
@Test
public void testReadStaticInConstructor() {
defaultCompilationHelper