diff --git a/test/src/main/java/org/springframework/security/test/web/support/WebTestUtils.java b/test/src/main/java/org/springframework/security/test/web/support/WebTestUtils.java index eb23d4b9468..c7290ca4d47 100644 --- a/test/src/main/java/org/springframework/security/test/web/support/WebTestUtils.java +++ b/test/src/main/java/org/springframework/security/test/web/support/WebTestUtils.java @@ -98,7 +98,7 @@ public static void setSecurityContextRepository(HttpServletRequest request, } SecurityContextHolderFilter holderFilter = findFilter(request, SecurityContextHolderFilter.class); if (holderFilter != null) { - ReflectionTestUtils.setField(holderFilter, "securityContextRepository", securityContextRepository); + holderFilter.setSecurityContextRepository(securityContextRepository); } } diff --git a/web/src/main/java/org/springframework/security/web/context/SecurityContextHolderFilter.java b/web/src/main/java/org/springframework/security/web/context/SecurityContextHolderFilter.java index cacad3e3b80..dfd8bbe30b2 100644 --- a/web/src/main/java/org/springframework/security/web/context/SecurityContextHolderFilter.java +++ b/web/src/main/java/org/springframework/security/web/context/SecurityContextHolderFilter.java @@ -49,7 +49,7 @@ public class SecurityContextHolderFilter extends GenericFilterBean { private static final String FILTER_APPLIED = SecurityContextHolderFilter.class.getName() + ".APPLIED"; - private final SecurityContextRepository securityContextRepository; + private SecurityContextRepository securityContextRepository; private SecurityContextHolderStrategy securityContextHolderStrategy = SecurityContextHolder .getContextHolderStrategy(); @@ -69,6 +69,10 @@ public void doFilter(ServletRequest request, ServletResponse response, FilterCha doFilter((HttpServletRequest) request, (HttpServletResponse) response, chain); } + public void setSecurityContextRepository(SecurityContextRepository securityContextRepository) { + Assert.notNull(securityContextRepository, "securityContextRepository cannot be null"); + this.securityContextRepository = securityContextRepository; + } private void doFilter(HttpServletRequest request, HttpServletResponse response, FilterChain chain) throws ServletException, IOException { if (request.getAttribute(FILTER_APPLIED) != null) { diff --git a/web/src/test/java/org/springframework/security/web/context/SecurityContextHolderFilterTests.java b/web/src/test/java/org/springframework/security/web/context/SecurityContextHolderFilterTests.java index 1b1fc1e8deb..b0fae6f9698 100644 --- a/web/src/test/java/org/springframework/security/web/context/SecurityContextHolderFilterTests.java +++ b/web/src/test/java/org/springframework/security/web/context/SecurityContextHolderFilterTests.java @@ -43,6 +43,7 @@ import org.springframework.security.core.context.SecurityContextImpl; import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatIllegalArgumentException; import static org.mockito.BDDMockito.given; import static org.mockito.Mockito.inOrder; import static org.mockito.Mockito.lenient; @@ -81,7 +82,12 @@ void setup() { void cleanup() { SecurityContextHolder.clearContext(); } - + @Test + void setSecurityContextRepositoryWhenNullThenThrowsIllegalArgumentException() { + assertThatIllegalArgumentException() + .isThrownBy(() -> this.filter.setSecurityContextRepository(null)) + .withMessage("securityContextRepository cannot be null"); + } @Test void doFilterThenSetsAndClearsSecurityContextHolder() throws Exception { Authentication authentication = TestAuthentication.authenticatedUser();