diff --git a/spring-cloud-gateway-server/src/main/java/org/springframework/cloud/gateway/filter/headers/RemoveHopByHopHeadersFilter.java b/spring-cloud-gateway-server/src/main/java/org/springframework/cloud/gateway/filter/headers/RemoveHopByHopHeadersFilter.java index 54332534..a04e9e4b 100644 --- a/spring-cloud-gateway-server/src/main/java/org/springframework/cloud/gateway/filter/headers/RemoveHopByHopHeadersFilter.java +++ b/spring-cloud-gateway-server/src/main/java/org/springframework/cloud/gateway/filter/headers/RemoveHopByHopHeadersFilter.java @@ -21,10 +21,12 @@ import java.util.HashSet; import java.util.List; import java.util.Map; import java.util.Set; +import java.util.stream.Collectors; import org.springframework.boot.context.properties.ConfigurationProperties; import org.springframework.core.Ordered; import org.springframework.http.HttpHeaders; +import org.springframework.util.Assert; import org.springframework.web.server.ServerWebExchange; @ConfigurationProperties("spring.cloud.gateway.filter.remove-hop-by-hop") @@ -51,7 +53,8 @@ public class RemoveHopByHopHeadersFilter implements HttpHeadersFilter, Ordered { } public void setHeaders(Set headers) { - this.headers = headers; + Assert.notNull(headers, "headers may not be null"); + this.headers = headers.stream().map(String::toLowerCase).collect(Collectors.toSet()); } @Override diff --git a/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/filter/headers/RemoveHopByHopHeadersFilterTests.java b/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/filter/headers/RemoveHopByHopHeadersFilterTests.java index 1be4ac28..2f211bcf 100644 --- a/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/filter/headers/RemoveHopByHopHeadersFilterTests.java +++ b/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/filter/headers/RemoveHopByHopHeadersFilterTests.java @@ -18,6 +18,7 @@ package org.springframework.cloud.gateway.filter.headers; import java.util.Arrays; import java.util.HashSet; +import java.util.LinkedHashSet; import java.util.Set; import org.junit.jupiter.api.Test; @@ -25,6 +26,7 @@ import org.junit.jupiter.api.Test; import org.springframework.http.HttpHeaders; import org.springframework.mock.http.server.reactive.MockServerHttpRequest; import org.springframework.mock.web.server.MockServerWebExchange; +import org.springframework.util.StringUtils; import static org.assertj.core.api.Assertions.assertThat; import static org.springframework.cloud.gateway.filter.headers.RemoveHopByHopHeadersFilter.HEADERS_REMOVED_ON_REQUEST; @@ -52,6 +54,23 @@ public class RemoveHopByHopHeadersFilterTests { testFilter(MockServerWebExchange.from(builder)); } + @Test + public void caseInsensitiveCustom() { + MockServerHttpRequest.BaseBuilder builder = MockServerHttpRequest.get("http://localhost/get"); + + HEADERS_REMOVED_ON_REQUEST + .forEach(header -> builder.header(StringUtils.capitalize(header.toLowerCase()), header + "1")); + + LinkedHashSet customHeaders = new LinkedHashSet<>(); + HEADERS_REMOVED_ON_REQUEST.forEach(header -> { + String newHeader = header.charAt(0) + StringUtils.capitalize(header.substring(1)); + customHeaders.add(newHeader); + }); + RemoveHopByHopHeadersFilter filter = new RemoveHopByHopHeadersFilter(); + filter.setHeaders(customHeaders); + testFilter(filter, MockServerWebExchange.from(builder)); + } + @Test public void removesHeadersListedInConnectionHeader() { MockServerHttpRequest.BaseBuilder builder = MockServerHttpRequest.get("http://localhost/get"); @@ -64,7 +83,11 @@ public class RemoveHopByHopHeadersFilterTests { } private void testFilter(MockServerWebExchange exchange, String... additionalHeaders) { - RemoveHopByHopHeadersFilter filter = new RemoveHopByHopHeadersFilter(); + testFilter(new RemoveHopByHopHeadersFilter(), exchange, additionalHeaders); + } + + private void testFilter(RemoveHopByHopHeadersFilter filter, MockServerWebExchange exchange, + String... additionalHeaders) { HttpHeaders headers = filter.filter(exchange.getRequest().getHeaders(), exchange); Set toRemove = new HashSet<>(HEADERS_REMOVED_ON_REQUEST);