diff --git a/spring-cloud-gateway-server/src/main/java/org/springframework/cloud/gateway/filter/headers/ForwardedHeadersFilter.java b/spring-cloud-gateway-server/src/main/java/org/springframework/cloud/gateway/filter/headers/ForwardedHeadersFilter.java index b95cb191..54d7c297 100644 --- a/spring-cloud-gateway-server/src/main/java/org/springframework/cloud/gateway/filter/headers/ForwardedHeadersFilter.java +++ b/spring-cloud-gateway-server/src/main/java/org/springframework/cloud/gateway/filter/headers/ForwardedHeadersFilter.java @@ -20,11 +20,17 @@ import java.net.Inet6Address; import java.net.InetAddress; import java.net.InetSocketAddress; import java.net.URI; +import java.net.UnknownHostException; import java.util.ArrayList; import java.util.HashMap; import java.util.List; import java.util.Map; +import org.apache.commons.logging.Log; +import org.apache.commons.logging.LogFactory; + +import org.springframework.beans.factory.annotation.Value; +import org.springframework.boot.context.properties.ConfigurationProperties; import org.springframework.core.Ordered; import org.springframework.http.HttpHeaders; import org.springframework.http.server.reactive.ServerHttpRequest; @@ -34,8 +40,21 @@ import org.springframework.util.ObjectUtils; import org.springframework.util.StringUtils; import org.springframework.web.server.ServerWebExchange; +/** + * @author Olga Maciaszek-Sharma + * @author Tillmann Heigel + */ +@ConfigurationProperties("spring.cloud.gateway.forwarded") public class ForwardedHeadersFilter implements HttpHeadersFilter, Ordered { + @Value("${server.port}") + private int serverPort; + + private final Log logger = LogFactory.getLog(getClass()); + + /** If Forwarded: by header is enabled. */ + private boolean byEnabled = true; + /** * Forwarded header. */ @@ -120,8 +139,7 @@ public class ForwardedHeadersFilter implements HttpHeadersFilter, Ordered { String forValue; if (remoteAddress.isUnresolved()) { forValue = remoteAddress.getHostName(); - } - else { + } else { InetAddress address = remoteAddress.getAddress(); forValue = remoteAddress.getAddress().getHostAddress(); if (address instanceof Inet6Address) { @@ -134,13 +152,37 @@ public class ForwardedHeadersFilter implements HttpHeadersFilter, Ordered { } forwarded.put("for", forValue); } - // TODO: support by? + + if (byEnabled) { + addForwardedByHeader(forwarded); + } updated.add(FORWARDED_HEADER, forwarded.toHeaderValue()); return updated; } + private void addForwardedByHeader(Forwarded forwarded) { + try { + addForwardedBy(forwarded, InetAddress.getLocalHost()); + } catch (UnknownHostException e) { + this.logger.warn("Can not resolve host address, skipping Forwarded 'by' header", e); + } + } + + /* visible for testing */ void addForwardedBy(Forwarded forwarded, InetAddress localAddress) { + if (localAddress != null) { + String byValue = localAddress.getHostAddress(); + if (localAddress instanceof Inet6Address) { + byValue = "[" + byValue + "]"; + } + if (serverPort > 0) { + byValue = byValue + ":" + serverPort; + } + forwarded.put("by", byValue); + } + } + /* for testing */ static class Forwarded { private static final char EQUALS = '='; @@ -195,4 +237,12 @@ public class ForwardedHeadersFilter implements HttpHeadersFilter, Ordered { } + public boolean isByEnabled() { + return byEnabled; + } + + public void setByEnabled(boolean byEnabled) { + this.byEnabled = byEnabled; + } + } diff --git a/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/filter/headers/ForwardedHeadersFilterTests.java b/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/filter/headers/ForwardedHeadersFilterTests.java index 73cb0003..ce2d57ee 100644 --- a/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/filter/headers/ForwardedHeadersFilterTests.java +++ b/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/filter/headers/ForwardedHeadersFilterTests.java @@ -16,6 +16,9 @@ package org.springframework.cloud.gateway.filter.headers; +import static org.assertj.core.api.Assertions.assertThat; +import static org.springframework.cloud.gateway.filter.headers.ForwardedHeadersFilter.FORWARDED_HEADER; + import java.net.InetAddress; import java.net.InetSocketAddress; import java.net.UnknownHostException; @@ -25,6 +28,7 @@ import java.util.HashMap; import java.util.List; import java.util.Map; +import org.assertj.core.api.Assertions; import org.junit.jupiter.api.Test; import org.springframework.cloud.gateway.filter.headers.ForwardedHeadersFilter.Forwarded; @@ -33,9 +37,6 @@ 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.ForwardedHeadersFilter.FORWARDED_HEADER; - /** * @author Spencer Gibb */ @@ -201,4 +202,42 @@ public class ForwardedHeadersFilterTests { } } + @Test + public void forwardedByForIpv4AddressIsAdded() throws UnknownHostException { + Forwarded forwarded = new Forwarded(); + InetAddress ipv4Address = InetAddress.getByName("216.103.69.111"); + ForwardedHeadersFilter forwardedHeadersFilter = new ForwardedHeadersFilter(); + forwardedHeadersFilter.setByEnabled(true); + + forwardedHeadersFilter.addForwardedBy(forwarded, ipv4Address); + + Assertions.assertThat(forwarded.getValues()).containsEntry("by", "216.103.69.111"); + + } + + @Test + public void forwardedByForIpv6AddressIsAdded() throws UnknownHostException { + Forwarded forwarded = new Forwarded(); + InetAddress ipv6Address = InetAddress.getByName("abc4:babf:955f:1724:11bc:0153:275c:d36e"); + ForwardedHeadersFilter forwardedHeadersFilter = new ForwardedHeadersFilter(); + forwardedHeadersFilter.setByEnabled(true); + + forwardedHeadersFilter.addForwardedBy(forwarded, ipv6Address); + + Assertions.assertThat(forwarded.getValues()).containsEntry("by", + "\"[abc4:babf:955f:1724:11bc:153:275c:d36e]\""); + + } + + @Test + public void forwardedByIsNotAddedIfFeatureIsDisabled() throws UnknownHostException { + Forwarded forwarded = new Forwarded(); + InetAddress ipv4Address = InetAddress.getByName("216.103.69.111"); + ForwardedHeadersFilter forwardedHeadersFilter = new ForwardedHeadersFilter(); + forwardedHeadersFilter.setByEnabled(false); + + forwardedHeadersFilter.addForwardedBy(forwarded, ipv4Address); + + Assertions.assertThat(forwarded.getValues()).containsEntry("by", "216.103.69.111"); + } } diff --git a/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/filter/headers/XForwardedHeadersFilterTests.java b/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/filter/headers/XForwardedHeadersFilterTests.java index 5bfcd7e8..b899986c 100644 --- a/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/filter/headers/XForwardedHeadersFilterTests.java +++ b/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/filter/headers/XForwardedHeadersFilterTests.java @@ -16,6 +16,15 @@ package org.springframework.cloud.gateway.filter.headers; +import static org.assertj.core.api.Assertions.assertThat; +import static org.springframework.cloud.gateway.filter.headers.XForwardedHeadersFilter.X_FORWARDED_FOR_HEADER; +import static org.springframework.cloud.gateway.filter.headers.XForwardedHeadersFilter.X_FORWARDED_HOST_HEADER; +import static org.springframework.cloud.gateway.filter.headers.XForwardedHeadersFilter.X_FORWARDED_PORT_HEADER; +import static org.springframework.cloud.gateway.filter.headers.XForwardedHeadersFilter.X_FORWARDED_PREFIX_HEADER; +import static org.springframework.cloud.gateway.filter.headers.XForwardedHeadersFilter.X_FORWARDED_PROTO_HEADER; +import static org.springframework.cloud.gateway.support.ServerWebExchangeUtils.GATEWAY_ORIGINAL_REQUEST_URL_ATTR; +import static org.springframework.cloud.gateway.support.ServerWebExchangeUtils.GATEWAY_REQUEST_URL_ATTR; + import java.net.InetAddress; import java.net.InetSocketAddress; import java.net.URI; @@ -29,22 +38,13 @@ import org.springframework.mock.web.server.MockServerWebExchange; import org.springframework.web.server.ServerWebExchange; import org.springframework.web.util.UriComponentsBuilder; -import static org.assertj.core.api.Assertions.assertThat; -import static org.springframework.cloud.gateway.filter.headers.XForwardedHeadersFilter.X_FORWARDED_FOR_HEADER; -import static org.springframework.cloud.gateway.filter.headers.XForwardedHeadersFilter.X_FORWARDED_HOST_HEADER; -import static org.springframework.cloud.gateway.filter.headers.XForwardedHeadersFilter.X_FORWARDED_PORT_HEADER; -import static org.springframework.cloud.gateway.filter.headers.XForwardedHeadersFilter.X_FORWARDED_PREFIX_HEADER; -import static org.springframework.cloud.gateway.filter.headers.XForwardedHeadersFilter.X_FORWARDED_PROTO_HEADER; -import static org.springframework.cloud.gateway.support.ServerWebExchangeUtils.GATEWAY_ORIGINAL_REQUEST_URL_ATTR; -import static org.springframework.cloud.gateway.support.ServerWebExchangeUtils.GATEWAY_REQUEST_URL_ATTR; - /** * @author Spencer Gibb */ public class XForwardedHeadersFilterTests { @Test - public void remoteAddressIsNull() throws Exception { + public void remoteAddressIsNull() { MockServerHttpRequest request = MockServerHttpRequest.get("http://localhost:8080/get") .header(HttpHeaders.HOST, "myhost") .build(); diff --git a/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/test/GatewayIntegrationTests.java b/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/test/GatewayIntegrationTests.java index f998d8e0..b91e9693 100644 --- a/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/test/GatewayIntegrationTests.java +++ b/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/test/GatewayIntegrationTests.java @@ -16,6 +16,10 @@ package org.springframework.cloud.gateway.test; +import static org.assertj.core.api.Assertions.assertThat; +import static org.springframework.boot.test.context.SpringBootTest.WebEnvironment.RANDOM_PORT; +import static org.springframework.cloud.gateway.test.TestUtils.getMap; + import java.util.ArrayList; import java.util.List; import java.util.Map; @@ -58,9 +62,7 @@ import org.springframework.web.bind.annotation.GetMapping; import org.springframework.web.bind.annotation.RestController; import org.springframework.web.server.ServerWebExchange; -import static org.assertj.core.api.Assertions.assertThat; -import static org.springframework.boot.test.context.SpringBootTest.WebEnvironment.RANDOM_PORT; -import static org.springframework.cloud.gateway.test.TestUtils.getMap; +import reactor.core.publisher.Mono; @SpringBootTest(webEnvironment = RANDOM_PORT) @DirtiesContext @@ -125,7 +127,8 @@ class GatewayIntegrationTests extends BaseWebClientTests { assertThat(headers.get(ForwardedHeadersFilter.FORWARDED_HEADER)).asString() .contains("proto=http") .contains("host=\"localhost:") - .contains("for=\"127.0.0.1:"); + .contains("for=\"127.0.0.1:") + .contains("by="); assertThat(headers.get(XForwardedHeadersFilter.X_FORWARDED_HOST_HEADER)).asString() .isEqualTo("localhost:" + this.port); assertThat(headers.get(XForwardedHeadersFilter.X_FORWARDED_PORT_HEADER)).asString()