diff --git a/spring-cloud-gateway-server/src/main/java/org/springframework/cloud/gateway/filter/factory/SecureHeadersGatewayFilterFactory.java b/spring-cloud-gateway-server/src/main/java/org/springframework/cloud/gateway/filter/factory/SecureHeadersGatewayFilterFactory.java index 5d672c67..702ef933 100644 --- a/spring-cloud-gateway-server/src/main/java/org/springframework/cloud/gateway/filter/factory/SecureHeadersGatewayFilterFactory.java +++ b/spring-cloud-gateway-server/src/main/java/org/springframework/cloud/gateway/filter/factory/SecureHeadersGatewayFilterFactory.java @@ -92,39 +92,39 @@ public class SecureHeadersGatewayFilterFactory List disabled = properties.getDisable(); Config config = originalConfig.withDefaults(properties); - if (isEnabled(disabled, X_XSS_PROTECTION_HEADER)) { - headers.add(X_XSS_PROTECTION_HEADER, config.getXssProtectionHeader()); - } + return chain.filter(exchange).then(Mono.fromRunnable(() -> { + if (isEnabled(disabled, X_XSS_PROTECTION_HEADER)) { + headers.addIfAbsent(X_XSS_PROTECTION_HEADER, config.getXssProtectionHeader()); + } - if (isEnabled(disabled, STRICT_TRANSPORT_SECURITY_HEADER)) { - headers.add(STRICT_TRANSPORT_SECURITY_HEADER, config.getStrictTransportSecurity()); - } + if (isEnabled(disabled, STRICT_TRANSPORT_SECURITY_HEADER)) { + headers.addIfAbsent(STRICT_TRANSPORT_SECURITY_HEADER, config.getStrictTransportSecurity()); + } - if (isEnabled(disabled, X_FRAME_OPTIONS_HEADER)) { - headers.add(X_FRAME_OPTIONS_HEADER, config.getFrameOptions()); - } + if (isEnabled(disabled, X_FRAME_OPTIONS_HEADER)) { + headers.addIfAbsent(X_FRAME_OPTIONS_HEADER, config.getFrameOptions()); + } - if (isEnabled(disabled, X_CONTENT_TYPE_OPTIONS_HEADER)) { - headers.add(X_CONTENT_TYPE_OPTIONS_HEADER, config.getContentTypeOptions()); - } + if (isEnabled(disabled, X_CONTENT_TYPE_OPTIONS_HEADER)) { + headers.addIfAbsent(X_CONTENT_TYPE_OPTIONS_HEADER, config.getContentTypeOptions()); + } - if (isEnabled(disabled, REFERRER_POLICY_HEADER)) { - headers.add(REFERRER_POLICY_HEADER, config.getReferrerPolicy()); - } + if (isEnabled(disabled, REFERRER_POLICY_HEADER)) { + headers.addIfAbsent(REFERRER_POLICY_HEADER, config.getReferrerPolicy()); + } - if (isEnabled(disabled, CONTENT_SECURITY_POLICY_HEADER)) { - headers.add(CONTENT_SECURITY_POLICY_HEADER, config.getContentSecurityPolicy()); - } + if (isEnabled(disabled, CONTENT_SECURITY_POLICY_HEADER)) { + headers.addIfAbsent(CONTENT_SECURITY_POLICY_HEADER, config.getContentSecurityPolicy()); + } - if (isEnabled(disabled, X_DOWNLOAD_OPTIONS_HEADER)) { - headers.add(X_DOWNLOAD_OPTIONS_HEADER, config.getDownloadOptions()); - } + if (isEnabled(disabled, X_DOWNLOAD_OPTIONS_HEADER)) { + headers.addIfAbsent(X_DOWNLOAD_OPTIONS_HEADER, config.getDownloadOptions()); + } - if (isEnabled(disabled, X_PERMITTED_CROSS_DOMAIN_POLICIES_HEADER)) { - headers.add(X_PERMITTED_CROSS_DOMAIN_POLICIES_HEADER, config.getPermittedCrossDomainPolicies()); - } - - return chain.filter(exchange); + if (isEnabled(disabled, X_PERMITTED_CROSS_DOMAIN_POLICIES_HEADER)) { + headers.addIfAbsent(X_PERMITTED_CROSS_DOMAIN_POLICIES_HEADER, config.getPermittedCrossDomainPolicies()); + } + })); } @Override diff --git a/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/filter/factory/SecureHeadersGatewayFilterFactoryTests.java b/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/filter/factory/SecureHeadersGatewayFilterFactoryTests.java index 84c898cd..7d4c72bd 100644 --- a/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/filter/factory/SecureHeadersGatewayFilterFactoryTests.java +++ b/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/filter/factory/SecureHeadersGatewayFilterFactoryTests.java @@ -28,6 +28,7 @@ import org.springframework.cloud.gateway.test.BaseWebClientTests; import org.springframework.context.annotation.Import; import org.springframework.http.HttpHeaders; import org.springframework.http.HttpStatus; +import org.springframework.http.MediaType; import org.springframework.test.annotation.DirtiesContext; import org.springframework.test.context.junit4.SpringRunner; import org.springframework.web.reactive.function.client.ClientResponse; @@ -73,6 +74,18 @@ public class SecureHeadersGatewayFilterFactoryTests extends BaseWebClientTests { }).expectComplete().verify(DURATION); } + @Test + public void addsSecureHeadersAfterResponseIsReceived() { + Mono result = webClient.patch().uri("/headers").header("Host", "www.secureheaders.org") + .contentType(MediaType.APPLICATION_JSON).bodyValue("{ \"X-Frame-Options\": \"sameorigin\" }") + .exchangeToMono(Mono::just); + + StepVerifier.create(result).consumeNextWith(response -> { + assertThat(response.statusCode()).isEqualTo(HttpStatus.OK); + assertThat(response.headers().header(X_FRAME_OPTIONS_HEADER)).containsOnly("sameorigin"); + }).expectComplete().verify(DURATION); + } + @EnableAutoConfiguration @SpringBootConfiguration @Import(DefaultTestConfig.class) diff --git a/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/filter/factory/SecureHeadersGatewayFilterFactoryUnitTests.java b/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/filter/factory/SecureHeadersGatewayFilterFactoryUnitTests.java index 68ad726f..b99b23eb 100644 --- a/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/filter/factory/SecureHeadersGatewayFilterFactoryUnitTests.java +++ b/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/filter/factory/SecureHeadersGatewayFilterFactoryUnitTests.java @@ -72,9 +72,9 @@ public class SecureHeadersGatewayFilterFactoryUnitTests { new SecureHeadersProperties()); filter = filterFactory.apply(new Config()); - filter.filter(exchange, filterChain); + filter.filter(exchange, filterChain).block(); - ServerHttpResponse response = captor.getValue().getResponse(); + ServerHttpResponse response = exchange.getResponse(); assertThat(response.getHeaders()).containsKeys(X_XSS_PROTECTION_HEADER, STRICT_TRANSPORT_SECURITY_HEADER, X_FRAME_OPTIONS_HEADER, X_CONTENT_TYPE_OPTIONS_HEADER, REFERRER_POLICY_HEADER, CONTENT_SECURITY_POLICY_HEADER, X_DOWNLOAD_OPTIONS_HEADER, X_PERMITTED_CROSS_DOMAIN_POLICIES_HEADER); @@ -90,7 +90,7 @@ public class SecureHeadersGatewayFilterFactoryUnitTests { SecureHeadersGatewayFilterFactory filterFactory = new SecureHeadersGatewayFilterFactory(properties); filter = filterFactory.apply(new Config()); - filter.filter(exchange, filterChain); + filter.filter(exchange, filterChain).block(); ServerHttpResponse response = captor.getValue().getResponse(); assertThat(response.getHeaders()).doesNotContainKeys(X_XSS_PROTECTION_HEADER, STRICT_TRANSPORT_SECURITY_HEADER, @@ -109,9 +109,9 @@ public class SecureHeadersGatewayFilterFactoryUnitTests { config.setReferrerPolicy("referrer"); filter = filterFactory.apply(config); - filter.filter(exchange, filterChain); + filter.filter(exchange, filterChain).block(); - ServerHttpResponse response = captor.getValue().getResponse(); + ServerHttpResponse response = exchange.getResponse(); assertThat(response.getHeaders()).containsKeys(X_XSS_PROTECTION_HEADER, STRICT_TRANSPORT_SECURITY_HEADER, X_FRAME_OPTIONS_HEADER, X_CONTENT_TYPE_OPTIONS_HEADER, REFERRER_POLICY_HEADER, CONTENT_SECURITY_POLICY_HEADER, X_DOWNLOAD_OPTIONS_HEADER, X_PERMITTED_CROSS_DOMAIN_POLICIES_HEADER); @@ -132,6 +132,31 @@ public class SecureHeadersGatewayFilterFactoryUnitTests { } + @Test + public void doesNotDuplicateHeaders() { + String originalHeaderValue = "original-header-value"; + SecureHeadersGatewayFilterFactory filterFactory = new SecureHeadersGatewayFilterFactory( + new SecureHeadersProperties()); + Config config = new Config(); + + String[] headers = { X_XSS_PROTECTION_HEADER, STRICT_TRANSPORT_SECURITY_HEADER, X_FRAME_OPTIONS_HEADER, + X_CONTENT_TYPE_OPTIONS_HEADER, REFERRER_POLICY_HEADER, CONTENT_SECURITY_POLICY_HEADER, + X_DOWNLOAD_OPTIONS_HEADER, X_PERMITTED_CROSS_DOMAIN_POLICIES_HEADER }; + + for (String header : headers) { + filter = filterFactory.apply(config); + + MockServerHttpRequest request = MockServerHttpRequest.get("http://localhost").build(); + exchange = MockServerWebExchange.from(request); + exchange.getResponse().getHeaders().set(header, originalHeaderValue); + + filter.filter(exchange, filterChain).block(); + + ServerHttpResponse response = captor.getValue().getResponse(); + assertThat(response.getHeaders().get(header)).containsOnly(originalHeaderValue); + } + } + @Test public void toStringFormat() { GatewayFilter filter = new SecureHeadersGatewayFilterFactory(new SecureHeadersProperties()).apply(new Config()); diff --git a/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/test/HttpBinCompatibleController.java b/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/test/HttpBinCompatibleController.java index f53ed868..381f6ec3 100644 --- a/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/test/HttpBinCompatibleController.java +++ b/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/test/HttpBinCompatibleController.java @@ -73,6 +73,17 @@ public class HttpBinCompatibleController { return result; } + @RequestMapping(path = "/headers", method = RequestMethod.PATCH) + public ResponseEntity> headersPatch(ServerWebExchange exchange, + @RequestBody Map headersToAdd) { + Map result = new HashMap<>(); + result.put("headers", getHeaders(exchange)); + ResponseEntity.BodyBuilder responseEntity = ResponseEntity.status(HttpStatus.OK); + headersToAdd.forEach(responseEntity::header); + + return responseEntity.body(result); + } + @RequestMapping(path = "/multivalueheaders", method = { RequestMethod.GET, RequestMethod.POST }, produces = MediaType.APPLICATION_JSON_VALUE) public Map multiValueHeaders(ServerWebExchange exchange) {