diff --git a/core/spring-boot-test/src/main/java/org/springframework/boot/test/context/assertj/ApplicationContextAssertProvider.java b/core/spring-boot-test/src/main/java/org/springframework/boot/test/context/assertj/ApplicationContextAssertProvider.java index 93e855a8af0..1db7a364f8f 100644 --- a/core/spring-boot-test/src/main/java/org/springframework/boot/test/context/assertj/ApplicationContextAssertProvider.java +++ b/core/spring-boot-test/src/main/java/org/springframework/boot/test/context/assertj/ApplicationContextAssertProvider.java @@ -56,8 +56,8 @@ public interface ApplicationContextAssertProvider extends ApplicationContext, AssertProvider>, Closeable { /** - * Return an assert for AspectJ. - * @return an AspectJ assert + * Return an assert for AssertJ. + * @return an AssertJ assert * @deprecated to prevent accidental use. Prefer standard AssertJ * {@code assertThat(context)...} calls instead. */ @@ -132,6 +132,7 @@ public interface ApplicationContextAssertProvider Assert.isTrue(type.isInterface(), "'type' must be an interface"); Assert.notNull(contextType, "'contextType' must not be null"); Assert.isTrue(contextType.isInterface(), "'contextType' must be an interface"); + Assert.notNull(contextSupplier, "'contextSupplier' must not be null"); Class[] interfaces = merge(new Class[] { type, contextType }, additionalContextInterfaces); return (T) Proxy.newProxyInstance(Thread.currentThread().getContextClassLoader(), interfaces, new AssertProviderApplicationContextInvocationHandler(contextType, contextSupplier)); diff --git a/core/spring-boot-test/src/test/java/org/springframework/boot/test/context/assertj/ApplicationContextAssertProviderTests.java b/core/spring-boot-test/src/test/java/org/springframework/boot/test/context/assertj/ApplicationContextAssertProviderTests.java index 9b11048bd60..f45941c26b4 100644 --- a/core/spring-boot-test/src/test/java/org/springframework/boot/test/context/assertj/ApplicationContextAssertProviderTests.java +++ b/core/spring-boot-test/src/test/java/org/springframework/boot/test/context/assertj/ApplicationContextAssertProviderTests.java @@ -20,9 +20,7 @@ import java.util.function.Supplier; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; -import org.junit.jupiter.api.extension.ExtendWith; import org.mockito.Mock; -import org.mockito.junit.jupiter.MockitoExtension; import org.springframework.context.ApplicationContext; import org.springframework.context.ConfigurableApplicationContext; @@ -32,6 +30,7 @@ import static org.assertj.core.api.Assertions.assertThat; import static org.assertj.core.api.Assertions.assertThatIllegalArgumentException; import static org.assertj.core.api.Assertions.assertThatIllegalStateException; import static org.mockito.BDDMockito.then; +import static org.mockito.Mockito.mock; /** * Tests for {@link ApplicationContextAssertProvider} and @@ -39,12 +38,10 @@ import static org.mockito.BDDMockito.then; * * @author Phillip Webb */ -@ExtendWith(MockitoExtension.class) class ApplicationContextAssertProviderTests { @Mock - @SuppressWarnings("NullAway.Init") - private ConfigurableApplicationContext mockContext; + private final ConfigurableApplicationContext mockContext = mock(); private RuntimeException startupFailure; @@ -70,15 +67,7 @@ class ApplicationContextAssertProviderTests { } @Test - @SuppressWarnings("NullAway") // Test null check void getWhenTypeIsClassShouldThrowException() { - assertThatIllegalArgumentException().isThrownBy( - () -> ApplicationContextAssertProvider.get(null, ApplicationContext.class, this.mockContextSupplier)) - .withMessageContaining("'type' must not be null"); - } - - @Test - void getWhenContextTypeIsNullShouldThrowException() { assertThatIllegalArgumentException() .isThrownBy(() -> ApplicationContextAssertProvider.get(TestAssertProviderApplicationContextClass.class, ApplicationContext.class, this.mockContextSupplier)) @@ -87,7 +76,7 @@ class ApplicationContextAssertProviderTests { @Test @SuppressWarnings("NullAway") // Test null check - void getWhenContextTypeIsClassShouldThrowException() { + void getWhenContextTypeIsNullShouldThrowException() { assertThatIllegalArgumentException() .isThrownBy(() -> ApplicationContextAssertProvider.get(TestAssertProviderApplicationContext.class, null, this.mockContextSupplier)) @@ -95,13 +84,22 @@ class ApplicationContextAssertProviderTests { } @Test - void getWhenSupplierIsNullShouldThrowException() { + void getWhenContextTypeIsClassShouldThrowException() { assertThatIllegalArgumentException() .isThrownBy(() -> ApplicationContextAssertProvider.get(TestAssertProviderApplicationContext.class, StaticApplicationContext.class, this.mockContextSupplier)) .withMessageContaining("'contextType' must be an interface"); } + @Test + @SuppressWarnings("NullAway") // Test null check + void getWhenSupplierIsNullShouldThrowException() { + assertThatIllegalArgumentException() + .isThrownBy(() -> ApplicationContextAssertProvider.get(TestAssertProviderApplicationContext.class, + ApplicationContext.class, null)) + .withMessageContaining("'contextSupplier' must not be null"); + } + @Test void getWhenContextStartsShouldReturnProxyThatCallsRealMethods() { ApplicationContextAssertProvider context = get(this.mockContextSupplier);