Adds Forwarded By Header

Fixes gh-2658
This commit is contained in:
Tillmann Heigel
2022-06-30 10:29:09 +02:00
committed by spencergibb
parent f6f90df404
commit c278a43108
4 changed files with 112 additions and 20 deletions

View File

@@ -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;
}
}

View File

@@ -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");
}
}

View File

@@ -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();

View File

@@ -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()