diff --git a/ruoyi-extend/ruoyi-isc-gateway/src/main/java/com/ruoyi/gateway/config/handler/GlobalErrorWebExceptionHandler.java b/ruoyi-extend/ruoyi-isc-gateway/src/main/java/com/ruoyi/gateway/config/handler/GlobalErrorWebExceptionHandler.java index fb8afeb61..e6684d255 100644 --- a/ruoyi-extend/ruoyi-isc-gateway/src/main/java/com/ruoyi/gateway/config/handler/GlobalErrorWebExceptionHandler.java +++ b/ruoyi-extend/ruoyi-isc-gateway/src/main/java/com/ruoyi/gateway/config/handler/GlobalErrorWebExceptionHandler.java @@ -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 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 result(Optional code, String message) { + private Map result(Optional code, String message, RequestPath path) { final HashMap 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; } diff --git a/ruoyi-extend/ruoyi-isc-gateway/src/main/java/com/ruoyi/gateway/filter/CustomerGlobalFilter.java b/ruoyi-extend/ruoyi-isc-gateway/src/main/java/com/ruoyi/gateway/filter/CustomerGlobalFilter.java index 1430d6585..247100ea2 100644 --- a/ruoyi-extend/ruoyi-isc-gateway/src/main/java/com/ruoyi/gateway/filter/CustomerGlobalFilter.java +++ b/ruoyi-extend/ruoyi-isc-gateway/src/main/java/com/ruoyi/gateway/filter/CustomerGlobalFilter.java @@ -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 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 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 queryParams) { + //处理 URL AccessKey + handleAccessKey(accessKey, () -> queryParams.get(accessKey.name), URI, false); + ServerRequest serverRequest = ServerRequest.create(exchange, HandlerStrategies.withDefaults().messageReaders()); final Mono 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 handlePostRequestFormUrlencoded(ServerWebExchange exchange, Route route, - ServerHttpRequest request, String accessKeyName, - String headerAk, String body) { + ServerHttpRequest request, AccessKey accessKey, + String body, MultiValueMap queryParams) { if (StrUtil.isBlank(body)) { - Assert.notBlank(headerAk, () -> new AkRequireException(accessKeyName)); + Assert.notBlank(accessKey.value, () -> new AkRequireException(accessKey.name)); } final List 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> rateLimiterAfterSupplier = () -> { - removeParam(headerAk, request, accessKeyName, null); - final List params = srcParams.stream().filter(param -> !accessKeyName.equals(param[0])) + removeParam(accessKey, request, null, queryParams); + final List 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 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 queryParams) { + final IscRule rule = handleRule(accessKey, () -> Collections.singletonList(jsonObj.get(accessKey.name, + String.class, true)), route, BODY); Supplier> 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 handleGetRequest(ServerWebExchange exchange, GatewayFilterChain chain, Route route, - ServerHttpRequest request, String accessKeyName, String headerAk) { - final MultiValueMap queryParams = new LinkedMultiValueMap<>(request.getQueryParams()); - final IscRule rule = handleRule(headerAk, () -> queryParams.get(accessKeyName), route, accessKeyName); + ServerHttpRequest request, AccessKey accessKey, + MultiValueMap queryParams) { + final IscRule rule = handleRule(accessKey, () -> queryParams.get(accessKey.name), route, URI); Supplier> 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, 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> valueSupplier, Route route, String accessKeyName) { - //获取AK - final String ak = GatewayUtils.getValue(headerAk, valueSupplier, () -> new AkRequireException(accessKeyName)); + private IscRule handleRule(AccessKey accessKey, Supplier> 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> 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 类型 */ - private void removeParam(String headerAk, ServerHttpRequest request, String accessKeyName, - Supplier removeSupplier) { + private void removeParam(AccessKey accessKey, ServerHttpRequest request, Supplier removeSupplier, + MultiValueMap 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; + } } diff --git a/ruoyi-extend/ruoyi-isc-gateway/src/main/java/com/ruoyi/gateway/utils/GatewayUtils.java b/ruoyi-extend/ruoyi-isc-gateway/src/main/java/com/ruoyi/gateway/utils/GatewayUtils.java index 107066969..58ebc3806 100644 --- a/ruoyi-extend/ruoyi-isc-gateway/src/main/java/com/ruoyi/gateway/utils/GatewayUtils.java +++ b/ruoyi-extend/ruoyi-isc-gateway/src/main/java/com/ruoyi/gateway/utils/GatewayUtils.java @@ -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 modifyBody(ServerWebExchange exchange, GatewayFilterChain chain, Mono publisher) { + public static Mono modifyBody(ServerWebExchange exchange, GatewayFilterChain chain, Mono 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 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;