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:
@@ -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() + "\"");
|
||||
}
|
||||
};
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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));
|
||||
}
|
||||
|
||||
@@ -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();
|
||||
|
||||
|
||||
@@ -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();
|
||||
|
||||
|
||||
@@ -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();
|
||||
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
}
|
||||
|
||||
}
|
||||
@@ -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);
|
||||
|
||||
@@ -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");
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user