ApiSignInterceptor.java 4.95 KB
package com.infoloop.tianting.intercepter;

import cn.dev33.satoken.sign.SaSignUtil;
import cn.hutool.core.convert.Convert;
import com.infoloop.tianting.annotation.ApiEnterpriseParam;
import com.infoloop.tianting.annotation.ApiSign;
import com.infoloop.tianting.context.LoginContextHolder;
import com.infoloop.tianting.exception.ClientEndExceptions;
import com.infoloop.tianting.filter.CachedBodyHttpServletRequest;
import com.infoloop.tianting.utils.RequestUtil;
import lombok.extern.slf4j.Slf4j;
import org.springframework.http.MediaType;
import org.springframework.stereotype.Component;
import org.springframework.web.method.HandlerMethod;
import org.springframework.web.servlet.HandlerInterceptor;

import javax.servlet.http.HttpServletRequest;
import javax.servlet.http.HttpServletResponse;
import java.io.IOException;
import java.util.HashMap;
import java.util.Map;
import java.util.TreeMap;

@Slf4j
@Component
@SuppressWarnings("all")
public class ApiSignInterceptor implements HandlerInterceptor {

    private static final String TIMESTAMP = "timestamp";
    private static final String NONCE = "nonce";
    private static final String SIGN = "sign";

    @Override
    public boolean preHandle(HttpServletRequest request, HttpServletResponse response, Object handler) throws IOException {
        if (!(handler instanceof HandlerMethod)) {
            return true;
        }
        HandlerMethod method = (HandlerMethod) handler;
        if (!method.hasMethodAnnotation(ApiSign.class)) {
            return true;
        }
        final var annotation = method.getMethodAnnotation(ApiSign.class);
        final var paramMap = getParamMap(annotation, request);
        SaSignUtil.checkParamMap(paramMap);
        setCurrentEnterpriseId(request, annotation.enterpriseParam(), paramMap);
        return true;
    }

    private void setCurrentEnterpriseId(HttpServletRequest request, ApiEnterpriseParam enterpriseParam, Map<String, String> paramMap) {
        var enterpriseId = paramMap.get(enterpriseParam.value());
        if (enterpriseId == null && enterpriseParam.isReadHeader()) {
            enterpriseId = request.getHeader(enterpriseParam.value());
        }
        if (enterpriseId == null && enterpriseParam.required()) {
            throw ClientEndExceptions.ParameterInvalid.build(enterpriseParam.value() + " cannot be null");
        }
        if (enterpriseId != null) {
            LoginContextHolder.setRequestEnterpriseId(Convert.toInt(enterpriseId));
        }
    }

    private Map<String, String> getParamMap(ApiSign annotation, HttpServletRequest request) throws IOException {
        Map<String, String> paramMap = new TreeMap<>();
        switch (annotation.readFrom()) {
            case ALL:
                paramMap.putAll(RequestUtil.getFormParameters(request));
                if (isJson(request)) {
                    final var bodyMap = ((CachedBodyHttpServletRequest) request).getBodyMap();
                    paramMap.putAll(convertMapToStringValue(bodyMap));
                }
                break;
            case QUERY:
                paramMap.putAll(RequestUtil.getFormParameters(request));
                break;
            case BODY:
                if (isJson(request)) {
                    final var bodyMap = ((CachedBodyHttpServletRequest) request).getBodyMap();
                    paramMap.putAll(convertMapToStringValue(bodyMap));
                }
                break;
        }
        if (annotation.paramNames().length != 0) {
            paramMap = takeRequestParam(paramMap, annotation.paramNames());
        }
        return paramMap;
    }

    private Map<String, String> convertMapToStringValue(Map<String, Object> originalMap) {
        Map<String, String> stringMap = new HashMap<>();
        for (Map.Entry<String, Object> entry : originalMap.entrySet()) {
            final var value = entry.getValue();
            if (value == null) {
                continue;
            }
            if (value instanceof String) {
                stringMap.put(entry.getKey(), (String) value);
            } else {
                stringMap.put(entry.getKey(), String.valueOf(value));
            }
        }
        return stringMap;
    }

    private boolean isJson(HttpServletRequest request) {
        if (request.getContentType() != null) {
            return request.getContentType().contains(MediaType.APPLICATION_JSON_VALUE);
        }
        return false;
    }

    private Map<String, String> takeRequestParam(Map<String, String> map, String[] paramNames) {
        Map<String, String> paramMap = new TreeMap<>();
        paramMap.put(TIMESTAMP, map.get(TIMESTAMP));
        paramMap.put(NONCE, map.get(NONCE));
        paramMap.put(SIGN, map.get(SIGN));
        for (String paramName : paramNames) {
            paramMap.put(paramName, map.get(paramName));
        }
        return paramMap;
    }

    @Override
    public void afterCompletion(HttpServletRequest request, HttpServletResponse response, Object handler, Exception ex) {
        LoginContextHolder.clear();
    }

}