diff --git a/spring-cloud-gateway-mvc/src/main/java/org/springframework/cloud/gateway/mvc/ProxyExchange.java b/spring-cloud-gateway-mvc/src/main/java/org/springframework/cloud/gateway/mvc/ProxyExchange.java index 28b3687d..394d99a8 100644 --- a/spring-cloud-gateway-mvc/src/main/java/org/springframework/cloud/gateway/mvc/ProxyExchange.java +++ b/spring-cloud-gateway-mvc/src/main/java/org/springframework/cloud/gateway/mvc/ProxyExchange.java @@ -32,6 +32,7 @@ import java.util.List; import java.util.Set; import java.util.Vector; import java.util.function.Function; +import java.util.stream.Collectors; import javax.servlet.ReadListener; import javax.servlet.ServletInputStream; @@ -215,6 +216,8 @@ public class ProxyExchange { if (this.sensitive == null) { this.sensitive = new HashSet<>(); } + + this.sensitive.clear(); for (String name : names) { this.sensitive.add(name.toLowerCase()); } @@ -342,20 +345,19 @@ public class ProxyExchange { } private BodyBuilder headers(BodyBuilder builder) { - Set sensitive = this.sensitive; - if (sensitive == null) { - sensitive = DEFAULT_SENSITIVE; - } proxy(); - for (String name : headers.keySet()) { - if (sensitive.contains(name.toLowerCase())) { - continue; - } + for (String name : filterHeaderKeys(headers)) { builder.header(name, headers.get(name).toArray(new String[0])); } return builder; } + private Set filterHeaderKeys(HttpHeaders headers) { + final Set sensitiveHeaders = this.sensitive != null ? this.sensitive : DEFAULT_SENSITIVE; + return headers.keySet().stream().filter(header -> !sensitiveHeaders.contains(header.toLowerCase())) + .collect(Collectors.toSet()); + } + private void proxy() { try { URI uri = new URI(webRequest.getNativeRequest(HttpServletRequest.class).getRequestURL().toString()); diff --git a/spring-cloud-gateway-mvc/src/test/java/org/springframework/cloud/gateway/mvc/ProductionConfigurationTests.java b/spring-cloud-gateway-mvc/src/test/java/org/springframework/cloud/gateway/mvc/ProductionConfigurationTests.java index bda865a7..c519e360 100644 --- a/spring-cloud-gateway-mvc/src/test/java/org/springframework/cloud/gateway/mvc/ProductionConfigurationTests.java +++ b/spring-cloud-gateway-mvc/src/test/java/org/springframework/cloud/gateway/mvc/ProductionConfigurationTests.java @@ -236,6 +236,29 @@ public class ProductionConfigurationTests { assertThat(deleteResponse.getBody().get("deleted")).isEqualToComparingFieldByField(foo); } + @Test + @SuppressWarnings({ "Duplicates", "unchecked" }) + public void testSensitiveHeadersOverride() throws Exception { + Map> headers = rest + .exchange( + RequestEntity.get(rest.getRestTemplate().getUriTemplateHandler().expand("/proxy/headers")) + .header("foo", "bar").header("abc", "xyz").header("cookie", "monster").build(), + Map.class) + .getBody(); + assertThat(headers).doesNotContainKey("foo").doesNotContainKey("hello").containsKeys("bar", "abc"); + + assertThat(headers.get("cookie")).containsOnly("monster"); + } + + @Test + public void testSensitiveHeadersDefault() throws Exception { + Map> headers = rest.exchange(RequestEntity + .get(rest.getRestTemplate().getUriTemplateHandler().expand("/proxy/sensitive-headers-default")) + .header("cookie", "monster").build(), Map.class).getBody(); + + assertThat(headers).doesNotContainKey("cookie"); + } + @Test @SuppressWarnings({ "Duplicates", "unchecked" }) public void headers() throws Exception { @@ -406,8 +429,16 @@ public class ProductionConfigurationTests { @GetMapping("/proxy/headers") @SuppressWarnings("Duplicates") public ResponseEntity>> headers(ProxyExchange>> proxy) { - proxy.sensitive("foo"); - proxy.sensitive("hello"); + proxy.sensitive("foo", "hello"); + proxy.header("bar", "hello"); + proxy.header("abc", "123"); + proxy.header("hello", "world"); + return proxy.uri(home.toString() + "/headers").get(); + } + + @GetMapping("/proxy/sensitive-headers-default") + public ResponseEntity>> defaultSensitiveHeaders( + ProxyExchange>> proxy) { proxy.header("bar", "hello"); proxy.header("abc", "123"); proxy.header("hello", "world"); diff --git a/spring-cloud-gateway-webflux/src/main/java/org/springframework/cloud/gateway/webflux/ProxyExchange.java b/spring-cloud-gateway-webflux/src/main/java/org/springframework/cloud/gateway/webflux/ProxyExchange.java index c8f7b5bb..5044ffcb 100644 --- a/spring-cloud-gateway-webflux/src/main/java/org/springframework/cloud/gateway/webflux/ProxyExchange.java +++ b/spring-cloud-gateway-webflux/src/main/java/org/springframework/cloud/gateway/webflux/ProxyExchange.java @@ -140,8 +140,6 @@ public class ProxyExchange { this.bindingContext = bindingContext; this.responseType = type; this.rest = rest; - this.sensitive = new HashSet<>(DEFAULT_SENSITIVE.size()); - this.sensitive.addAll(DEFAULT_SENSITIVE); this.httpMethod = exchange.getRequest().getMethod(); } @@ -206,6 +204,8 @@ public class ProxyExchange { if (this.sensitive == null) { this.sensitive = new HashSet<>(); } + + this.sensitive.clear(); for (String name : names) { this.sensitive.add(name.toLowerCase()); } @@ -376,7 +376,8 @@ public class ProxyExchange { } private Set filterHeaderKeys(HttpHeaders headers) { - return headers.keySet().stream().filter(header -> !sensitive.contains(header.toLowerCase())) + final Set sensitiveHeaders = this.sensitive != null ? this.sensitive : DEFAULT_SENSITIVE; + return headers.keySet().stream().filter(header -> !sensitiveHeaders.contains(header.toLowerCase())) .collect(Collectors.toSet()); } diff --git a/spring-cloud-gateway-webflux/src/test/java/org/springframework/cloud/gateway/webflux/ProductionConfigurationTests.java b/spring-cloud-gateway-webflux/src/test/java/org/springframework/cloud/gateway/webflux/ProductionConfigurationTests.java index abe0df55..a508d2df 100644 --- a/spring-cloud-gateway-webflux/src/test/java/org/springframework/cloud/gateway/webflux/ProductionConfigurationTests.java +++ b/spring-cloud-gateway-webflux/src/test/java/org/springframework/cloud/gateway/webflux/ProductionConfigurationTests.java @@ -174,6 +174,29 @@ public class ProductionConfigurationTests { .isEqualTo("host=localhost:" + port + ";foobar"); } + @Test + @SuppressWarnings({ "Duplicates", "unchecked" }) + public void testSensitiveHeadersOverride() throws Exception { + Map> headers = rest + .exchange(RequestEntity.get(rest.getRestTemplate().getUriTemplateHandler().expand("/proxy/headers")) + .header("foo", "bar").header("abc", "xyz").header("cookie", "monster").build(), Map.class) + .getBody(); + assertThat(headers).doesNotContainKey("foo").doesNotContainKey("hello").containsKeys("bar", "abc"); + + assertThat(headers.get("cookie")).containsOnly("monster"); + } + + @Test + public void testSensitiveHeadersDefault() throws Exception { + Map> headers = rest + .exchange(RequestEntity.get(rest.getRestTemplate().getUriTemplateHandler().expand("/proxy/sensitive-headers-default")) + .header("cookie", "monster") + .build(), Map.class) + .getBody(); + + assertThat(headers).doesNotContainKey("cookie"); + } + @Test @SuppressWarnings({ "Duplicates", "unchecked" }) public void headers() throws Exception { @@ -317,8 +340,16 @@ public class ProductionConfigurationTests { @GetMapping("/proxy/headers") public Mono>>> headers( ProxyExchange>> proxy) { - proxy.sensitive("foo"); - proxy.sensitive("hello"); + proxy.sensitive("foo", "hello"); + proxy.header("bar", "hello"); + proxy.header("abc", "123"); + proxy.header("hello", "world"); + return proxy.uri(home.toString() + "/headers").get(); + } + + @GetMapping("/proxy/sensitive-headers-default") + public Mono>>> defaultSensitiveHeaders( + ProxyExchange>> proxy) { proxy.header("bar", "hello"); proxy.header("abc", "123"); proxy.header("hello", "world");