Don't decode query parameters.

When query parameters or paths of request uri are mutated make sure
they are not decoded.

This means bypassing the default ServerHttpRequest.mutate()

fixes gh-147
This commit is contained in:
Spencer Gibb
2018-01-30 00:55:07 -05:00
parent 84bced0411
commit 7f92ea1213
9 changed files with 287 additions and 21 deletions

View File

@@ -26,6 +26,7 @@ import org.springframework.http.server.reactive.ServerHttpRequest;
import org.springframework.tuple.Tuple;
import org.springframework.util.StringUtils;
import org.springframework.cloud.gateway.filter.GatewayFilter;
import org.springframework.web.util.UriComponentsBuilder;
/**
* @author Spencer Gibb
@@ -49,7 +50,7 @@ public class AddRequestParameterGatewayFilterFactory implements GatewayFilterFac
URI uri = exchange.getRequest().getURI();
StringBuilder query = new StringBuilder();
String originalQuery = uri.getQuery();
String originalQuery = uri.getRawQuery();
if (StringUtils.hasText(originalQuery)) {
query.append(originalQuery);
@@ -64,13 +65,15 @@ public class AddRequestParameterGatewayFilterFactory implements GatewayFilterFac
query.append(value);
try {
URI newUri = new URI(uri.getScheme(), uri.getUserInfo(), uri.getHost(), uri.getPort(),
uri.getPath(), query.toString(), uri.getFragment());
URI newUri = UriComponentsBuilder.fromUri(uri)
.replaceQuery(query.toString())
.build(true)
.toUri();
ServerHttpRequest request = exchange.getRequest().mutate().uri(newUri).build();
ServerHttpRequest request = mutate(exchange.getRequest()).uri(newUri).build();
return chain.filter(exchange.mutate().request(request).build());
} catch (URISyntaxException ex) {
} catch (RuntimeException ex) {
throw new IllegalStateException("Invalid URI query: \"" + query.toString() + "\"");
}
};

View File

@@ -19,7 +19,9 @@ package org.springframework.cloud.gateway.filter.factory;
import org.springframework.cloud.gateway.filter.GatewayFilter;
import org.springframework.cloud.gateway.support.ArgumentHints;
import org.springframework.cloud.gateway.support.GatewayServerHttpRequestBuilder;
import org.springframework.cloud.gateway.support.NameUtils;
import org.springframework.http.server.reactive.ServerHttpRequest;
import org.springframework.tuple.Tuple;
/**
@@ -36,4 +38,9 @@ public interface GatewayFilterFactory extends ArgumentHints {
default String name() {
return NameUtils.normalizeFilterName(getClass());
}
default ServerHttpRequest.Builder mutate(ServerHttpRequest request) {
return new GatewayServerHttpRequestBuilder(request);
}
}

View File

@@ -148,6 +148,7 @@ public class HystrixGatewayFilterFactory implements GatewayFilterFactory {
//TODO: copied from RouteToRequestUrlFilter
URI uri = exchange.getRequest().getURI();
//TODO: assume always?
boolean encoded = containsEncodedQuery(uri);
URI requestUrl = UriComponentsBuilder.fromUri(uri)
.host(null)
@@ -157,7 +158,7 @@ public class HystrixGatewayFilterFactory implements GatewayFilterFactory {
.toUri();
exchange.getAttributes().put(GATEWAY_REQUEST_URL_ATTR, requestUrl);
ServerHttpRequest request = this.exchange.getRequest().mutate().uri(requestUrl).build();
ServerHttpRequest request = mutate(this.exchange.getRequest()).uri(requestUrl).build();
ServerWebExchange mutated = exchange.mutate().request(request).build();
return RxReactiveStreams.toObservable(HystrixGatewayFilterFactory.this.dispatcherHandler.handle(mutated));
}

View File

@@ -55,7 +55,7 @@ public class PrefixPathGatewayFilterFactory implements GatewayFilterFactory {
addOriginalRequestUrl(exchange, req.getURI());
String newPath = prefix + req.getURI().getPath();
ServerHttpRequest request = req.mutate()
ServerHttpRequest request = mutate(req)
.path(newPath)
.build();

View File

@@ -20,9 +20,9 @@ package org.springframework.cloud.gateway.filter.factory;
import java.util.Arrays;
import java.util.List;
import org.springframework.cloud.gateway.filter.GatewayFilter;
import org.springframework.http.server.reactive.ServerHttpRequest;
import org.springframework.tuple.Tuple;
import org.springframework.cloud.gateway.filter.GatewayFilter;
import static org.springframework.cloud.gateway.support.ServerWebExchangeUtils.GATEWAY_REQUEST_URL_ATTR;
import static org.springframework.cloud.gateway.support.ServerWebExchangeUtils.addOriginalRequestUrl;
@@ -54,7 +54,7 @@ public class RewritePathGatewayFilterFactory implements GatewayFilterFactory {
String path = req.getURI().getPath();
String newPath = path.replaceAll(regex, replacement);
ServerHttpRequest request = req.mutate()
ServerHttpRequest request = mutate(req)
.path(newPath)
.build();

View File

@@ -73,7 +73,7 @@ public class SetPathGatewayFilterFactory implements GatewayFilterFactory {
exchange.getAttributes().put(GATEWAY_REQUEST_URL_ATTR, uri);
ServerHttpRequest request = req.mutate()
ServerHttpRequest request = mutate(req)
.path(newPath)
.build();

View File

@@ -0,0 +1,210 @@
package org.springframework.cloud.gateway.support;
import java.net.InetSocketAddress;
import java.net.URI;
import java.util.LinkedList;
import java.util.List;
import java.util.Map;
import java.util.function.Consumer;
import org.springframework.core.io.buffer.DataBuffer;
import org.springframework.http.HttpCookie;
import org.springframework.http.HttpHeaders;
import org.springframework.http.HttpMethod;
import org.springframework.http.server.reactive.AbstractServerHttpRequest;
import org.springframework.http.server.reactive.ServerHttpRequest;
import org.springframework.http.server.reactive.SslInfo;
import org.springframework.lang.Nullable;
import org.springframework.util.Assert;
import org.springframework.util.LinkedMultiValueMap;
import org.springframework.util.MultiValueMap;
import org.springframework.web.util.UriComponentsBuilder;
import reactor.core.publisher.Flux;
/**
* Package-private default implementation of {@link ServerHttpRequest.Builder}.
*
* @author Rossen Stoyanchev
* @author Sebastien Deleuze
* @since 5.0
*/
public class GatewayServerHttpRequestBuilder implements ServerHttpRequest.Builder {
private boolean encoded;
private URI uri;
private HttpHeaders httpHeaders;
private String httpMethodValue;
private final MultiValueMap<String, HttpCookie> cookies;
@Nullable
private String uriPath;
@Nullable
private String contextPath;
private Flux<DataBuffer> body;
private final ServerHttpRequest originalRequest;
public GatewayServerHttpRequestBuilder(ServerHttpRequest original) {
this(original, true);
}
public GatewayServerHttpRequestBuilder(ServerHttpRequest original, boolean encoded) {
Assert.notNull(original, "ServerHttpRequest is required");
this.uri = original.getURI();
this.httpMethodValue = original.getMethodValue();
this.body = original.getBody();
this.httpHeaders = new HttpHeaders();
copyMultiValueMap(original.getHeaders(), this.httpHeaders);
this.cookies = new LinkedMultiValueMap<>(original.getCookies().size());
copyMultiValueMap(original.getCookies(), this.cookies);
this.originalRequest = original;
this.encoded = encoded;
}
private static <K, V> void copyMultiValueMap(MultiValueMap<K,V> source,
MultiValueMap<K,V> destination) {
for (Map.Entry<K, List<V>> entry : source.entrySet()) {
K key = entry.getKey();
List<V> values = new LinkedList<>(entry.getValue());
destination.put(key, values);
}
}
@Override
public ServerHttpRequest.Builder method(HttpMethod httpMethod) {
this.httpMethodValue = httpMethod.name();
return this;
}
@Override
public ServerHttpRequest.Builder uri(URI uri) {
this.uri = uri;
return this;
}
@Override
public ServerHttpRequest.Builder path(String path) {
this.uriPath = path;
return this;
}
@Override
public ServerHttpRequest.Builder contextPath(String contextPath) {
this.contextPath = contextPath;
return this;
}
@Override
public ServerHttpRequest.Builder header(String key, String value) {
this.httpHeaders.add(key, value);
return this;
}
@Override
public ServerHttpRequest.Builder headers(Consumer<HttpHeaders> headersConsumer) {
Assert.notNull(headersConsumer, "'headersConsumer' must not be null");
headersConsumer.accept(this.httpHeaders);
return this;
}
@Override
public ServerHttpRequest build() {
URI uriToUse = getUriToUse();
return new GatewayServerHttpRequest(uriToUse, this.contextPath, this.httpHeaders,
this.httpMethodValue, this.cookies, this.body, this.originalRequest);
}
private URI getUriToUse() {
if (this.uriPath == null) {
return this.uri;
}
try {
return UriComponentsBuilder.fromUri(this.uri)
.replacePath(uriPath)
.build(encoded).toUri();
}
catch (RuntimeException ex) {
throw new IllegalStateException("Invalid URI path: \"" + this.uriPath + "\"");
}
}
private static class GatewayServerHttpRequest extends AbstractServerHttpRequest {
private final String methodValue;
private final MultiValueMap<String, HttpCookie> cookies;
@Nullable
private final InetSocketAddress remoteAddress;
@Nullable
private final SslInfo sslInfo;
private final Flux<DataBuffer> body;
private final ServerHttpRequest originalRequest;
public GatewayServerHttpRequest(URI uri, @Nullable String contextPath,
HttpHeaders headers, String methodValue, MultiValueMap<String, HttpCookie> cookies,
Flux<DataBuffer> body, ServerHttpRequest originalRequest) {
super(uri, contextPath, headers);
this.methodValue = methodValue;
this.cookies = cookies;
this.remoteAddress = originalRequest.getRemoteAddress();
this.sslInfo = originalRequest.getSslInfo();
this.body = body;
this.originalRequest = originalRequest;
}
@Override
public String getMethodValue() {
return this.methodValue;
}
@Override
protected MultiValueMap<String, HttpCookie> initCookies() {
return this.cookies;
}
@Nullable
@Override
public InetSocketAddress getRemoteAddress() {
return this.remoteAddress;
}
@Nullable
@Override
protected SslInfo initSslInfo() {
return this.sslInfo;
}
@Override
public Flux<DataBuffer> getBody() {
return this.body;
}
@SuppressWarnings("unchecked")
@Override
public <T> T getNativeRequest() {
return (T) this.originalRequest;
}
}
}

View File

@@ -17,6 +17,11 @@
package org.springframework.cloud.gateway.filter.factory;
import java.io.UnsupportedEncodingException;
import java.net.URI;
import java.net.URLDecoder;
import java.util.Map;
import org.junit.Test;
import org.junit.runner.RunWith;
import org.springframework.boot.SpringBootConfiguration;
@@ -27,16 +32,17 @@ import org.springframework.context.annotation.Import;
import org.springframework.test.annotation.DirtiesContext;
import org.springframework.test.context.ActiveProfiles;
import org.springframework.test.context.junit4.SpringRunner;
import reactor.core.publisher.Mono;
import reactor.test.StepVerifier;
import java.util.Map;
import org.springframework.web.util.UriComponentsBuilder;
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.support.ServerWebExchangeUtils.containsEncodedQuery;
import static org.springframework.cloud.gateway.test.TestUtils.getMap;
import static org.springframework.web.reactive.function.BodyExtractors.toMono;
import reactor.core.publisher.Mono;
import reactor.test.StepVerifier;
@RunWith(SpringRunner.class)
@SpringBootTest(webEnvironment = RANDOM_PORT)
@DirtiesContext
@@ -45,17 +51,30 @@ public class AddRequestParameterGatewayFilterFactoryTests extends BaseWebClientT
@Test
public void addRequestParameterFilterWorksBlankQuery() {
testRequestParameterFilter("");
testRequestParameterFilter(null, null);
}
@Test
public void addRequestParameterFilterWorksNonBlankQuery() {
testRequestParameterFilter("?baz=bam");
testRequestParameterFilter("baz", "bam");
}
private void testRequestParameterFilter(String query) {
@Test
public void addRequestParameterFilterWorksEncodedQuery() {
testRequestParameterFilter("name", "%E6%89%8E%E6%A0%B9");
}
private void testRequestParameterFilter(String name, String value) {
String query;
if (name != null) {
query = "?" + name + "=" + value;
} else {
query = "";
}
URI uri = UriComponentsBuilder.fromUriString(this.baseUri+"/get" + query).build(true).toUri();
boolean checkForEncodedValue = containsEncodedQuery(uri);
Mono<Map> result = webClient.get()
.uri("/get" + query)
.uri(uri)
.header("Host", "www.addrequestparameter.org")
.exchange()
.flatMap(response -> response.body(toMono(Map.class)));
@@ -65,6 +84,17 @@ public class AddRequestParameterGatewayFilterFactoryTests extends BaseWebClientT
response -> {
Map<String, Object> args = getMap(response, "args");
assertThat(args).containsEntry("foo", "bar");
if (name != null) {
if (checkForEncodedValue) {
try {
assertThat(args).containsEntry(name, URLDecoder.decode(value, "UTF-8"));
} catch (UnsupportedEncodingException e) {
throw new RuntimeException(e);
}
} else {
assertThat(args).containsEntry(name, value);
}
}
})
.expectComplete()
.verify(DURATION);

View File

@@ -23,9 +23,12 @@ import java.util.LinkedHashSet;
import org.junit.Test;
import org.mockito.ArgumentCaptor;
import org.springframework.cloud.gateway.filter.GatewayFilter;
import org.springframework.cloud.gateway.filter.GatewayFilterChain;
import org.springframework.http.HttpMethod;
import org.springframework.mock.http.server.reactive.MockServerHttpRequest;
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.mockito.Mockito.mock;
@@ -36,7 +39,6 @@ import static org.springframework.cloud.gateway.support.ServerWebExchangeUtils.G
import static org.springframework.cloud.gateway.support.ServerWebExchangeUtils.GATEWAY_REQUEST_URL_ATTR;
import static org.springframework.tuple.TupleBuilder.tuple;
import org.springframework.cloud.gateway.filter.GatewayFilterChain;
import reactor.core.publisher.Mono;
/**
@@ -54,11 +56,12 @@ public class RewritePathGatewayFilterFactoryTests {
testRewriteFilter("/foo/(?<id>\\d.*)", "/bar/baz/$\\{id}", "/foo/123", "/bar/baz/123");
}
private void testRewriteFilter(String regex, String replacement, String actualPath, String expectedPath) {
private ServerWebExchange testRewriteFilter(String regex, String replacement, String actualPath, String expectedPath) {
GatewayFilter filter = new RewritePathGatewayFilterFactory().apply(tuple().of(REGEXP_KEY, regex, REPLACEMENT_KEY, replacement));
URI url = UriComponentsBuilder.fromUriString("http://localhost"+ actualPath).build(true).toUri();
MockServerHttpRequest request = MockServerHttpRequest
.get("http://localhost"+ actualPath)
.method(HttpMethod.GET, url)
.build();
ServerWebExchange exchange = MockServerWebExchange.from(request);
@@ -78,5 +81,17 @@ public class RewritePathGatewayFilterFactoryTests {
assertThat(requestUrl).hasScheme("http").hasHost("localhost").hasNoPort().hasPath(expectedPath);
LinkedHashSet<URI> uris = webExchange.getRequiredAttribute(GATEWAY_ORIGINAL_REQUEST_URL_ATTR);
assertThat(uris).contains(request.getURI());
return webExchange;
}
@Test
public void rewritePathWithEncodedParams() {
ServerWebExchange exchange = testRewriteFilter("/foo", "/baz",
"/foo/bar?name=%E6%89%8E%E6%A0%B9",
"/baz/bar");
URI uri = exchange.getRequest().getURI();
assertThat(uri.getRawQuery()).isEqualTo("name=%E6%89%8E%E6%A0%B9");
}
}