Adapt to upstream Spring Security changes

This commit is contained in:
Phillip Webb
2023-12-14 20:33:09 -08:00
parent 5915db09e6
commit 1d10e51755
3 changed files with 86 additions and 24 deletions

View File

@@ -18,7 +18,9 @@ package org.springframework.boot.actuate.autoconfigure.cloudfoundry.servlet;
import java.util.Arrays;
import java.util.Collection;
import java.util.List;
import jakarta.servlet.Filter;
import org.junit.jupiter.api.Test;
import org.springframework.boot.actuate.autoconfigure.endpoint.EndpointAutoConfiguration;
@@ -43,6 +45,7 @@ import org.springframework.boot.autoconfigure.security.servlet.SecurityAutoConfi
import org.springframework.boot.autoconfigure.web.client.RestTemplateAutoConfiguration;
import org.springframework.boot.autoconfigure.web.servlet.DispatcherServletAutoConfiguration;
import org.springframework.boot.autoconfigure.web.servlet.WebMvcAutoConfiguration;
import org.springframework.boot.test.context.assertj.AssertableWebApplicationContext;
import org.springframework.boot.test.context.runner.WebApplicationContextRunner;
import org.springframework.context.ApplicationContext;
import org.springframework.http.HttpMethod;
@@ -55,6 +58,7 @@ import org.springframework.test.web.servlet.MockMvc;
import org.springframework.test.web.servlet.setup.MockMvcBuilders;
import org.springframework.web.client.RestTemplate;
import org.springframework.web.cors.CorsConfiguration;
import org.springframework.web.filter.CompositeFilter;
import static org.assertj.core.api.Assertions.assertThat;
import static org.springframework.test.web.servlet.request.MockMvcRequestBuilders.get;
@@ -173,9 +177,7 @@ class CloudFoundryActuatorAutoConfigurationTests {
this.contextRunner.withBean(TestEndpoint.class, TestEndpoint::new)
.withPropertyValues("VCAP_APPLICATION:---", "vcap.application.application_id:my-app-id")
.run((context) -> {
FilterChainProxy securityFilterChain = (FilterChainProxy) context
.getBean(BeanIds.SPRING_SECURITY_FILTER_CHAIN);
SecurityFilterChain chain = securityFilterChain.getFilterChains().get(0);
SecurityFilterChain chain = getSecurityFilterChain(context);
assertThat(chain.getFilters()).isEmpty();
MockHttpServletRequest request = new MockHttpServletRequest();
testCloudFoundrySecurity(request, BASE_PATH, chain);
@@ -189,6 +191,27 @@ class CloudFoundryActuatorAutoConfigurationTests {
});
}
private SecurityFilterChain getSecurityFilterChain(AssertableWebApplicationContext context) {
Filter springSecurityFilterChain = context.getBean(BeanIds.SPRING_SECURITY_FILTER_CHAIN, Filter.class);
FilterChainProxy filterChainProxy = getFilterChainProxy(springSecurityFilterChain);
SecurityFilterChain securityFilterChain = filterChainProxy.getFilterChains().get(0);
return securityFilterChain;
}
private FilterChainProxy getFilterChainProxy(Filter filter) {
if (filter instanceof FilterChainProxy filterChainProxy) {
return filterChainProxy;
}
if (filter instanceof CompositeFilter) {
List<?> filters = (List<?>) ReflectionTestUtils.getField(filter, "filters");
return (FilterChainProxy) filters.stream()
.filter(FilterChainProxy.class::isInstance)
.findFirst()
.orElseThrow();
}
throw new IllegalStateException("No FilterChainProxy found");
}
private static void testCloudFoundrySecurity(MockHttpServletRequest request, String servletPath,
SecurityFilterChain chain) {
request.setServletPath(servletPath);

View File

@@ -48,6 +48,7 @@ import org.springframework.security.web.FilterChainProxy;
import org.springframework.security.web.SecurityFilterChain;
import org.springframework.test.util.ReflectionTestUtils;
import org.springframework.util.ObjectUtils;
import org.springframework.web.filter.CompositeFilter;
import static org.assertj.core.api.Assertions.assertThat;
@@ -68,7 +69,7 @@ class OAuth2WebSecurityConfigurationTests {
.run((context) -> {
ClientRegistrationRepository expected = context.getBean(ClientRegistrationRepository.class);
ClientRegistrationRepository actual = (ClientRegistrationRepository) ReflectionTestUtils.getField(
getFilters(context, OAuth2LoginAuthenticationFilter.class).get(0),
getSecurityFilters(context, OAuth2LoginAuthenticationFilter.class).get(0),
"clientRegistrationRepository");
assertThat(isEqual(expected.findByRegistrationId("first"), actual.findByRegistrationId("first")))
.isTrue();
@@ -85,7 +86,7 @@ class OAuth2WebSecurityConfigurationTests {
.run((context) -> {
ClientRegistrationRepository expected = context.getBean(ClientRegistrationRepository.class);
ClientRegistrationRepository actual = (ClientRegistrationRepository) ReflectionTestUtils.getField(
getFilters(context, OAuth2AuthorizationCodeGrantFilter.class).get(0),
getSecurityFilters(context, OAuth2AuthorizationCodeGrantFilter.class).get(0),
"clientRegistrationRepository");
assertThat(isEqual(expected.findByRegistrationId("first"), actual.findByRegistrationId("first")))
.isTrue();
@@ -98,8 +99,8 @@ class OAuth2WebSecurityConfigurationTests {
void securityConfigurerBacksOffWhenClientRegistrationBeanAbsent() {
this.contextRunner.withUserConfiguration(TestConfig.class, OAuth2WebSecurityConfiguration.class)
.run((context) -> {
assertThat(getFilters(context, OAuth2LoginAuthenticationFilter.class)).isEmpty();
assertThat(getFilters(context, OAuth2AuthorizationCodeGrantFilter.class)).isEmpty();
assertThat(getSecurityFilters(context, OAuth2LoginAuthenticationFilter.class)).isEmpty();
assertThat(getSecurityFilters(context, OAuth2AuthorizationCodeGrantFilter.class)).isEmpty();
});
}
@@ -124,8 +125,8 @@ class OAuth2WebSecurityConfigurationTests {
this.contextRunner.withConfiguration(AutoConfigurations.of(WebMvcAutoConfiguration.class))
.withUserConfiguration(TestSecurityFilterChainConfiguration.class, OAuth2WebSecurityConfiguration.class)
.run((context) -> {
assertThat(getFilters(context, OAuth2LoginAuthenticationFilter.class)).isEmpty();
assertThat(getFilters(context, OAuth2AuthorizationCodeGrantFilter.class)).isEmpty();
assertThat(getSecurityFilters(context, OAuth2LoginAuthenticationFilter.class)).isEmpty();
assertThat(getSecurityFilters(context, OAuth2AuthorizationCodeGrantFilter.class)).isEmpty();
assertThat(context).getBean(OAuth2AuthorizedClientService.class).isNotNull();
});
}
@@ -137,8 +138,8 @@ class OAuth2WebSecurityConfigurationTests {
OAuth2WebSecurityConfiguration.class)
.withClassLoader(new FilteredClassLoader(SecurityFilterChain.class))
.run((context) -> {
assertThat(getFilters(context, OAuth2LoginAuthenticationFilter.class)).isEmpty();
assertThat(getFilters(context, OAuth2AuthorizationCodeGrantFilter.class)).isEmpty();
assertThat(getSecurityFilters(context, OAuth2LoginAuthenticationFilter.class)).isEmpty();
assertThat(getSecurityFilters(context, OAuth2AuthorizationCodeGrantFilter.class)).isEmpty();
});
}
@@ -164,11 +165,29 @@ class OAuth2WebSecurityConfigurationTests {
});
}
private List<Filter> getFilters(AssertableWebApplicationContext context, Class<? extends Filter> filter) {
FilterChainProxy filterChain = (FilterChainProxy) context.getBean(BeanIds.SPRING_SECURITY_FILTER_CHAIN);
List<SecurityFilterChain> filterChains = filterChain.getFilterChains();
List<Filter> filters = filterChains.get(0).getFilters();
return filters.stream().filter(filter::isInstance).toList();
private List<Filter> getSecurityFilters(AssertableWebApplicationContext context, Class<? extends Filter> filter) {
return getSecurityFilterChain(context).getFilters().stream().filter(filter::isInstance).toList();
}
private SecurityFilterChain getSecurityFilterChain(AssertableWebApplicationContext context) {
Filter springSecurityFilterChain = context.getBean(BeanIds.SPRING_SECURITY_FILTER_CHAIN, Filter.class);
FilterChainProxy filterChainProxy = getFilterChainProxy(springSecurityFilterChain);
SecurityFilterChain securityFilterChain = filterChainProxy.getFilterChains().get(0);
return securityFilterChain;
}
private FilterChainProxy getFilterChainProxy(Filter filter) {
if (filter instanceof FilterChainProxy filterChainProxy) {
return filterChainProxy;
}
if (filter instanceof CompositeFilter) {
List<?> filters = (List<?>) ReflectionTestUtils.getField(filter, "filters");
return (FilterChainProxy) filters.stream()
.filter(FilterChainProxy.class::isInstance)
.findFirst()
.orElseThrow();
}
throw new IllegalStateException("No FilterChainProxy found");
}
private boolean isEqual(ClientRegistration reg1, ClientRegistration reg2) {

View File

@@ -46,6 +46,8 @@ import org.springframework.security.saml2.provider.service.web.authentication.Sa
import org.springframework.security.saml2.provider.service.web.authentication.logout.Saml2LogoutRequestFilter;
import org.springframework.security.web.FilterChainProxy;
import org.springframework.security.web.SecurityFilterChain;
import org.springframework.test.util.ReflectionTestUtils;
import org.springframework.web.filter.CompositeFilter;
import static org.assertj.core.api.Assertions.assertThat;
import static org.mockito.Mockito.mock;
@@ -208,7 +210,7 @@ class Saml2RelyingPartyAutoConfigurationTests {
@Test
void samlLoginShouldBeConfigured() {
this.contextRunner.withPropertyValues(getPropertyValues())
.run((context) -> assertThat(hasFilter(context, Saml2WebSsoAuthenticationFilter.class)).isTrue());
.run((context) -> assertThat(hasSecurityFilter(context, Saml2WebSsoAuthenticationFilter.class)).isTrue());
}
@Test
@@ -216,7 +218,7 @@ class Saml2RelyingPartyAutoConfigurationTests {
this.contextRunner.withConfiguration(AutoConfigurations.of(WebMvcAutoConfiguration.class))
.withUserConfiguration(TestSecurityFilterChainConfig.class)
.withPropertyValues(getPropertyValues())
.run((context) -> assertThat(hasFilter(context, Saml2WebSsoAuthenticationFilter.class)).isFalse());
.run((context) -> assertThat(hasSecurityFilter(context, Saml2WebSsoAuthenticationFilter.class)).isFalse());
}
@Test
@@ -229,7 +231,7 @@ class Saml2RelyingPartyAutoConfigurationTests {
@Test
void samlLogoutShouldBeConfigured() {
this.contextRunner.withPropertyValues(getPropertyValues())
.run((context) -> assertThat(hasFilter(context, Saml2LogoutRequestFilter.class)).isTrue());
.run((context) -> assertThat(hasSecurityFilter(context, Saml2LogoutRequestFilter.class)).isTrue());
}
private String[] getPropertyValuesWithoutSigningCredentials(boolean signRequests) {
@@ -323,11 +325,29 @@ class Saml2RelyingPartyAutoConfigurationTests {
PREFIX + ".foo.acs.binding=redirect" };
}
private boolean hasFilter(AssertableWebApplicationContext context, Class<? extends Filter> filter) {
FilterChainProxy filterChain = (FilterChainProxy) context.getBean(BeanIds.SPRING_SECURITY_FILTER_CHAIN);
List<SecurityFilterChain> filterChains = filterChain.getFilterChains();
List<Filter> filters = filterChains.get(0).getFilters();
return filters.stream().anyMatch(filter::isInstance);
private boolean hasSecurityFilter(AssertableWebApplicationContext context, Class<? extends Filter> filter) {
return getSecurityFilterChain(context).getFilters().stream().anyMatch(filter::isInstance);
}
private SecurityFilterChain getSecurityFilterChain(AssertableWebApplicationContext context) {
Filter springSecurityFilterChain = context.getBean(BeanIds.SPRING_SECURITY_FILTER_CHAIN, Filter.class);
FilterChainProxy filterChainProxy = getFilterChainProxy(springSecurityFilterChain);
SecurityFilterChain securityFilterChain = filterChainProxy.getFilterChains().get(0);
return securityFilterChain;
}
private FilterChainProxy getFilterChainProxy(Filter filter) {
if (filter instanceof FilterChainProxy filterChainProxy) {
return filterChainProxy;
}
if (filter instanceof CompositeFilter) {
List<?> filters = (List<?>) ReflectionTestUtils.getField(filter, "filters");
return (FilterChainProxy) filters.stream()
.filter(FilterChainProxy.class::isInstance)
.findFirst()
.orElseThrow();
}
throw new IllegalStateException("No FilterChainProxy found");
}
private void setupMockResponse(MockWebServer server, Resource resourceBody) throws Exception {