POST 请求支持 AccessKey 放在 URI请求参数中

This commit is contained in:
Wenchao Gong 2021-10-20 23:52:17 +08:00
parent 1317da09ca
commit 685fab4eed
3 changed files with 161 additions and 69 deletions

View File

@ -9,6 +9,7 @@ import org.springframework.core.annotation.Order;
import org.springframework.core.io.buffer.DataBufferFactory; import org.springframework.core.io.buffer.DataBufferFactory;
import org.springframework.http.HttpStatus; import org.springframework.http.HttpStatus;
import org.springframework.http.MediaType; import org.springframework.http.MediaType;
import org.springframework.http.server.RequestPath;
import org.springframework.http.server.reactive.ServerHttpResponse; import org.springframework.http.server.reactive.ServerHttpResponse;
import org.springframework.web.server.ResponseStatusException; import org.springframework.web.server.ResponseStatusException;
import org.springframework.web.server.ServerWebExchange; import org.springframework.web.server.ServerWebExchange;
@ -29,6 +30,10 @@ import java.util.Optional;
@Order(-1) @Order(-1)
public class GlobalErrorWebExceptionHandler implements ErrorWebExceptionHandler { public class GlobalErrorWebExceptionHandler implements ErrorWebExceptionHandler {
public static final String NOT_FOUND_MESSAGE = "服务不存在";
public static final String SERVICE_UNAVAILABLE_MESSAGE = "服务暂时不可用";
public static final String GATEWAY_TIMEOUT_MESSAGE = "服务响应超时";
@Override @Override
public Mono<Void> handle(ServerWebExchange exchange, Throwable ex) public Mono<Void> handle(ServerWebExchange exchange, Throwable ex)
{ {
@ -49,17 +54,31 @@ public class GlobalErrorWebExceptionHandler implements ErrorWebExceptionHandler
return response.writeWith(Mono.fromSupplier(() -> { return response.writeWith(Mono.fromSupplier(() -> {
DataBufferFactory bufferFactory = response.bufferFactory(); DataBufferFactory bufferFactory = response.bufferFactory();
return bufferFactory.wrap(JSONUtil.toJsonStr(result(code, message)) return bufferFactory.wrap(JSONUtil.toJsonStr(result(code, message, exchange.getRequest().getPath()))
.getBytes(StandardCharsets.UTF_8)); .getBytes(StandardCharsets.UTF_8));
})); }));
} }
private String getMessage(Throwable ex) private String getMessage(Throwable ex)
{ {
String reason;
if(ex instanceof ResponseStatusException) { if(ex instanceof ResponseStatusException) {
String reason = ((ResponseStatusException) ex).getReason(); ResponseStatusException statusException = (ResponseStatusException) ex;
switch (statusException.getStatus()) {
case NOT_FOUND:
reason = NOT_FOUND_MESSAGE;
break;
case SERVICE_UNAVAILABLE:
reason = SERVICE_UNAVAILABLE_MESSAGE;
break;
case GATEWAY_TIMEOUT:
reason = GATEWAY_TIMEOUT_MESSAGE;
break;
default:
reason = statusException.getReason();
}
if(log.isDebugEnabled()){ if(log.isDebugEnabled()){
log.debug(ex.getMessage()); log.debug(reason, ex);
} }
if(StrUtil.isNotBlank(reason)) { if(StrUtil.isNotBlank(reason)) {
return reason; return reason;
@ -70,10 +89,11 @@ public class GlobalErrorWebExceptionHandler implements ErrorWebExceptionHandler
return message; return message;
} }
private Map<String, Object> result(Optional<Integer> code, String message) { private Map<String, Object> result(Optional<Integer> code, String message, RequestPath path) {
final HashMap<String, Object> result = MapUtil.newHashMap(3); final HashMap<String, Object> result = MapUtil.newHashMap(3);
result.put("code", code.orElse(HttpStatus.INTERNAL_SERVER_ERROR.value())); result.put("code", code.orElse(HttpStatus.INTERNAL_SERVER_ERROR.value()));
result.put("message", message); result.put("message", message);
result.put("path", path.value());
result.put("data", null); result.put("data", null);
return result; return result;
} }

View File

@ -3,6 +3,7 @@ package com.ruoyi.gateway.filter;
import cn.hutool.core.lang.Assert; import cn.hutool.core.lang.Assert;
import cn.hutool.core.util.StrUtil; import cn.hutool.core.util.StrUtil;
import cn.hutool.json.JSONArray; import cn.hutool.json.JSONArray;
import cn.hutool.json.JSONException;
import cn.hutool.json.JSONObject; import cn.hutool.json.JSONObject;
import cn.hutool.json.JSONUtil; import cn.hutool.json.JSONUtil;
import com.ruoyi.gateway.exception.*; import com.ruoyi.gateway.exception.*;
@ -30,9 +31,12 @@ import java.net.URI;
import java.util.*; import java.util.*;
import java.util.concurrent.TimeUnit; import java.util.concurrent.TimeUnit;
import java.util.function.BiConsumer; import java.util.function.BiConsumer;
import java.util.function.BiFunction;
import java.util.function.Supplier; import java.util.function.Supplier;
import java.util.stream.Collectors; import java.util.stream.Collectors;
import static com.ruoyi.gateway.filter.CustomerGlobalFilter.AccessKey.AccessKeyType;
import static com.ruoyi.gateway.filter.CustomerGlobalFilter.AccessKey.AccessKeyType.*;
import static org.springframework.util.CollectionUtils.unmodifiableMultiValueMap; import static org.springframework.util.CollectionUtils.unmodifiableMultiValueMap;
/** /**
@ -53,6 +57,7 @@ public class CustomerGlobalFilter implements GlobalFilter, Ordered {
this.rateLimiter = rateLimiter; this.rateLimiter = rateLimiter;
} }
private static final String ACCESS_KEY_NAME_DEFAULT = "ak";
private static final TimeUnit[] TIME_UNITS = {TimeUnit.SECONDS, TimeUnit.MINUTES, TimeUnit.HOURS, TimeUnit.DAYS}; private static final TimeUnit[] TIME_UNITS = {TimeUnit.SECONDS, TimeUnit.MINUTES, TimeUnit.HOURS, TimeUnit.DAYS};
@Override @Override
@ -62,12 +67,15 @@ public class CustomerGlobalFilter implements GlobalFilter, Ordered {
final ServerHttpRequest request = exchange.getRequest(); final ServerHttpRequest request = exchange.getRequest();
//ak 是否存在 //ak 是否存在
final HttpMethod httpMethod = request.getMethod(); final HttpMethod httpMethod = request.getMethod();
final String accessKeyName = String.valueOf(metadata.get(GatewayUtils.CONFIG_ACCESS_KEY_NAME_KEY)); AccessKey accessKey = new AccessKey(String.valueOf(metadata.getOrDefault(GatewayUtils.CONFIG_ACCESS_KEY_NAME_KEY,
String headerAk = GatewayUtils.getValue(null, () -> request.getHeaders().get(accessKeyName)); ACCESS_KEY_NAME_DEFAULT)));
accessKey.set(GatewayUtils.getValue(null, () -> request.getHeaders().get(accessKey.name)), HEADER);
MultiValueMap<String, String> queryParams = request.getQueryParams().containsKey(accessKey.name) ?
new LinkedMultiValueMap<>(request.getQueryParams()) : request.getQueryParams();
if (HttpMethod.GET.equals(httpMethod)) { if (HttpMethod.GET.equals(httpMethod)) {
return handleGetRequest(exchange, chain, route, request, accessKeyName, headerAk); return handleGetRequest(exchange, chain, route, request, accessKey, queryParams);
} else if (HttpMethod.POST.equals(httpMethod)) { } else if (HttpMethod.POST.equals(httpMethod)) {
return handlePostRequest(exchange, chain, route, request, accessKeyName, headerAk); return handlePostRequest(exchange, chain, route, request, accessKey, queryParams);
} }
throw new MethodNotSupportedException(httpMethod); throw new MethodNotSupportedException(httpMethod);
} }
@ -79,25 +87,36 @@ public class CustomerGlobalFilter implements GlobalFilter, Ordered {
* @param chain 当前过滤链 * @param chain 当前过滤链
* @param route 路由信息 * @param route 路由信息
* @param request 请求信息 * @param request 请求信息
* @param accessKeyName AK键名 * @param accessKey AK信息
* @param headerAk 头部AK信息 * @param queryParams URL查询参数
* @return 指示请求处理何时完成 * @return 指示请求处理何时完成
*/ */
private Mono<Void> handlePostRequest(ServerWebExchange exchange, GatewayFilterChain chain, Route route, private Mono<Void> handlePostRequest(ServerWebExchange exchange, GatewayFilterChain chain, Route route,
ServerHttpRequest request, String accessKeyName, String headerAk) { ServerHttpRequest request, AccessKey accessKey,
final ServerRequest serverRequest = ServerRequest.create(exchange, MultiValueMap<String, String> queryParams) {
HandlerStrategies.withDefaults().messageReaders()); //处理 URL AccessKey
handleAccessKey(accessKey, () -> queryParams.get(accessKey.name), URI, false);
ServerRequest serverRequest = ServerRequest.create(exchange, HandlerStrategies.withDefaults().messageReaders());
final Mono<String> modifiedBody = serverRequest.bodyToMono(String.class).defaultIfEmpty(StrUtil.EMPTY) final Mono<String> modifiedBody = serverRequest.bodyToMono(String.class).defaultIfEmpty(StrUtil.EMPTY)
.flatMap(body -> { .flatMap(body -> {
MediaType mediaType = request.getHeaders().getContentType(); MediaType mediaType = request.getHeaders().getContentType();
if (MediaType.APPLICATION_JSON.equals(mediaType)) { if (Objects.isNull(mediaType) || MediaType.APPLICATION_JSON.equals(mediaType)) {
return handlePostRequestJson(exchange, route, request, accessKeyName, headerAk, body); try {
JSONObject jsonObj = StrUtil.isBlank(body) ? new JSONObject() : JSONUtil.parseObj(body);
return handlePostRequestJson(exchange, route, request, accessKey, jsonObj, queryParams);
}catch (JSONException e) {
if(Objects.nonNull(mediaType)) {
return Mono.error(e);
}
}
} else if (MediaType.APPLICATION_FORM_URLENCODED.equals(mediaType)) { } else if (MediaType.APPLICATION_FORM_URLENCODED.equals(mediaType)) {
return handlePostRequestFormUrlencoded(exchange, route, request, accessKeyName, headerAk, body); return handlePostRequestFormUrlencoded(exchange, route, request, accessKey, body, queryParams);
} }
return Mono.error(() -> new ContentTypeNotSupportedException(mediaType)); return Mono.error(() -> new ContentTypeNotSupportedException(mediaType));
}); });
return GatewayUtils.modifyBody(exchange, chain, modifiedBody); //如果 AccessKey 在URI部分 URI需要重新生成
URI uri = URI.equals(accessKey.type) ? getNewUri().apply(request, queryParams) : null;
return GatewayUtils.modifyBody(exchange, chain, modifiedBody, uri);
} }
/** /**
@ -106,26 +125,25 @@ public class CustomerGlobalFilter implements GlobalFilter, Ordered {
* @param exchange 当前服务交换器 * @param exchange 当前服务交换器
* @param route 路由信息 * @param route 路由信息
* @param request 请求信息 * @param request 请求信息
* @param accessKeyName AK键名 * @param accessKey AK信息
* @param headerAk 头部AK信息
* @param body 请求Body信息 * @param body 请求Body信息
* @return 处理后的Body信息 * @return 处理后的Body信息
*/ */
private Mono<? extends String> handlePostRequestFormUrlencoded(ServerWebExchange exchange, Route route, private Mono<? extends String> handlePostRequestFormUrlencoded(ServerWebExchange exchange, Route route,
ServerHttpRequest request, String accessKeyName, ServerHttpRequest request, AccessKey accessKey,
String headerAk, String body) { String body, MultiValueMap<String, String> queryParams) {
if (StrUtil.isBlank(body)) { if (StrUtil.isBlank(body)) {
Assert.notBlank(headerAk, () -> new AkRequireException(accessKeyName)); Assert.notBlank(accessKey.value, () -> new AkRequireException(accessKey.name));
} }
final List<String[]> srcParams = Arrays.stream(body.split("&")).map(param -> param.split("=")) final List<String[]> srcParams = Arrays.stream(body.split("&")).map(param -> param.split("="))
.collect(Collectors.toList()); .collect(Collectors.toList());
final IscRule rule = handleRule(headerAk, () -> srcParams.stream() final IscRule rule = handleRule(accessKey, () -> srcParams.stream()
.filter(param -> param.length > 0 && accessKeyName.equals(param[0])) .filter(param -> param.length > 0 && accessKey.name.equals(param[0]))
.map(param -> param[1]).collect(Collectors.toList()), route, accessKeyName); .map(param -> param[1]).collect(Collectors.toList()), route, BODY);
Supplier<Mono<String>> rateLimiterAfterSupplier = () -> { Supplier<Mono<String>> rateLimiterAfterSupplier = () -> {
removeParam(headerAk, request, accessKeyName, null); removeParam(accessKey, request, null, queryParams);
final List<String[]> params = srcParams.stream().filter(param -> !accessKeyName.equals(param[0])) final List<String[]> params = srcParams.stream().filter(param -> !accessKey.name.equals(param[0]))
.collect(Collectors.toList()); .collect(Collectors.toList());
handleHiddenParams(route, params, (next, list) -> { handleHiddenParams(route, params, (next, list) -> {
Object value; Object value;
@ -150,19 +168,17 @@ public class CustomerGlobalFilter implements GlobalFilter, Ordered {
* @param exchange 当前服务交换器 * @param exchange 当前服务交换器
* @param route 路由信息 * @param route 路由信息
* @param request 请求信息 * @param request 请求信息
* @param accessKeyName AK键名 * @param accessKey AK信息
* @param headerAk 头部AK信息 * @param jsonObj 请求Body JSON信息
* @param body 请求Body信息
* @return 处理后的Body信息 * @return 处理后的Body信息
*/ */
private Mono<? extends String> handlePostRequestJson(ServerWebExchange exchange, Route route, private Mono<? extends String> handlePostRequestJson(ServerWebExchange exchange, Route route,
ServerHttpRequest request, String accessKeyName, ServerHttpRequest request, AccessKey accessKey,
String headerAk, String body) { JSONObject jsonObj, MultiValueMap<String, String> queryParams) {
JSONObject jsonObj = StrUtil.isBlank(body) ? new JSONObject() : JSONUtil.parseObj(body); final IscRule rule = handleRule(accessKey, () -> Collections.singletonList(jsonObj.get(accessKey.name,
final IscRule rule = handleRule(headerAk, () -> Collections.singletonList(jsonObj.get(accessKeyName, String.class, true)), route, BODY);
String.class, true)), route, accessKeyName);
Supplier<Mono<String>> rateLimiterAfterSupplier = () -> { Supplier<Mono<String>> rateLimiterAfterSupplier = () -> {
removeParam(headerAk, request, accessKeyName, () -> jsonObj.remove(accessKeyName)); removeParam(accessKey, request, () -> jsonObj.remove(accessKey.name), queryParams);
handleHiddenParams(route, jsonObj, (next, map) -> map.set(next.getKey(), next.getValue())); handleHiddenParams(route, jsonObj, (next, map) -> map.set(next.getKey(), next.getValue()));
return Mono.just(jsonObj.toString()); return Mono.just(jsonObj.toString());
}; };
@ -177,16 +193,16 @@ public class CustomerGlobalFilter implements GlobalFilter, Ordered {
* @param chain 当前过滤链 * @param chain 当前过滤链
* @param route 路由信息 * @param route 路由信息
* @param request 请求信息 * @param request 请求信息
* @param accessKeyName AK键名 * @param accessKey AK信息
* @param headerAk 头部AK信息 * @param queryParams URL查询参数
* @return 指示请求处理何时完成 * @return 指示请求处理何时完成
*/ */
private Mono<Void> handleGetRequest(ServerWebExchange exchange, GatewayFilterChain chain, Route route, private Mono<Void> handleGetRequest(ServerWebExchange exchange, GatewayFilterChain chain, Route route,
ServerHttpRequest request, String accessKeyName, String headerAk) { ServerHttpRequest request, AccessKey accessKey,
final MultiValueMap<String, String> queryParams = new LinkedMultiValueMap<>(request.getQueryParams()); MultiValueMap<String, String> queryParams) {
final IscRule rule = handleRule(headerAk, () -> queryParams.get(accessKeyName), route, accessKeyName); final IscRule rule = handleRule(accessKey, () -> queryParams.get(accessKey.name), route, URI);
Supplier<Mono<Void>> rateLimiterAfterSupplier = () -> { Supplier<Mono<Void>> rateLimiterAfterSupplier = () -> {
removeParam(headerAk, request, accessKeyName, () -> queryParams.remove(accessKeyName)); removeParam(accessKey, request, null, queryParams);
handleHiddenParams(route, queryParams, (next, map) -> { handleHiddenParams(route, queryParams, (next, map) -> {
Object value; Object value;
if (Objects.isNull(value = next.getValue())) { if (Objects.isNull(value = next.getValue())) {
@ -197,14 +213,17 @@ public class CustomerGlobalFilter implements GlobalFilter, Ordered {
map.add(next.getKey(), value.toString()); map.add(next.getKey(), value.toString());
} }
}); });
URI newUri = UriComponentsBuilder.fromUri(request.getURI()) ServerHttpRequest req = exchange.getRequest().mutate().uri(getNewUri().apply(request, queryParams)).build();
.replaceQueryParams(unmodifiableMultiValueMap(queryParams)).build().toUri(); return chain.filter(exchange.mutate().request(req).build());
ServerHttpRequest updatedRequest = exchange.getRequest().mutate().uri(newUri).build();
return chain.filter(exchange.mutate().request(updatedRequest).build());
}; };
return rateLimiter(exchange, rule, route, 0, rateLimiterAfterSupplier); return rateLimiter(exchange, rule, route, 0, rateLimiterAfterSupplier);
} }
private BiFunction<ServerHttpRequest, MultiValueMap<String, String>, URI> getNewUri() {
return (request, queryParams) -> UriComponentsBuilder.fromUri(request.getURI())
.replaceQueryParams(unmodifiableMultiValueMap(queryParams)).build().toUri();
}
@Override @Override
public int getOrder() { public int getOrder() {
return -1; return -1;
@ -213,42 +232,64 @@ public class CustomerGlobalFilter implements GlobalFilter, Ordered {
/** /**
* 处理规则 * 处理规则
* *
* @param headerAk 头部AK * @param accessKey AK信息
* @param valueSupplier valueList 生产者 * @param valueSupplier valueList 生产者
* @param route 路由信息 * @param route 路由信息
* @param accessKeyName AK键名 * @param type AK类型
* @return AK对应服务规则信息 * @return AK对应服务规则信息
*/ */
private IscRule handleRule(String headerAk, Supplier<List<String>> valueSupplier, Route route, String accessKeyName) { private IscRule handleRule(AccessKey accessKey, Supplier<List<String>> valueSupplier, Route route, AccessKeyType type) {
//获取AK if(Objects.nonNull(valueSupplier) && Objects.nonNull(type)) {
final String ak = GatewayUtils.getValue(headerAk, valueSupplier, () -> new AkRequireException(accessKeyName)); handleAccessKey(accessKey, valueSupplier, type, true);
}
//获取规则 //获取规则
final IscRule rule = GatewayUtils.getRequiredValue(() -> GatewayUtils.getRule(ak, route.getId()), final IscRule rule = GatewayUtils.getRequiredValue(() -> GatewayUtils.getRule(accessKey.value, route.getId()),
() -> new RuleNotExistException(ak, route.getId())); () -> new RuleNotExistException(accessKey.value, route.getId()));
//是否到期 //是否到期
GatewayUtils.isBefore(rule, () -> new RuleExpiredException(rule.getExpire())); GatewayUtils.isBefore(rule, () -> new RuleExpiredException(rule.getExpire()));
//设置AK 到ID 为了传参方便 //设置AK 到ID 为了传参方便
rule.setId(ak); rule.setId(accessKey.value);
return rule; return rule;
} }
/**
* 处理 accessKey
*
* @param accessKey AccessKey
* @param valueSupplier valueList 生产者
* @param type AccessKeyType
* @param throwException 是否抛出异常
*/
private void handleAccessKey(AccessKey accessKey, Supplier<List<String>> valueSupplier, AccessKeyType type, boolean throwException) {
//获取AK
String value;
if(throwException) {
value = GatewayUtils.getValue(accessKey.value, valueSupplier, () -> new AkRequireException(accessKey.name));
}else {
value = GatewayUtils.getValue(accessKey.value, valueSupplier, null);
}
accessKey.set(value, type);
}
/** /**
* 删除参数 * 删除参数
* *
* @param headerAk 头部AK * @param accessKey AK信息
* @param request 请求 * @param request 请求
* @param accessKeyName AK名称
* @param removeSupplier 删除参数提供者 * @param removeSupplier 删除参数提供者
* @param queryParams URI请求参数
* @param <T> 类型 * @param <T> 类型
*/ */
private <T> void removeParam(String headerAk, ServerHttpRequest request, String accessKeyName, private <T> void removeParam(AccessKey accessKey, ServerHttpRequest request, Supplier<T> removeSupplier,
Supplier<T> removeSupplier) { MultiValueMap<String, String> queryParams) {
//如果header中有AK,则删除 //如果header中有AK,则删除
if (Objects.nonNull(headerAk)) { if (HEADER.equals(accessKey.type)) {
final HttpHeaders headers = request.getHeaders(); final HttpHeaders headers = request.getHeaders();
if(headers.containsKey(accessKeyName)) { if(headers.containsKey(accessKey.name)) {
request.mutate().headers(headMap -> headMap.remove(accessKeyName)).build(); request.mutate().headers(headMap -> headMap.remove(accessKey.name)).build();
} }
}else if(URI.equals(accessKey.type)) {
queryParams.remove(accessKey.name);
} }
if (Objects.nonNull(removeSupplier)) { if (Objects.nonNull(removeSupplier)) {
removeSupplier.get(); removeSupplier.get();
@ -311,4 +352,27 @@ public class CustomerGlobalFilter implements GlobalFilter, Ordered {
return Mono.error(() -> new RateLimitException(limit, timeUnit)); return Mono.error(() -> new RateLimitException(limit, timeUnit));
}); });
} }
/**
* AK 信息
*/
public static class AccessKey {
public AccessKey(String name) {
this.name = name;
}
public void set(String value, AccessKeyType type) {
if(StrUtil.isNotBlank(value)) {
this.value = value;
this.type = type;
}
}
public enum AccessKeyType {
HEADER,
URI,
BODY
}
private final String name;
private String value;
private AccessKeyType type;
}
} }

View File

@ -22,6 +22,7 @@ import org.springframework.web.server.ServerWebExchange;
import reactor.core.publisher.Flux; import reactor.core.publisher.Flux;
import reactor.core.publisher.Mono; import reactor.core.publisher.Mono;
import java.net.URI;
import java.sql.Date; import java.sql.Date;
import java.time.Instant; import java.time.Instant;
import java.util.List; import java.util.List;
@ -50,9 +51,10 @@ public class GatewayUtils {
* @param exchange 当前服务交换器 * @param exchange 当前服务交换器
* @param chain 当前过滤链 * @param chain 当前过滤链
* @param publisher body体提供者 * @param publisher body体提供者
* @param newUri 新的URI去掉AccessKey
* @return 指示请求处理何时完成 * @return 指示请求处理何时完成
*/ */
public static Mono<Void> modifyBody(ServerWebExchange exchange, GatewayFilterChain chain, Mono<String> publisher) { public static Mono<Void> modifyBody(ServerWebExchange exchange, GatewayFilterChain chain, Mono<String> publisher, URI newUri) {
BodyInserter bodyInserter = BodyInserters.fromPublisher(publisher, String.class); BodyInserter bodyInserter = BodyInserters.fromPublisher(publisher, String.class);
HttpHeaders headers = new HttpHeaders(); HttpHeaders headers = new HttpHeaders();
headers.putAll(exchange.getRequest().getHeaders()); headers.putAll(exchange.getRequest().getHeaders());
@ -60,14 +62,14 @@ public class GatewayUtils {
CachedBodyOutputMessage outputMessage = new CachedBodyOutputMessage(exchange, headers); CachedBodyOutputMessage outputMessage = new CachedBodyOutputMessage(exchange, headers);
return bodyInserter.insert(outputMessage, new BodyInserterContext()) return bodyInserter.insert(outputMessage, new BodyInserterContext())
.then(Mono.defer(() -> { .then(Mono.defer(() -> {
ServerHttpRequest request = decorate(exchange, headers, outputMessage); ServerHttpRequest request = decorate(exchange, headers, outputMessage, newUri);
return chain.filter(exchange.mutate().request(request).build()); return chain.filter(exchange.mutate().request(request).build());
})); }));
} }
public static ServerHttpRequestDecorator decorate(ServerWebExchange exchange, HttpHeaders headers, public static ServerHttpRequestDecorator decorate(ServerWebExchange exchange, HttpHeaders headers,
CachedBodyOutputMessage outputMessage) { CachedBodyOutputMessage outputMessage, URI newUri) {
return new ServerHttpRequestDecorator(exchange.getRequest()) { return new ServerHttpRequestDecorator(exchange.getRequest()) {
@Override @Override
public HttpHeaders getHeaders() { public HttpHeaders getHeaders() {
@ -86,6 +88,12 @@ public class GatewayUtils {
public Flux<DataBuffer> getBody() { public Flux<DataBuffer> getBody() {
return outputMessage.getBody(); return outputMessage.getBody();
} }
@Override
public URI getURI()
{
return Objects.isNull(newUri) ? super.getURI() : newUri;
}
}; };
} }
@ -123,7 +131,7 @@ public class GatewayUtils {
} }
} }
} }
if (errorMsgSupplier != null) { if (Objects.nonNull(errorMsgSupplier)) {
Assert.notBlank(before, errorMsgSupplier); Assert.notBlank(before, errorMsgSupplier);
} }
return before; return before;