diff --git a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/factory/AddRequestParameterGatewayFilterFactory.java b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/factory/AddRequestParameterGatewayFilterFactory.java index dbf41643..7a6ba84c 100644 --- a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/factory/AddRequestParameterGatewayFilterFactory.java +++ b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/factory/AddRequestParameterGatewayFilterFactory.java @@ -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() + "\""); } }; diff --git a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/factory/GatewayFilterFactory.java b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/factory/GatewayFilterFactory.java index a49f0c51..160caf8b 100644 --- a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/factory/GatewayFilterFactory.java +++ b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/factory/GatewayFilterFactory.java @@ -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); + } } diff --git a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/factory/HystrixGatewayFilterFactory.java b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/factory/HystrixGatewayFilterFactory.java index 0b60fb88..03d2a929 100644 --- a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/factory/HystrixGatewayFilterFactory.java +++ b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/factory/HystrixGatewayFilterFactory.java @@ -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)); } diff --git a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/factory/PrefixPathGatewayFilterFactory.java b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/factory/PrefixPathGatewayFilterFactory.java index 3045677f..3817219b 100644 --- a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/factory/PrefixPathGatewayFilterFactory.java +++ b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/factory/PrefixPathGatewayFilterFactory.java @@ -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(); diff --git a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/factory/RewritePathGatewayFilterFactory.java b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/factory/RewritePathGatewayFilterFactory.java index fa36c96c..3fd797ed 100644 --- a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/factory/RewritePathGatewayFilterFactory.java +++ b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/factory/RewritePathGatewayFilterFactory.java @@ -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(); diff --git a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/factory/SetPathGatewayFilterFactory.java b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/factory/SetPathGatewayFilterFactory.java index 0d5a8adb..2477c704 100644 --- a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/factory/SetPathGatewayFilterFactory.java +++ b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/factory/SetPathGatewayFilterFactory.java @@ -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(); diff --git a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/support/GatewayServerHttpRequestBuilder.java b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/support/GatewayServerHttpRequestBuilder.java new file mode 100644 index 00000000..9c85fdb7 --- /dev/null +++ b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/support/GatewayServerHttpRequestBuilder.java @@ -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 cookies; + + @Nullable + private String uriPath; + + @Nullable + private String contextPath; + + private Flux 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 void copyMultiValueMap(MultiValueMap source, + MultiValueMap destination) { + + for (Map.Entry> entry : source.entrySet()) { + K key = entry.getKey(); + List 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 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 cookies; + + @Nullable + private final InetSocketAddress remoteAddress; + + @Nullable + private final SslInfo sslInfo; + + private final Flux body; + + private final ServerHttpRequest originalRequest; + + + public GatewayServerHttpRequest(URI uri, @Nullable String contextPath, + HttpHeaders headers, String methodValue, MultiValueMap cookies, + Flux 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 initCookies() { + return this.cookies; + } + + @Nullable + @Override + public InetSocketAddress getRemoteAddress() { + return this.remoteAddress; + } + + @Nullable + @Override + protected SslInfo initSslInfo() { + return this.sslInfo; + } + + @Override + public Flux getBody() { + return this.body; + } + + @SuppressWarnings("unchecked") + @Override + public T getNativeRequest() { + return (T) this.originalRequest; + } + } + +} diff --git a/spring-cloud-gateway-core/src/test/java/org/springframework/cloud/gateway/filter/factory/AddRequestParameterGatewayFilterFactoryTests.java b/spring-cloud-gateway-core/src/test/java/org/springframework/cloud/gateway/filter/factory/AddRequestParameterGatewayFilterFactoryTests.java index 16fdbca3..d8e1b972 100644 --- a/spring-cloud-gateway-core/src/test/java/org/springframework/cloud/gateway/filter/factory/AddRequestParameterGatewayFilterFactoryTests.java +++ b/spring-cloud-gateway-core/src/test/java/org/springframework/cloud/gateway/filter/factory/AddRequestParameterGatewayFilterFactoryTests.java @@ -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 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 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); diff --git a/spring-cloud-gateway-core/src/test/java/org/springframework/cloud/gateway/filter/factory/RewritePathGatewayFilterFactoryTests.java b/spring-cloud-gateway-core/src/test/java/org/springframework/cloud/gateway/filter/factory/RewritePathGatewayFilterFactoryTests.java index d5fc7cfc..d0c66fce 100644 --- a/spring-cloud-gateway-core/src/test/java/org/springframework/cloud/gateway/filter/factory/RewritePathGatewayFilterFactoryTests.java +++ b/spring-cloud-gateway-core/src/test/java/org/springframework/cloud/gateway/filter/factory/RewritePathGatewayFilterFactoryTests.java @@ -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/(?\\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 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"); } }