Add filters to MockMvc with their init params and dispatcher types

Closes gh-37835
This commit is contained in:
Andy Wilkinson
2023-10-17 17:54:36 +01:00
parent 02c49b0287
commit daa903ab31
3 changed files with 47 additions and 22 deletions

View File

@@ -116,12 +116,8 @@ public class SpringBootMockMvcBuilderCustomizer implements MockMvcBuilderCustomi
private void addFilter(ConfigurableMockMvcBuilder<?> builder, AbstractFilterRegistrationBean<?> registration) {
Filter filter = registration.getFilter();
Collection<String> urls = registration.getUrlPatterns();
if (urls.isEmpty()) {
builder.addFilters(filter);
}
else {
builder.addFilter(filter, StringUtils.toStringArray(urls));
}
builder.addFilter(filter, registration.getInitParameters(), registration.determineDispatcherTypes(),
StringUtils.toStringArray(urls));
}
public void setAddFilters(boolean addFilters) {

View File

@@ -18,16 +18,21 @@ package org.springframework.boot.test.autoconfigure.web.servlet;
import java.util.ArrayList;
import java.util.Arrays;
import java.util.Collections;
import java.util.EnumSet;
import java.util.List;
import java.util.Map;
import java.util.concurrent.CountDownLatch;
import java.util.concurrent.TimeUnit;
import jakarta.servlet.DispatcherType;
import jakarta.servlet.Filter;
import jakarta.servlet.FilterChain;
import jakarta.servlet.FilterConfig;
import jakarta.servlet.ServletRequest;
import jakarta.servlet.ServletResponse;
import jakarta.servlet.http.HttpServlet;
import org.assertj.core.api.InstanceOfAssertFactories;
import org.junit.jupiter.api.Test;
import org.springframework.boot.test.autoconfigure.web.servlet.SpringBootMockMvcBuilderCustomizer.DeferredLinesWriter;
@@ -37,11 +42,12 @@ import org.springframework.boot.web.servlet.context.AnnotationConfigServletWebAp
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;
import org.springframework.mock.web.MockServletContext;
import org.springframework.test.util.ReflectionTestUtils;
import org.springframework.test.web.servlet.setup.DefaultMockMvcBuilder;
import org.springframework.test.web.servlet.setup.MockMvcBuilders;
import static org.assertj.core.api.Assertions.as;
import static org.assertj.core.api.Assertions.assertThat;
import static org.assertj.core.api.Assertions.tuple;
/**
* Tests for {@link SpringBootMockMvcBuilderCustomizer}.
@@ -51,7 +57,6 @@ import static org.assertj.core.api.Assertions.assertThat;
class SpringBootMockMvcBuilderCustomizerTests {
@Test
@SuppressWarnings("unchecked")
void customizeShouldAddFilters() {
AnnotationConfigServletWebApplicationContext context = new AnnotationConfigServletWebApplicationContext();
MockServletContext servletContext = new MockServletContext();
@@ -65,8 +70,11 @@ class SpringBootMockMvcBuilderCustomizerTests {
.getBean("filterRegistrationBean");
Filter testFilter = (Filter) context.getBean("testFilter");
Filter otherTestFilter = registrationBean.getFilter();
List<Filter> filters = (List<Filter>) ReflectionTestUtils.getField(builder, "filters");
assertThat(filters).containsExactlyInAnyOrder(testFilter, otherTestFilter);
assertThat(builder).extracting("filters", as(InstanceOfAssertFactories.LIST))
.extracting("delegate", "initParams", "dispatcherTypes")
.containsExactlyInAnyOrder(tuple(testFilter, Collections.emptyMap(), EnumSet.of(DispatcherType.REQUEST)),
tuple(otherTestFilter, Map.of("a", "alpha", "b", "bravo"),
EnumSet.of(DispatcherType.REQUEST, DispatcherType.ERROR)));
}
@Test
@@ -130,7 +138,11 @@ class SpringBootMockMvcBuilderCustomizerTests {
@Bean
FilterRegistrationBean<OtherTestFilter> filterRegistrationBean() {
return new FilterRegistrationBean<>(new OtherTestFilter());
FilterRegistrationBean<OtherTestFilter> filterRegistrationBean = new FilterRegistrationBean<>(
new OtherTestFilter());
filterRegistrationBean.setInitParameters(Map.of("a", "alpha", "b", "bravo"));
filterRegistrationBean.setDispatcherTypes(EnumSet.of(DispatcherType.REQUEST, DispatcherType.ERROR));
return filterRegistrationBean;
}
@Bean
@@ -182,4 +194,9 @@ class SpringBootMockMvcBuilderCustomizerTests {
}
static record RegisteredFilter(Filter filter, Map<String, String> initParameters,
EnumSet<DispatcherType> dispatcherTypes) {
}
}

View File

@@ -158,6 +158,28 @@ public abstract class AbstractFilterRegistrationBean<T extends Filter> extends D
Collections.addAll(this.urlPatterns, urlPatterns);
}
/**
* Determines the {@link DispatcherType dispatcher types} for which the filter should
* be registered. Applies defaults based on the type of filter being registered if
* none have been configured. Modifications to the returned {@link EnumSet} will have
* no effect on the registration.
* @return the dispatcher types, never {@code null}
* @since 3.2.0
*/
public EnumSet<DispatcherType> determineDispatcherTypes() {
if (this.dispatcherTypes == null) {
T filter = getFilter();
if (ClassUtils.isPresent("org.springframework.web.filter.OncePerRequestFilter",
filter.getClass().getClassLoader()) && filter instanceof OncePerRequestFilter) {
return EnumSet.allOf(DispatcherType.class);
}
else {
return EnumSet.of(DispatcherType.REQUEST);
}
}
return EnumSet.copyOf(this.dispatcherTypes);
}
/**
* Convenience method to {@link #setDispatcherTypes(EnumSet) set dispatcher types}
* using the specified elements.
@@ -216,17 +238,7 @@ public abstract class AbstractFilterRegistrationBean<T extends Filter> extends D
@Override
protected void configure(FilterRegistration.Dynamic registration) {
super.configure(registration);
EnumSet<DispatcherType> dispatcherTypes = this.dispatcherTypes;
if (dispatcherTypes == null) {
T filter = getFilter();
if (ClassUtils.isPresent("org.springframework.web.filter.OncePerRequestFilter",
filter.getClass().getClassLoader()) && filter instanceof OncePerRequestFilter) {
dispatcherTypes = EnumSet.allOf(DispatcherType.class);
}
else {
dispatcherTypes = EnumSet.of(DispatcherType.REQUEST);
}
}
EnumSet<DispatcherType> dispatcherTypes = determineDispatcherTypes();
Set<String> servletNames = new LinkedHashSet<>();
for (ServletRegistrationBean<?> servletRegistrationBean : this.servletRegistrationBeans) {
servletNames.add(servletRegistrationBean.getServletName());