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.http.HttpStatus;
import org.springframework.http.MediaType;
import org.springframework.http.server.RequestPath;
import org.springframework.http.server.reactive.ServerHttpResponse;
import org.springframework.web.server.ResponseStatusException;
import org.springframework.web.server.ServerWebExchange;
@ -29,6 +30,10 @@ import java.util.Optional;
@Order(-1)
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
public Mono<Void> handle(ServerWebExchange exchange, Throwable ex)
{
@ -49,17 +54,31 @@ public class GlobalErrorWebExceptionHandler implements ErrorWebExceptionHandler
return response.writeWith(Mono.fromSupplier(() -> {
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));
}));
}
private String getMessage(Throwable ex)
{
String reason;
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()){
log.debug(ex.getMessage());
log.debug(reason, ex);
}
if(StrUtil.isNotBlank(reason)) {
return reason;
@ -70,10 +89,11 @@ public class GlobalErrorWebExceptionHandler implements ErrorWebExceptionHandler
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);
result.put("code", code.orElse(HttpStatus.INTERNAL_SERVER_ERROR.value()));
result.put("message", message);
result.put("path", path.value());
result.put("data", null);
return result;
}

View File

@ -3,6 +3,7 @@ package com.ruoyi.gateway.filter;
import cn.hutool.core.lang.Assert;
import cn.hutool.core.util.StrUtil;
import cn.hutool.json.JSONArray;
import cn.hutool.json.JSONException;
import cn.hutool.json.JSONObject;
import cn.hutool.json.JSONUtil;
import com.ruoyi.gateway.exception.*;
@ -30,9 +31,12 @@ import java.net.URI;
import java.util.*;
import java.util.concurrent.TimeUnit;
import java.util.function.BiConsumer;
import java.util.function.BiFunction;
import java.util.function.Supplier;
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;
/**
@ -53,6 +57,7 @@ public class CustomerGlobalFilter implements GlobalFilter, Ordered {
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};
@Override
@ -62,12 +67,15 @@ public class CustomerGlobalFilter implements GlobalFilter, Ordered {
final ServerHttpRequest request = exchange.getRequest();
//ak 是否存在
final HttpMethod httpMethod = request.getMethod();
final String accessKeyName = String.valueOf(metadata.get(GatewayUtils.CONFIG_ACCESS_KEY_NAME_KEY));
String headerAk = GatewayUtils.getValue(null, () -> request.getHeaders().get(accessKeyName));
AccessKey accessKey = new AccessKey(String.valueOf(metadata.getOrDefault(GatewayUtils.CONFIG_ACCESS_KEY_NAME_KEY,
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)) {
return handleGetRequest(exchange, chain, route, request, accessKeyName, headerAk);
return handleGetRequest(exchange, chain, route, request, accessKey, queryParams);
} 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);
}
@ -75,29 +83,40 @@ public class CustomerGlobalFilter implements GlobalFilter, Ordered {
/**
* 处理POST 请求
*
* @param exchange 当前服务交换器
* @param chain 当前过滤链
* @param route 路由信息
* @param request 请求信息
* @param accessKeyName AK键名
* @param headerAk 头部AK信息
* @param exchange 当前服务交换器
* @param chain 当前过滤链
* @param route 路由信息
* @param request 请求信息
* @param accessKey AK信息
* @param queryParams URL查询参数
* @return 指示请求处理何时完成
*/
private Mono<Void> handlePostRequest(ServerWebExchange exchange, GatewayFilterChain chain, Route route,
ServerHttpRequest request, String accessKeyName, String headerAk) {
final ServerRequest serverRequest = ServerRequest.create(exchange,
HandlerStrategies.withDefaults().messageReaders());
ServerHttpRequest request, AccessKey accessKey,
MultiValueMap<String, String> queryParams) {
//处理 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)
.flatMap(body -> {
MediaType mediaType = request.getHeaders().getContentType();
if (MediaType.APPLICATION_JSON.equals(mediaType)) {
return handlePostRequestJson(exchange, route, request, accessKeyName, headerAk, body);
if (Objects.isNull(mediaType) || MediaType.APPLICATION_JSON.equals(mediaType)) {
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)) {
return handlePostRequestFormUrlencoded(exchange, route, request, accessKeyName, headerAk, body);
return handlePostRequestFormUrlencoded(exchange, route, request, accessKey, body, queryParams);
}
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 route 路由信息
* @param request 请求信息
* @param accessKeyName AK键名
* @param headerAk 头部AK信息
* @param accessKey AK信息
* @param body 请求Body信息
* @return 处理后的Body信息
*/
private Mono<? extends String> handlePostRequestFormUrlencoded(ServerWebExchange exchange, Route route,
ServerHttpRequest request, String accessKeyName,
String headerAk, String body) {
ServerHttpRequest request, AccessKey accessKey,
String body, MultiValueMap<String, String> queryParams) {
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("="))
.collect(Collectors.toList());
final IscRule rule = handleRule(headerAk, () -> srcParams.stream()
.filter(param -> param.length > 0 && accessKeyName.equals(param[0]))
.map(param -> param[1]).collect(Collectors.toList()), route, accessKeyName);
final IscRule rule = handleRule(accessKey, () -> srcParams.stream()
.filter(param -> param.length > 0 && accessKey.name.equals(param[0]))
.map(param -> param[1]).collect(Collectors.toList()), route, BODY);
Supplier<Mono<String>> rateLimiterAfterSupplier = () -> {
removeParam(headerAk, request, accessKeyName, null);
final List<String[]> params = srcParams.stream().filter(param -> !accessKeyName.equals(param[0]))
removeParam(accessKey, request, null, queryParams);
final List<String[]> params = srcParams.stream().filter(param -> !accessKey.name.equals(param[0]))
.collect(Collectors.toList());
handleHiddenParams(route, params, (next, list) -> {
Object value;
@ -150,19 +168,17 @@ public class CustomerGlobalFilter implements GlobalFilter, Ordered {
* @param exchange 当前服务交换器
* @param route 路由信息
* @param request 请求信息
* @param accessKeyName AK键名
* @param headerAk 头部AK信息
* @param body 请求Body信息
* @param accessKey AK信息
* @param jsonObj 请求Body JSON信息
* @return 处理后的Body信息
*/
private Mono<? extends String> handlePostRequestJson(ServerWebExchange exchange, Route route,
ServerHttpRequest request, String accessKeyName,
String headerAk, String body) {
JSONObject jsonObj = StrUtil.isBlank(body) ? new JSONObject() : JSONUtil.parseObj(body);
final IscRule rule = handleRule(headerAk, () -> Collections.singletonList(jsonObj.get(accessKeyName,
String.class, true)), route, accessKeyName);
ServerHttpRequest request, AccessKey accessKey,
JSONObject jsonObj, MultiValueMap<String, String> queryParams) {
final IscRule rule = handleRule(accessKey, () -> Collections.singletonList(jsonObj.get(accessKey.name,
String.class, true)), route, BODY);
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()));
return Mono.just(jsonObj.toString());
};
@ -177,16 +193,16 @@ public class CustomerGlobalFilter implements GlobalFilter, Ordered {
* @param chain 当前过滤链
* @param route 路由信息
* @param request 请求信息
* @param accessKeyName AK键名
* @param headerAk 头部AK信息
* @param accessKey AK信息
* @param queryParams URL查询参数
* @return 指示请求处理何时完成
*/
private Mono<Void> handleGetRequest(ServerWebExchange exchange, GatewayFilterChain chain, Route route,
ServerHttpRequest request, String accessKeyName, String headerAk) {
final MultiValueMap<String, String> queryParams = new LinkedMultiValueMap<>(request.getQueryParams());
final IscRule rule = handleRule(headerAk, () -> queryParams.get(accessKeyName), route, accessKeyName);
ServerHttpRequest request, AccessKey accessKey,
MultiValueMap<String, String> queryParams) {
final IscRule rule = handleRule(accessKey, () -> queryParams.get(accessKey.name), route, URI);
Supplier<Mono<Void>> rateLimiterAfterSupplier = () -> {
removeParam(headerAk, request, accessKeyName, () -> queryParams.remove(accessKeyName));
removeParam(accessKey, request, null, queryParams);
handleHiddenParams(route, queryParams, (next, map) -> {
Object value;
if (Objects.isNull(value = next.getValue())) {
@ -197,14 +213,17 @@ public class CustomerGlobalFilter implements GlobalFilter, Ordered {
map.add(next.getKey(), value.toString());
}
});
URI newUri = UriComponentsBuilder.fromUri(request.getURI())
.replaceQueryParams(unmodifiableMultiValueMap(queryParams)).build().toUri();
ServerHttpRequest updatedRequest = exchange.getRequest().mutate().uri(newUri).build();
return chain.filter(exchange.mutate().request(updatedRequest).build());
ServerHttpRequest req = exchange.getRequest().mutate().uri(getNewUri().apply(request, queryParams)).build();
return chain.filter(exchange.mutate().request(req).build());
};
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
public int getOrder() {
return -1;
@ -213,42 +232,64 @@ public class CustomerGlobalFilter implements GlobalFilter, Ordered {
/**
* 处理规则
*
* @param headerAk 头部AK
* @param accessKey AK信息
* @param valueSupplier valueList 生产者
* @param route 路由信息
* @param accessKeyName AK键名
* @param type AK类型
* @return AK对应服务规则信息
*/
private IscRule handleRule(String headerAk, Supplier<List<String>> valueSupplier, Route route, String accessKeyName) {
//获取AK
final String ak = GatewayUtils.getValue(headerAk, valueSupplier, () -> new AkRequireException(accessKeyName));
private IscRule handleRule(AccessKey accessKey, Supplier<List<String>> valueSupplier, Route route, AccessKeyType type) {
if(Objects.nonNull(valueSupplier) && Objects.nonNull(type)) {
handleAccessKey(accessKey, valueSupplier, type, true);
}
//获取规则
final IscRule rule = GatewayUtils.getRequiredValue(() -> GatewayUtils.getRule(ak, route.getId()),
() -> new RuleNotExistException(ak, route.getId()));
final IscRule rule = GatewayUtils.getRequiredValue(() -> GatewayUtils.getRule(accessKey.value, route.getId()),
() -> new RuleNotExistException(accessKey.value, route.getId()));
//是否到期
GatewayUtils.isBefore(rule, () -> new RuleExpiredException(rule.getExpire()));
//设置AK 到ID 为了传参方便
rule.setId(ak);
rule.setId(accessKey.value);
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 accessKeyName AK名称
* @param removeSupplier 删除参数提供者
* @param queryParams URI请求参数
* @param <T> 类型
*/
private <T> void removeParam(String headerAk, ServerHttpRequest request, String accessKeyName,
Supplier<T> removeSupplier) {
private <T> void removeParam(AccessKey accessKey, ServerHttpRequest request, Supplier<T> removeSupplier,
MultiValueMap<String, String> queryParams) {
//如果header中有AK,则删除
if (Objects.nonNull(headerAk)) {
if (HEADER.equals(accessKey.type)) {
final HttpHeaders headers = request.getHeaders();
if(headers.containsKey(accessKeyName)) {
request.mutate().headers(headMap -> headMap.remove(accessKeyName)).build();
if(headers.containsKey(accessKey.name)) {
request.mutate().headers(headMap -> headMap.remove(accessKey.name)).build();
}
}else if(URI.equals(accessKey.type)) {
queryParams.remove(accessKey.name);
}
if (Objects.nonNull(removeSupplier)) {
removeSupplier.get();
@ -311,4 +352,27 @@ public class CustomerGlobalFilter implements GlobalFilter, Ordered {
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.Mono;
import java.net.URI;
import java.sql.Date;
import java.time.Instant;
import java.util.List;
@ -50,9 +51,10 @@ public class GatewayUtils {
* @param exchange 当前服务交换器
* @param chain 当前过滤链
* @param publisher body体提供者
* @param newUri 新的URI去掉AccessKey
* @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);
HttpHeaders headers = new HttpHeaders();
headers.putAll(exchange.getRequest().getHeaders());
@ -60,14 +62,14 @@ public class GatewayUtils {
CachedBodyOutputMessage outputMessage = new CachedBodyOutputMessage(exchange, headers);
return bodyInserter.insert(outputMessage, new BodyInserterContext())
.then(Mono.defer(() -> {
ServerHttpRequest request = decorate(exchange, headers, outputMessage);
ServerHttpRequest request = decorate(exchange, headers, outputMessage, newUri);
return chain.filter(exchange.mutate().request(request).build());
}));
}
public static ServerHttpRequestDecorator decorate(ServerWebExchange exchange, HttpHeaders headers,
CachedBodyOutputMessage outputMessage) {
CachedBodyOutputMessage outputMessage, URI newUri) {
return new ServerHttpRequestDecorator(exchange.getRequest()) {
@Override
public HttpHeaders getHeaders() {
@ -86,6 +88,12 @@ public class GatewayUtils {
public Flux<DataBuffer> 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);
}
return before;