package com.h5diversion.util;

import com.alibaba.fastjson.JSON;
import com.alibaba.fastjson.JSONObject;
import com.alibaba.fastjson.TypeReference;
import org.apache.commons.codec.binary.Base64;
import org.apache.commons.lang3.RandomStringUtils;
import org.apache.commons.lang3.StringUtils;

import javax.crypto.Cipher;
import javax.crypto.spec.SecretKeySpec;
import java.io.ByteArrayOutputStream;
import java.nio.charset.StandardCharsets;
import java.security.*;
import java.security.interfaces.RSAPrivateKey;
import java.security.interfaces.RSAPublicKey;
import java.security.spec.PKCS8EncodedKeySpec;
import java.security.spec.X509EncodedKeySpec;
import java.util.ArrayList;
import java.util.Arrays;
import java.util.HashMap;
import java.util.Map;
import java.util.TreeMap;
import java.util.stream.Collectors;

public class SecurityUtil {

  private static final String KEY_ALGORITHM = "RSA";

  private static final String SIGNATURE_ALGORITHM = "SHA256WithRSA";

  public static final String SIGN_DEFAULT_SALT = "fdsjkf43jds9432few";

  private static final String HEXDIGITS[] = { "0", "1", "2", "3", "4", "5", "6", "7", "8", "9", "a", "b", "c", "d", "e", "f" };

  /**
   * 加密块大小 - 如果内容大于117字节需要分段加密
   */
  private static final int MAX_ENCRYPT_BLOCK = 117;

  /**
   * RSA最大解密密文大小
   */
  private static final int MAX_DECRYPT_BLOCK = 128;


//====================== 渠道方生成 RSA 公私钥方法 ===========
//=========================      开始     ===================

  public static void main(String[] args) throws Exception {
    newRsaKeys();
  }

//====================== 渠道方生成 RSA 公私钥方法 ===========
//=========================      结束      ===================


//====================== 请求数据/响应数据  加密/解密/加签/验签 ===========
//=========================      开始     ===================


  /**
   *
   * 加密请求数据
   *
   * @param channelId     渠道方ID
   * @param method        方法名称
   * @param data          加密前的业务数据
   * @param rsaPublicKey  rsa公钥
   * @param rsaPrivateKey rsa私钥
   * @return 加密后的请求数据
   */
  public static String encryptRequest(String channelId, String method, String data, String rsaPublicKey, String rsaPrivateKey) {
    // 随机字符串 作为 “AES秘钥” (注意AES加解密128/192/256 bits.对应秘钥16/24/32位)
    String aesKey = RandomStringUtils.randomAlphanumeric(16);
    String params, key;
    try {
      // 使用AES秘钥加密得到data
      params = aesEncrypt(data, aesKey);
      // 使用RSA公钥 加密 “AES秘钥” 得到 key
      key = rsaEncryptWithPublicKeyToBase64(getPublicKey(rsaPublicKey), aesKey.getBytes(StandardCharsets.UTF_8));
    } catch (Exception e) {
      throw new RuntimeException("加密失败", e);
    }
    // 组装数据
    Map<String, String> request = new HashMap<>(6, 1F);
    request.put("channelId", channelId);
    request.put("t", String.valueOf(System.currentTimeMillis()));
    request.put("method", method);
    request.put("params", params);
    request.put("key", key);
    // 加签
    String signContent = getSignContent(request);
    String sign = signByPrivateKey(signContent, rsaPrivateKey);
    request.put("sign", sign);

    return JSON.toJSONString(request);
  }



  /**
   *
   * 解密响应数据
   *
   * @param response      响应数据
   * @param rsaPublicKey  rsa公钥
   * @param rsaPrivateKey rsa私钥
   * @return 解密后的业务数据
   */
  public static String decryptResponse(String response, String rsaPublicKey, String rsaPrivateKey) {
    Map<String, String> paramObj = JSON.parseObject(response, new TypeReference<Map<String, String>>() {});

    String sign = paramObj.get("sign");
    String key = paramObj.get("key");
    String data = paramObj.get("data");

    // 验签
    if (!verifySignByPublicKey(getSignContent(paramObj), sign, rsaPublicKey)) {
      throw new RuntimeException("验签失败");
    }

    try {
      // 使用私钥解密得到AES秘钥
      String aesKey = decryptByPrivateKey(key, rsaPrivateKey);
      // 使用AES秘钥解密params得到业务参数
      return aesDecrypt(data, aesKey);
    } catch (Exception e) {
      throw new RuntimeException("解密失败", e);
    }
  }



  /**
   * 解密请求数据
   *
   * @param request       原始请求数据
   * @param rsaPublicKey  rsa公钥
   * @param rsaPrivateKey rsa私钥
   * @return 解密后的业务数据
   */
  public static String decryptRequest(String request, String rsaPublicKey, String rsaPrivateKey) {
    Map<String, String> paramObj = JSON.parseObject(request, new TypeReference<Map<String, String>>() {});
    String sign = paramObj.get("sign");
    String key = paramObj.get("key");
    String params = paramObj.get("params");

    // 验签
    if (!verifySignByPublicKey(getSignContent(paramObj), sign, rsaPublicKey)) {
      throw new RuntimeException("验签失败");
    }

    try {
      // 使用私钥解密得到AES秘钥
      String aesKey = decryptByPrivateKey(key, rsaPrivateKey);
      // 使用AES秘钥解密params得到业务参数
      return aesDecrypt(params, aesKey);
    } catch (Exception e) {
      throw new RuntimeException("解密失败", e);
    }
  }

  /**
   * 加密响应数据
   *
   * @param code          返回码, 响应体中的code字段
   * @param msg           响应信息, 响应体中的msg字段
   * @param responseData  加密前的业务数据
   * @param rsaPublicKey  rsa公钥
   * @param rsaPrivateKey rsa私钥
   * @return 加密后的响应数据
   */
  public static String encryptResponse(String code, String msg, String responseData,
                                       String rsaPublicKey, String rsaPrivateKey) {
    // 随机字符串 作为 “AES秘钥” (注意AES加解密128/192/256 bits.对应秘钥16/24/32位)
    String aesKey = RandomStringUtils.randomAlphanumeric(16);
    String data, key;
    try {
      // 使用AES秘钥加密得到data
      data = aesEncrypt(responseData, aesKey);
      // 使用RSA公钥 加密 “AES秘钥” 得到 key
      key = rsaEncryptWithPublicKeyToBase64(getPublicKey(rsaPublicKey), aesKey.getBytes(StandardCharsets.UTF_8));
    } catch (Exception e) {
      throw new RuntimeException("加密失败", e);
    }
    // 组装数据
    Map<String, String> response = new HashMap<>(5, 1F);
    response.put("code", code);
    response.put("msg", msg);
    response.put("data", data);
    response.put("key", key);
    // 加签
    String signContent = getSignContent(response);
    String sign = signByPrivateKey(signContent, rsaPrivateKey);
    response.put("sign", sign);

    return JSON.toJSONString(response);
  }

//====================== 请求数据/响应数据  加密/解密/加签/验签 ===========
//=========================      结束    ===================


  // ========================  签名 相关方法  ==========================
  //=========================      开始    ===================
  public static String sign(Map<String, String> paraMap, String salt) {
    ArrayList<String> list = new ArrayList<String>();
    for (Map.Entry<String, String> entry : paraMap.entrySet()) {
      if (entry.getValue() != null && entry.getValue() != "") {
        list.add(entry.getKey() + "=" + entry.getValue() + "&");
      }
    }
    int size = list.size();
    String[] arrayToSort = list.toArray(new String[size]);
    Arrays.sort(arrayToSort, String.CASE_INSENSITIVE_ORDER);
    StringBuilder sb = new StringBuilder();
    for (int i = 0; i < size; i++) {
      sb.append(arrayToSort[i]);
    }
    String result = sb.toString();
    // 加盐
    result += "salt=" + (StringUtils.isBlank(salt) ? SIGN_DEFAULT_SALT : salt);
    result = md5Encode(result, "UTF-8");
    return result;
  }

  public static String sign(JSONObject jsonObj, String salt) {
    ArrayList<String> list = new ArrayList<String>();
    for (Map.Entry<String, Object> entry : jsonObj.entrySet()) {
      if (entry.getValue() != null && entry.getValue() != "") {
        list.add(entry.getKey() + "=" + entry.getValue() + "&");
      }
    }
    int size = list.size();
    String[] arrayToSort = list.toArray(new String[size]);
    Arrays.sort(arrayToSort, String.CASE_INSENSITIVE_ORDER);
    StringBuilder sb = new StringBuilder();
    for (int i = 0; i < size; i++) {
      sb.append(arrayToSort[i]);
    }
    String result = sb.toString();
    // 加盐
    result += "salt=" + (StringUtils.isBlank(salt) ? SIGN_DEFAULT_SALT : salt);
    result = md5Encode(result, "UTF-8");
    return result;
  }


  public static boolean verifySign(Map<String, String> paraMap, String sign, String salt) {
    paraMap.keySet().removeIf(key -> key.equals("sign"));
    return sign.equals(sign(paraMap, salt));
  }

  public static boolean verifySign(JSONObject jsonObj, String sign, String salt) {
    jsonObj.remove("sign");
    return sign.equals(sign(jsonObj, salt));
  }





//=========================      结束    ===================






//=====================  一下为内部方法，无需主动调用  ====================

  /**
   * 对签名转换升序排序
   */
  private static String getSignContent(Map<String, String> paramObj) {
    TreeMap<String, String> treeMap = convertMapToTreeMap(paramObj);
    return treeMap.entrySet().stream()
      .map(entry -> {
        if (!"sign".equals(entry.getKey())) {
          if (StringUtils.isNotEmpty(entry.getValue()) && !"null".equalsIgnoreCase(entry.getValue())) {
            return entry.getKey() + "=" + entry.getValue();
          }
        }
        return "";
      })
      .collect(Collectors.joining("&"));
  }

  private static TreeMap<String, String> convertMapToTreeMap(Map<String, String> paramMap) {
    TreeMap<String, String> treeMap = new TreeMap<>();
    if (paramMap == null) {
      return treeMap;
    }
    treeMap.putAll(paramMap);
    // 本身的签名信息不参与签名
    treeMap.remove("sign");
    return treeMap;
  }

  private static String signByPrivateKey(String content, String privateKey) {
    try {
      PrivateKey priKey = getPrivateKey(privateKey);
      Signature signature = Signature.getInstance(SIGNATURE_ALGORITHM);
      signature.initSign(priKey);
      signature.update(content.getBytes(StandardCharsets.UTF_8));
      byte[] signed = signature.sign();
      return new String(Base64.encodeBase64(signed), StandardCharsets.UTF_8);
    } catch (Exception e) {
      throw new RuntimeException("加签失败");
    }
  }

  private static boolean verifySignByPublicKey(String content, String sign, String publicKey) {
    try {
      KeyFactory keyFactory = KeyFactory.getInstance(KEY_ALGORITHM);
      byte[] encodedKey = Base64.decodeBase64(publicKey.getBytes(StandardCharsets.UTF_8));
      PublicKey pubKey = keyFactory.generatePublic(new X509EncodedKeySpec(encodedKey));
      Signature signature = Signature.getInstance(SIGNATURE_ALGORITHM);
      signature.initVerify(pubKey);
      signature.update(content.getBytes(StandardCharsets.UTF_8));
      return signature.verify(Base64.decodeBase64(sign.getBytes(StandardCharsets.UTF_8)));
    } catch (Exception e) {
      return false;
    }
  }

  /**
   * 获取私钥
   */
  private static PrivateKey getPrivateKey(String key) throws Exception {
    byte[] keyBytes = Base64.decodeBase64(key);
    PKCS8EncodedKeySpec keySpec = new PKCS8EncodedKeySpec(keyBytes);
    KeyFactory keyFactory = KeyFactory.getInstance(KEY_ALGORITHM);
    return keyFactory.generatePrivate(keySpec);
  }

  /**
   * 获取公钥
   */
  private static PublicKey getPublicKey(String key) throws Exception {
    byte[] keyBytes = Base64.decodeBase64(key);
    X509EncodedKeySpec keySpec = new X509EncodedKeySpec(keyBytes);
    KeyFactory keyFactory = KeyFactory.getInstance(KEY_ALGORITHM);
    return keyFactory.generatePublic(keySpec);
  }

  public static String aesEncrypt(String content, String key) throws Exception {
    Cipher cipher = Cipher.getInstance("AES/ECB/PKCS5Padding");
    byte[] keyBytes = key.getBytes(StandardCharsets.UTF_8);
    cipher.init(Cipher.ENCRYPT_MODE, new SecretKeySpec(keyBytes, "AES"));
    byte[] bytes = cipher.doFinal(content.getBytes(StandardCharsets.UTF_8));
    return Base64.encodeBase64String(bytes);
  }

  private static String rsaEncryptWithPublicKeyToBase64(PublicKey publicKey, byte[] rawData) throws Exception {
    final byte[] encryptBytes = encryptBytes(publicKey, rawData, Cipher.ENCRYPT_MODE);
    return Base64.encodeBase64String(encryptBytes);
  }

  private static byte[] encryptBytes(Key key, byte[] rawData, int encryptMode) throws Exception {
    String algorithm = key.getAlgorithm();
    final Cipher cipher = Cipher.getInstance(algorithm.toUpperCase());
    cipher.init(encryptMode, key);
    return cipher.doFinal(rawData);
  }

  private static String decryptByPrivateKey(String value, String key) throws Exception {
    byte[] data = Base64.decodeBase64(value);
    // 对私钥解密
    byte[] keyBytes = Base64.decodeBase64(key);
    PKCS8EncodedKeySpec pkcs8EncodedKeySpec = new PKCS8EncodedKeySpec(keyBytes);
    KeyFactory keyFactory = KeyFactory.getInstance(KEY_ALGORITHM);
    Key privateKey = keyFactory.generatePrivate(pkcs8EncodedKeySpec);
    // 对数据解密
    Cipher cipher = Cipher.getInstance(keyFactory.getAlgorithm());
    cipher.init(Cipher.DECRYPT_MODE, privateKey);
    return new String(doFinalBySegment(cipher, data, false), StandardCharsets.UTF_8);
  }

  private static byte[] doFinalBySegment(Cipher cipher, byte[] source, boolean isEncode) throws Exception {
    try (ByteArrayOutputStream out = new ByteArrayOutputStream()) {
      int blockSize = isEncode ? MAX_ENCRYPT_BLOCK : MAX_DECRYPT_BLOCK;
      if (source.length <= blockSize) {
        return cipher.doFinal(source);
      }
      int offsetIndex = 0, offset = 0, sourceLength = source.length;
      while (sourceLength - offset > 0) {
        int size = Math.min(sourceLength - offset, blockSize);
        byte[] buffer = cipher.doFinal(source, offset, size);
        out.write(buffer, 0, buffer.length);
        offsetIndex++;
        offset = offsetIndex * blockSize;
      }
      return out.toByteArray();
    }
  }


  /**
   * aes解密
   *
   * @param encryptBytes 密文
   * @param key          秘钥
   */
  public static String aesDecrypt(String encryptBytes, String key) throws Exception {
    Cipher cipher = Cipher.getInstance("AES/ECB/PKCS5Padding");
    byte[] keyBytes = key.getBytes(StandardCharsets.UTF_8);
    cipher.init(Cipher.DECRYPT_MODE, new SecretKeySpec(keyBytes, "AES"));
    byte[] decryptBytes = cipher.doFinal(Base64.decodeBase64(encryptBytes));
    return new String(decryptBytes);
  }



  /**
   * 创建RSA公钥和私钥对
   * <p>
   * 公钥：RSAUtils.PUBLIC_KEY
   * 私钥：RSAUtils.PRIVATE_KEY
   *
   * @throws NoSuchAlgorithmException 创建异常
   */
  private static void newRsaKeys() throws NoSuchAlgorithmException {
    KeyPairGenerator keyPairGen = KeyPairGenerator.getInstance(KEY_ALGORITHM);
    keyPairGen.initialize(1024);
    KeyPair keyPair = keyPairGen.generateKeyPair();
    RSAPublicKey publicKey = (RSAPublicKey) keyPair.getPublic();
    RSAPrivateKey privateKey = (RSAPrivateKey) keyPair.getPrivate();
    String publicKeyStr = Base64.encodeBase64String(publicKey.getEncoded());
    String privateKeyStr = Base64.encodeBase64String(privateKey.getEncoded());
    System.out.println("RSA::PrivateKey::" + privateKeyStr);
    System.out.println("RSA::PublicKey::" + publicKeyStr);
  }



  private static String md5Encode(String origin, String charsetname) {
    String resultString = null;
    try {
      resultString = new String(origin);
      MessageDigest md = MessageDigest.getInstance("MD5");
      if (charsetname == null || "".equals(charsetname)) {
        resultString = byteArrayToHexString(md.digest(resultString.getBytes()));
      } else {
        resultString = byteArrayToHexString(md.digest(resultString.getBytes(charsetname)));
      }
    } catch (Exception exception) {
    }
    return resultString;
  }

  private static String byteArrayToHexString(byte b[]) {
    StringBuffer resultSb = new StringBuffer();
    for (int i = 0; i < b.length; i++) {
      resultSb.append(byteToHexString(b[i]));
    }

    return resultSb.toString();
  }

  private static String byteToHexString(byte b) {
    int n = b;
    if (n < 0) {
      n += 256;
    }
    int d1 = n / 16;
    int d2 = n % 16;
    return HEXDIGITS[d1] + HEXDIGITS[d2];
  }

}
