diff --git a/spring-cloud-gateway-server-mvc/src/main/java/org/springframework/cloud/gateway/server/mvc/common/MvcUtils.java b/spring-cloud-gateway-server-mvc/src/main/java/org/springframework/cloud/gateway/server/mvc/common/MvcUtils.java index 5eec2dd6..05c073ed 100644 --- a/spring-cloud-gateway-server-mvc/src/main/java/org/springframework/cloud/gateway/server/mvc/common/MvcUtils.java +++ b/spring-cloud-gateway-server-mvc/src/main/java/org/springframework/cloud/gateway/server/mvc/common/MvcUtils.java @@ -256,13 +256,11 @@ public abstract class MvcUtils { request.servletRequest().setAttribute(GATEWAY_REQUEST_URL_ATTR, url); } + @SuppressWarnings("unchecked") public static void addOriginalRequestUrl(ServerRequest request, URI url) { - LinkedHashSet urls = getAttribute(request, GATEWAY_ORIGINAL_REQUEST_URL_ATTR); - if (urls == null) { - urls = new LinkedHashSet<>(); - } + LinkedHashSet urls = (LinkedHashSet) request.attributes() + .computeIfAbsent(GATEWAY_ORIGINAL_REQUEST_URL_ATTR, s -> new LinkedHashSet<>()); urls.add(url); - putAttribute(request, GATEWAY_ORIGINAL_REQUEST_URL_ATTR, urls); } private record ByteArrayInputMessage(ServerRequest request, ByteArrayInputStream body) implements HttpInputMessage { diff --git a/spring-cloud-gateway-server-mvc/src/main/java/org/springframework/cloud/gateway/server/mvc/filter/BeforeFilterFunctions.java b/spring-cloud-gateway-server-mvc/src/main/java/org/springframework/cloud/gateway/server/mvc/filter/BeforeFilterFunctions.java index 83904cb5..3d81e508 100644 --- a/spring-cloud-gateway-server-mvc/src/main/java/org/springframework/cloud/gateway/server/mvc/filter/BeforeFilterFunctions.java +++ b/spring-cloud-gateway-server-mvc/src/main/java/org/springframework/cloud/gateway/server/mvc/filter/BeforeFilterFunctions.java @@ -189,12 +189,14 @@ public abstract class BeforeFilterFunctions { final UriTemplate uriTemplate = new UriTemplate(prefix); return request -> { + MvcUtils.addOriginalRequestUrl(request, request.uri()); Map uriVariables = MvcUtils.getUriTemplateVariables(request); URI uri = uriTemplate.expand(uriVariables); String newPath = uri.getRawPath() + request.uri().getRawPath(); URI prefixedUri = UriComponentsBuilder.fromUri(request.uri()).replacePath(newPath).build().toUri(); + MvcUtils.setRequestUrl(request, prefixedUri); return ServerRequest.from(request).uri(prefixedUri).build(); }; } @@ -326,7 +328,7 @@ public abstract class BeforeFilterFunctions { String normalizedReplacement = replacement.replace("$\\", "$"); Pattern pattern = Pattern.compile(regexp); return request -> { - // TODO: original request url + MvcUtils.addOriginalRequestUrl(request, request.uri()); String path = request.uri().getRawPath(); String newPath = pattern.matcher(path).replaceAll(normalizedReplacement); @@ -334,8 +336,7 @@ public abstract class BeforeFilterFunctions { ServerRequest modified = ServerRequest.from(request).uri(rewrittenUri).build(); - // TODO: can this be restored at some point? - // MvcUtils.setRequestUrl(modified, modified.uri()); + MvcUtils.setRequestUrl(request, rewrittenUri); return modified; }; } @@ -372,14 +373,13 @@ public abstract class BeforeFilterFunctions { UriTemplate uriTemplate = new UriTemplate(path); return request -> { + MvcUtils.addOriginalRequestUrl(request, request.uri()); Map uriVariables = MvcUtils.getUriTemplateVariables(request); URI uri = uriTemplate.expand(uriVariables); - URI prefixedUri = UriComponentsBuilder.fromUri(request.uri()) - .replacePath(uri.getRawPath()) - .build(true) - .toUri(); - return ServerRequest.from(request).uri(prefixedUri).build(); + URI newUri = UriComponentsBuilder.fromUri(request.uri()).replacePath(uri.getRawPath()).build(true).toUri(); + MvcUtils.setRequestUrl(request, newUri); + return ServerRequest.from(request).uri(newUri).build(); }; } diff --git a/spring-cloud-gateway-server-mvc/src/main/java/org/springframework/cloud/gateway/server/mvc/filter/LoadBalancerFilterFunctions.java b/spring-cloud-gateway-server-mvc/src/main/java/org/springframework/cloud/gateway/server/mvc/filter/LoadBalancerFilterFunctions.java index ab6d47dc..db7b1464 100644 --- a/spring-cloud-gateway-server-mvc/src/main/java/org/springframework/cloud/gateway/server/mvc/filter/LoadBalancerFilterFunctions.java +++ b/spring-cloud-gateway-server-mvc/src/main/java/org/springframework/cloud/gateway/server/mvc/filter/LoadBalancerFilterFunctions.java @@ -63,6 +63,8 @@ public abstract class LoadBalancerFilterFunctions { public static HandlerFilterFunction lb(String serviceId, BiFunction reconstructUriFunction) { return (request, next) -> { + MvcUtils.addOriginalRequestUrl(request, request.uri()); + LoadBalancerClientFactory clientFactory = getApplicationContext(request) .getBean(LoadBalancerClientFactory.class); Set supportedLifecycleProcessors = LoadBalancerLifecycleValidator diff --git a/spring-cloud-gateway-server-mvc/src/test/java/org/springframework/cloud/gateway/server/mvc/ServerMvcIntegrationTests.java b/spring-cloud-gateway-server-mvc/src/test/java/org/springframework/cloud/gateway/server/mvc/ServerMvcIntegrationTests.java index 02a89446..38bfa681 100644 --- a/spring-cloud-gateway-server-mvc/src/test/java/org/springframework/cloud/gateway/server/mvc/ServerMvcIntegrationTests.java +++ b/spring-cloud-gateway-server-mvc/src/test/java/org/springframework/cloud/gateway/server/mvc/ServerMvcIntegrationTests.java @@ -239,6 +239,7 @@ public class ServerMvcIntegrationTests { public void stripPrefixWorks() { restClient.get() .uri("/long/path/to/get") + .header("Host", "www.stripprefix.org") .exchange() .expectStatus() .isOk() @@ -246,14 +247,13 @@ public class ServerMvcIntegrationTests { .consumeWith(res -> { Map map = res.getResponseBody(); Map headers = getMap(map, "headers"); - assertThat(headers).containsKeys( - XForwardedRequestHeadersFilter.X_FORWARDED_PREFIX_HEADER, + assertThat(headers).containsKeys(XForwardedRequestHeadersFilter.X_FORWARDED_PREFIX_HEADER, XForwardedRequestHeadersFilter.X_FORWARDED_HOST_HEADER, XForwardedRequestHeadersFilter.X_FORWARDED_PORT_HEADER, XForwardedRequestHeadersFilter.X_FORWARDED_PROTO_HEADER, XForwardedRequestHeadersFilter.X_FORWARDED_FOR_HEADER); - assertThat(headers).containsEntry( - XForwardedRequestHeadersFilter.X_FORWARDED_PREFIX_HEADER, "/long/path/to"); + assertThat(headers).containsEntry(XForwardedRequestHeadersFilter.X_FORWARDED_PREFIX_HEADER, + "/long/path/to"); assertThat(headers).containsEntry("X-Test", "stripPrefix"); }); } @@ -272,18 +272,40 @@ public class ServerMvcIntegrationTests { Map map = res.getResponseBody(); assertThat(map).containsEntry("data", "hello"); Map headers = getMap(map, "headers"); - assertThat(headers).containsKeys( - XForwardedRequestHeadersFilter.X_FORWARDED_PREFIX_HEADER, + assertThat(headers).containsKeys(XForwardedRequestHeadersFilter.X_FORWARDED_PREFIX_HEADER, XForwardedRequestHeadersFilter.X_FORWARDED_HOST_HEADER, XForwardedRequestHeadersFilter.X_FORWARDED_PORT_HEADER, XForwardedRequestHeadersFilter.X_FORWARDED_PROTO_HEADER, XForwardedRequestHeadersFilter.X_FORWARDED_FOR_HEADER); - assertThat(headers).containsEntry( - XForwardedRequestHeadersFilter.X_FORWARDED_PREFIX_HEADER, "/long/path/to"); + assertThat(headers).containsEntry(XForwardedRequestHeadersFilter.X_FORWARDED_PREFIX_HEADER, + "/long/path/to"); assertThat(headers).containsEntry("X-Test", "stripPrefixPost"); }); } + @Test + public void stripPrefixLbWorks() { + restClient.get() + .uri("/long/path/to/get") + .header("Host", "www.stripprefixlb.org") + .exchange() + .expectStatus() + .isOk() + .expectBody(Map.class) + .consumeWith(res -> { + Map map = res.getResponseBody(); + Map headers = getMap(map, "headers"); + assertThat(headers).containsKeys(XForwardedRequestHeadersFilter.X_FORWARDED_PREFIX_HEADER, + XForwardedRequestHeadersFilter.X_FORWARDED_HOST_HEADER, + XForwardedRequestHeadersFilter.X_FORWARDED_PORT_HEADER, + XForwardedRequestHeadersFilter.X_FORWARDED_PROTO_HEADER, + XForwardedRequestHeadersFilter.X_FORWARDED_FOR_HEADER); + assertThat(headers).containsEntry(XForwardedRequestHeadersFilter.X_FORWARDED_PREFIX_HEADER, + "/long/path/to"); + assertThat(headers).containsEntry("X-Test", "stripPrefix"); + }); + } + @Test public void setStatusGatewayRouterFunctionWorks() { restClient.get() @@ -1083,8 +1105,8 @@ public class ServerMvcIntegrationTests { // @formatter:off return route("testsetpath") .route(POST("/mycustompath{extra}").and(host("**.setpathpost.org")), http()) - .filter(new HttpbinUriResolver()) .filter(setPath("/{extra}")) + .filter(new HttpbinUriResolver()) .build(); // @formatter:on } @@ -1092,11 +1114,12 @@ public class ServerMvcIntegrationTests { @Bean public RouterFunction gatewayRouterFunctionsStripPrefix() { // @formatter:off - return route(GET("/long/path/to/get"), http()) + return route("teststripprefix") + .route(GET("/long/path/to/get").and(host("**.stripprefix.org")), http()) .filter(stripPrefix(3)) .filter(addRequestHeader("X-Test", "stripPrefix")) - .filter(new HttpbinUriResolver(true)) - .withAttribute(MvcUtils.GATEWAY_ROUTE_ID_ATTR, "teststripprefix"); + .filter(new HttpbinUriResolver()) + .build(); // @formatter:on } @@ -1107,7 +1130,19 @@ public class ServerMvcIntegrationTests { .route(POST("/long/path/to/post").and(host("**.stripprefixpost.org")), http()) .filter(stripPrefix(3)) .filter(addRequestHeader("X-Test", "stripPrefixPost")) - .filter(new HttpbinUriResolver(true)) + .filter(new HttpbinUriResolver()) + .build(); + // @formatter:on + } + + @Bean + public RouterFunction gatewayRouterFunctionsStripPrefixLb() { + // @formatter:off + return route("teststripprefix") + .route(GET("/long/path/to/get").and(host("**.stripprefixlb.org")), http()) + .filter(stripPrefix(3)) + .filter(addRequestHeader("X-Test", "stripPrefix")) + .filter(lb("httpbin")) .build(); // @formatter:on } @@ -1442,8 +1477,8 @@ public class ServerMvcIntegrationTests { return route("requestheadertorequesturi") .route(cloudFoundryRouteService().and(host("**.requestheadertorequesturi.org")), http()) //.before(new HttpbinUriResolver()) NO URI RESOLVER! - .before(requestHeaderToRequestUri("X-CF-Forwarded-Url")) .filter(setPath("/hello")) + .before(requestHeaderToRequestUri("X-CF-Forwarded-Url")) .build(); // @formatter:on } diff --git a/spring-cloud-gateway-server-mvc/src/test/java/org/springframework/cloud/gateway/server/mvc/test/HttpbinUriResolver.java b/spring-cloud-gateway-server-mvc/src/test/java/org/springframework/cloud/gateway/server/mvc/test/HttpbinUriResolver.java index 0d327185..93b6bb51 100644 --- a/spring-cloud-gateway-server-mvc/src/test/java/org/springframework/cloud/gateway/server/mvc/test/HttpbinUriResolver.java +++ b/spring-cloud-gateway-server-mvc/src/test/java/org/springframework/cloud/gateway/server/mvc/test/HttpbinUriResolver.java @@ -16,6 +16,7 @@ package org.springframework.cloud.gateway.server.mvc.test; +import java.lang.reflect.UndeclaredThrowableException; import java.net.URI; import java.net.URISyntaxException; import java.util.function.Function; @@ -31,33 +32,21 @@ import org.springframework.web.servlet.function.ServerResponse; public class HttpbinUriResolver implements Function, HandlerFilterFunction { - private final boolean preservePath; - - public HttpbinUriResolver(boolean preservePath) { - this.preservePath = preservePath; - } - - public HttpbinUriResolver() { - this(false); - } - protected URI uri(ServerRequest request) { ApplicationContext context = MvcUtils.getApplicationContext(request); Integer port = context.getEnvironment().getProperty("httpbin.port", Integer.class); String host = context.getEnvironment().getProperty("httpbin.host"); Assert.hasText(host, "httpbin.host is not set, did you initialize HttpbinTestcontainers?"); Assert.notNull(port, "httpbin.port is not set, did you initialize HttpbinTestcontainers?"); - if (preservePath) { - URI original = request.uri(); - try { - return new URI("http", original.getUserInfo(), host, port, original.getPath(), - original.getQuery(), original.getFragment()); - } catch (URISyntaxException e) { - throw new IllegalArgumentException(e.getMessage(), e); - } + URI original = request.uri(); + try { + return new URI("http", original.getUserInfo(), host, port, original.getPath(), original.getQuery(), + original.getFragment()); + } + catch (URISyntaxException e) { + throw new UndeclaredThrowableException(e); } - return URI.create(String.format("http://%s:%d", host, port)); } @Override