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..45cfc56c6af 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(); @@ -63,6 +63,16 @@ public SecurityContextHolderFilter(SecurityContextRepository securityContextRepo this.securityContextRepository = securityContextRepository; } + /** + * Sets the {@link SecurityContextRepository} to use. + * @param securityContextRepository the repository to use. Cannot be null. + * @since 7.1 + */ + public void setSecurityContextRepository(SecurityContextRepository securityContextRepository) { + Assert.notNull(securityContextRepository, "securityContextRepository cannot be null"); + this.securityContextRepository = securityContextRepository; + } + @Override public void doFilter(ServletRequest request, ServletResponse response, FilterChain chain) throws IOException, ServletException { 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..7d361d4a1be 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; @@ -77,6 +78,11 @@ void setup() { this.filter = new SecurityContextHolderFilter(this.repository); } + @Test + void setSecurityContextRepositoryWhenNullThenThrowsIllegalArgumentException() { + assertThatIllegalArgumentException().isThrownBy(() -> this.filter.setSecurityContextRepository(null)); + } + @AfterEach void cleanup() { SecurityContextHolder.clearContext();