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