ymt_v3-20260723170706/ymt_v3/crypto.go

550 lines
15 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

package ymt_v3
import (
"bytes"
"crypto"
"crypto/aes"
"crypto/rand"
"crypto/rsa"
"crypto/sha256"
"crypto/x509"
"encoding/base64"
"encoding/json"
"encoding/pem"
"fmt"
"io"
"reflect"
"sort"
"strings"
"time"
)
// ============================================================
// 时间戳工具
// ============================================================
// GenerateTimestamp 生成格式为 yyyy-MM-dd HH:mm:ss 的时间戳
func GenerateTimestamp() string {
return time.Now().Format("2006-01-02 15:04:05")
}
// ============================================================
// PKCS7 填充/去填充
// ============================================================
func pkcs7Padding(data []byte, blockSize int) []byte {
padding := blockSize - len(data)%blockSize
padText := bytes.Repeat([]byte{byte(padding)}, padding)
return append(data, padText...)
}
func pkcs7UnPadding(data []byte) ([]byte, error) {
length := len(data)
if length == 0 {
return nil, fmt.Errorf("数据为空")
}
padding := int(data[length-1])
if padding > length || padding == 0 {
return nil, fmt.Errorf("无效的填充")
}
for i := length - padding; i < length; i++ {
if data[i] != byte(padding) {
return nil, fmt.Errorf("无效的填充")
}
}
return data[:length-padding], nil
}
// ============================================================
// AES-ECB 加密/解密
// ============================================================
// AESECBEncrypt AES-ECB模式加密返回base64编码
func AESECBEncrypt(plaintext []byte, key []byte) (string, error) {
block, err := aes.NewCipher(key)
if err != nil {
return "", fmt.Errorf("创建AES cipher失败: %v", err)
}
padded := pkcs7Padding(plaintext, aes.BlockSize)
ciphertext := make([]byte, len(padded))
for i := 0; i < len(padded); i += aes.BlockSize {
block.Encrypt(ciphertext[i:i+aes.BlockSize], padded[i:i+aes.BlockSize])
}
return base64.StdEncoding.EncodeToString(ciphertext), nil
}
// AESECBDecrypt AES-ECB模式解密输入base64编码
func AESECBDecrypt(encryptedData string, key []byte) ([]byte, error) {
ciphertext, err := base64.StdEncoding.DecodeString(encryptedData)
if err != nil {
return nil, fmt.Errorf("base64解码失败: %v", err)
}
block, err := aes.NewCipher(key)
if err != nil {
return nil, fmt.Errorf("创建AES cipher失败: %v", err)
}
if len(ciphertext)%aes.BlockSize != 0 {
return nil, fmt.Errorf("密文长度不是块大小的整数倍")
}
plaintext := make([]byte, len(ciphertext))
for i := 0; i < len(ciphertext); i += aes.BlockSize {
block.Decrypt(plaintext[i:i+aes.BlockSize], ciphertext[i:i+aes.BlockSize])
}
plaintext, err = pkcs7UnPadding(plaintext)
if err != nil {
return nil, fmt.Errorf("去除填充失败: %v", err)
}
return plaintext, nil
}
// ============================================================
// SM4-CBC 加密/解密纯Go实现无外部依赖
// ============================================================
// SM4CBCEncrypt SM4-CBC模式加密返回base64编码IV前置
func SM4CBCEncrypt(plaintext []byte, key []byte) (string, error) {
if len(key) != 16 {
return "", fmt.Errorf("SM4密钥长度必须为16字节")
}
block, err := newSM4Cipher(key)
if err != nil {
return "", fmt.Errorf("创建SM4 cipher失败: %v", err)
}
padded := pkcs7Padding(plaintext, block.BlockSize())
iv := make([]byte, block.BlockSize())
if _, err := io.ReadFull(rand.Reader, iv); err != nil {
return "", fmt.Errorf("生成IV失败: %v", err)
}
mode := newCBCEncrypter(block, iv)
ciphertext := make([]byte, len(padded))
mode.CryptBlocks(ciphertext, padded)
result := append(iv, ciphertext...)
return base64.StdEncoding.EncodeToString(result), nil
}
// SM4CBCDecrypt SM4-CBC模式解密输入base64编码
func SM4CBCDecrypt(encryptedData string, key []byte) ([]byte, error) {
if len(key) != 16 {
return nil, fmt.Errorf("SM4密钥长度必须为16字节")
}
data, err := base64.StdEncoding.DecodeString(encryptedData)
if err != nil {
return nil, fmt.Errorf("base64解码失败: %v", err)
}
block, err := newSM4Cipher(key)
if err != nil {
return nil, fmt.Errorf("创建SM4 cipher失败: %v", err)
}
blockSize := block.BlockSize()
if len(data) < blockSize {
return nil, fmt.Errorf("数据长度不足")
}
iv := data[:blockSize]
ciphertext := data[blockSize:]
mode := newCBCDecrypter(block, iv)
plaintext := make([]byte, len(ciphertext))
mode.CryptBlocks(plaintext, ciphertext)
plaintext, err = pkcs7UnPadding(plaintext)
if err != nil {
return nil, fmt.Errorf("去除填充失败: %v", err)
}
return plaintext, nil
}
// ============================================================
// SM4 纯Go实现
// ============================================================
type sm4Cipher struct {
sk [32]uint32
}
func newSM4Cipher(key []byte) (*sm4Cipher, error) {
if len(key) != 16 {
return nil, fmt.Errorf("SM4密钥长度必须为16字节")
}
c := &sm4Cipher{}
c.sm4KeyInit(key)
return c, nil
}
func (c *sm4Cipher) BlockSize() int { return 16 }
func (c *sm4Cipher) Encrypt(dst, src []byte) {
c.sm4OneRound(dst, src, c.sk)
}
func (c *sm4Cipher) Decrypt(dst, src []byte) {
var rk [32]uint32
for i := 0; i < 32; i++ {
rk[i] = c.sk[31-i]
}
c.sm4OneRound(dst, src, rk)
}
type sm4CBCEncrypter struct {
b *sm4Cipher
iv []byte
}
func newCBCEncrypter(b *sm4Cipher, iv []byte) *sm4CBCEncrypter {
return &sm4CBCEncrypter{b: b, iv: append([]byte{}, iv...)}
}
func (c *sm4CBCEncrypter) CryptBlocks(dst, src []byte) {
blockSize := c.b.BlockSize()
iv := make([]byte, blockSize)
copy(iv, c.iv)
for i := 0; i < len(src); i += blockSize {
for j := 0; j < blockSize; j++ {
dst[i+j] = src[i+j] ^ iv[j]
}
c.b.Encrypt(dst[i:i+blockSize], dst[i:i+blockSize])
copy(iv, dst[i:i+blockSize])
}
}
type sm4CBCDecrypter struct {
b *sm4Cipher
iv []byte
}
func newCBCDecrypter(b *sm4Cipher, iv []byte) *sm4CBCDecrypter {
return &sm4CBCDecrypter{b: b, iv: append([]byte{}, iv...)}
}
func (c *sm4CBCDecrypter) CryptBlocks(dst, src []byte) {
blockSize := c.b.BlockSize()
iv := make([]byte, blockSize)
copy(iv, c.iv)
for i := 0; i < len(src); i += blockSize {
c.b.Decrypt(dst[i:i+blockSize], src[i:i+blockSize])
for j := 0; j < blockSize; j++ {
dst[i+j] ^= iv[j]
}
copy(iv, src[i:i+blockSize])
}
}
var sm4Sbox = [256]byte{
0xd6, 0x90, 0xe9, 0xfe, 0xcc, 0xe1, 0x3d, 0xb7, 0x16, 0xb6, 0x14, 0xc2, 0x28, 0xfb, 0x2c, 0x05,
0x2b, 0x67, 0x9a, 0x76, 0x2a, 0xbe, 0x04, 0xc3, 0xaa, 0x44, 0x13, 0x26, 0x49, 0x86, 0x06, 0x99,
0x9c, 0x42, 0x50, 0xf4, 0x91, 0xef, 0x98, 0x7a, 0x33, 0x54, 0x0b, 0x43, 0xed, 0xcf, 0xac, 0x62,
0xe4, 0xb3, 0x1c, 0xa9, 0xc9, 0x08, 0xe8, 0x95, 0x80, 0xdf, 0x94, 0xfa, 0x75, 0x8f, 0x3f, 0xa6,
0x47, 0x07, 0xa7, 0xfc, 0xf3, 0x73, 0x17, 0xba, 0x83, 0x59, 0x3c, 0x19, 0xe6, 0x85, 0x4f, 0xa8,
0x68, 0x6b, 0x81, 0xb2, 0x71, 0x64, 0xda, 0x8b, 0xf8, 0xeb, 0x0f, 0x4b, 0x70, 0x56, 0x9d, 0x35,
0x1e, 0x24, 0x0e, 0x5e, 0x63, 0x58, 0xd1, 0xa2, 0x25, 0x22, 0x7c, 0x3b, 0x01, 0x21, 0x78, 0x87,
0xd4, 0x00, 0x46, 0x57, 0x9f, 0xd3, 0x27, 0x52, 0x4c, 0x36, 0x02, 0xe7, 0xa0, 0xc4, 0xc8, 0x9e,
0xea, 0xbf, 0x8a, 0xd2, 0x40, 0xc7, 0x38, 0xb5, 0xa3, 0xf7, 0xf2, 0xce, 0xf9, 0x61, 0x15, 0xa1,
0xe0, 0xae, 0x5d, 0xa4, 0x9b, 0x34, 0x1a, 0x55, 0xad, 0x93, 0x32, 0x30, 0xf5, 0x8c, 0xb1, 0xe3,
0x1d, 0xf6, 0xe2, 0x2e, 0x82, 0x66, 0xca, 0x60, 0xc0, 0x29, 0x23, 0xab, 0x0d, 0x53, 0x4e, 0x6f,
0xd5, 0xdb, 0x37, 0x45, 0xde, 0xfd, 0x8e, 0x2f, 0x03, 0xff, 0x6a, 0x72, 0x6d, 0x6c, 0x5b, 0x51,
0x8d, 0x1b, 0xaf, 0x92, 0xbb, 0xdd, 0xbc, 0x7f, 0x11, 0xd9, 0x5c, 0x41, 0x1f, 0x10, 0x5a, 0xd8,
0x0a, 0xc1, 0x31, 0x88, 0xa5, 0xcd, 0x7b, 0xbd, 0x2d, 0x74, 0xd0, 0x12, 0xb8, 0xe5, 0xb4, 0xb0,
0x89, 0x69, 0x97, 0x4a, 0x0c, 0x96, 0x77, 0x7e, 0x65, 0xb9, 0xf1, 0x09, 0xc5, 0x6e, 0xc6, 0x84,
0x18, 0xf0, 0x7d, 0xec, 0x3a, 0xdc, 0x4d, 0x20, 0x79, 0xee, 0x5f, 0x3e, 0xd7, 0xcb, 0x39, 0x48,
}
var sm4FK = [4]uint32{0xa3b1bac6, 0x56aa3350, 0x677d9197, 0xb27022dc}
var sm4CK = [32]uint32{
0x00070e15, 0x1c232a31, 0x383f464d, 0x545b6269,
0x70777e85, 0x8c939aa1, 0xa8afb6bd, 0xc4cbd2d9,
0xe0e7eef5, 0xfc030a11, 0x181f262d, 0x343b4249,
0x50575e65, 0x6c737a81, 0x888f969d, 0xa4abb2b9,
0xc0c7ced5, 0xdce3eaf1, 0xf8ff060d, 0x141b2229,
0x30373e45, 0x4c535a61, 0x686f767d, 0x848b9299,
0xa0a7aeb5, 0xbcc3cad1, 0xd8dfe6ed, 0xf4fb0209,
0x10171e25, 0x2c333a41, 0x484f565d, 0x646b7279,
}
func (c *sm4Cipher) sm4KeyInit(key []byte) {
var mk [4]uint32
for i := 0; i < 4; i++ {
mk[i] = uint32(key[4*i])<<24 | uint32(key[4*i+1])<<16 | uint32(key[4*i+2])<<8 | uint32(key[4*i+3])
}
var k [36]uint32
for i := 0; i < 4; i++ {
k[i] = mk[i] ^ sm4FK[i]
}
for i := 0; i < 32; i++ {
k[i+4] = k[i] ^ sm4L1(k[i+1]^k[i+2]^k[i+3]^sm4CK[i])
c.sk[i] = k[i+4]
}
}
func (c *sm4Cipher) sm4OneRound(dst, src []byte, sk [32]uint32) {
var x [36]uint32
for i := 0; i < 4; i++ {
x[i] = uint32(src[4*i])<<24 | uint32(src[4*i+1])<<16 | uint32(src[4*i+2])<<8 | uint32(src[4*i+3])
}
for i := 0; i < 32; i++ {
x[i+4] = x[i] ^ sm4L2(x[i+1]^x[i+2]^x[i+3]^sk[i])
}
for i := 0; i < 4; i++ {
dst[4*i] = byte(x[35-i] >> 24)
dst[4*i+1] = byte(x[35-i] >> 16)
dst[4*i+2] = byte(x[35-i] >> 8)
dst[4*i+3] = byte(x[35-i])
}
}
func sm4L1(b uint32) uint32 {
return b ^ sm4Rotl(b, 2) ^ sm4Rotl(b, 10) ^ sm4Rotl(b, 18) ^ sm4Rotl(b, 24)
}
func sm4L2(b uint32) uint32 {
return b ^ sm4Rotl(b, 13) ^ sm4Rotl(b, 23)
}
func sm4Rotl(x uint32, n uint32) uint32 {
return (x << n) | (x >> (32 - n))
}
// ============================================================
// RSA 签名与验签
// ============================================================
// SignWithRSA 使用RSA私钥对字符串进行签名返回base64编码
func SignWithRSA(signStr string, privateKeyPEM string) (string, error) {
block, _ := pem.Decode([]byte(privateKeyPEM))
if block == nil {
return "", fmt.Errorf("解析PEM私钥失败")
}
privateKey, err := x509.ParsePKCS8PrivateKey(block.Bytes)
if err != nil {
privateKey, err = x509.ParsePKCS1PrivateKey(block.Bytes)
if err != nil {
return "", fmt.Errorf("解析私钥失败: %v", err)
}
}
rsaPrivateKey, ok := privateKey.(*rsa.PrivateKey)
if !ok {
return "", fmt.Errorf("不是RSA私钥")
}
h := sha256.New()
h.Write([]byte(signStr))
hashed := h.Sum(nil)
signature, err := rsa.SignPKCS1v15(rand.Reader, rsaPrivateKey, crypto.SHA256, hashed)
if err != nil {
return "", fmt.Errorf("签名失败: %v", err)
}
return base64.StdEncoding.EncodeToString(signature), nil
}
// VerifyWithRSA 使用RSA公钥验证签名
func VerifyWithRSA(signStr string, sign string, publicKeyPEM string) error {
block, _ := pem.Decode([]byte(publicKeyPEM))
if block == nil {
return fmt.Errorf("解析PEM公钥失败")
}
publicKey, err := x509.ParsePKIXPublicKey(block.Bytes)
if err != nil {
publicKey, err = x509.ParsePKCS1PublicKey(block.Bytes)
if err != nil {
return fmt.Errorf("解析公钥失败: %v", err)
}
}
rsaPublicKey, ok := publicKey.(*rsa.PublicKey)
if !ok {
return fmt.Errorf("不是RSA公钥")
}
signBytes, err := base64.StdEncoding.DecodeString(sign)
if err != nil {
return fmt.Errorf("base64解码签名失败: %v", err)
}
h := sha256.New()
h.Write([]byte(signStr))
hashed := h.Sum(nil)
return rsa.VerifyPKCS1v15(rsaPublicKey, crypto.SHA256, hashed, signBytes)
}
// ============================================================
// 业务参数加密/解密
// ============================================================
// EncryptBizParams 加密业务参数
// 将业务参数去掉零值按字母排序转JSON然后用指定算法加密
func EncryptBizParams(params interface{}, key []byte, encryptType string) (string, error) {
// 将结构体转为map去掉零值
m, err := structToMap(params)
if err != nil {
return "", fmt.Errorf("转换参数失败: %v", err)
}
// 按key排序
keys := make([]string, 0, len(m))
for k := range m {
keys = append(keys, k)
}
sort.Strings(keys)
// 构建有序map
orderedMap := make(map[string]interface{})
for _, k := range keys {
orderedMap[k] = m[k]
}
// 转JSON
plaintext, err := json.Marshal(orderedMap)
if err != nil {
return "", fmt.Errorf("JSON序列化失败: %v", err)
}
// 加密
switch encryptType {
case "aes":
return AESECBEncrypt(plaintext, key)
case "sm4":
return SM4CBCEncrypt(plaintext, key)
default:
return "", fmt.Errorf("不支持的加密类型: %s", encryptType)
}
}
// DecryptBizParams 解密业务参数
func DecryptBizParams(ciphertext string, key []byte, encryptType string) ([]byte, error) {
switch encryptType {
case "aes":
return AESECBDecrypt(ciphertext, key)
case "sm4":
return SM4CBCDecrypt(ciphertext, key)
default:
return nil, fmt.Errorf("不支持的加密类型: %s", encryptType)
}
}
// structToMap 将结构体转为map去掉零值
func structToMap(obj interface{}) (map[string]interface{}, error) {
result := make(map[string]interface{})
v := reflect.ValueOf(obj)
if v.Kind() == reflect.Ptr {
v = v.Elem()
}
if v.Kind() != reflect.Struct {
return nil, fmt.Errorf("不是结构体")
}
t := v.Type()
for i := 0; i < t.NumField(); i++ {
field := t.Field(i)
value := v.Field(i)
// 获取json tag
jsonTag := field.Tag.Get("json")
if jsonTag == "" || jsonTag == "-" {
continue
}
name := strings.Split(jsonTag, ",")[0]
// 检查omitempty
opts := strings.Split(jsonTag, ",")
omitempty := false
for _, opt := range opts[1:] {
if opt == "omitempty" {
omitempty = true
break
}
}
// 获取实际值
var val interface{}
switch value.Kind() {
case reflect.String:
val = value.String()
case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64:
val = value.Int()
case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64:
val = value.Uint()
case reflect.Float32, reflect.Float64:
val = value.Float()
case reflect.Bool:
val = value.Bool()
case reflect.Slice, reflect.Map:
if value.IsNil() {
if omitempty {
continue
}
val = value.Interface()
} else {
val = value.Interface()
}
case reflect.Ptr, reflect.Interface:
if value.IsNil() {
if omitempty {
continue
}
val = nil
} else {
val = value.Elem().Interface()
}
default:
val = value.Interface()
}
// 检查零值
if omitempty && isZeroValue(val) {
continue
}
result[name] = val
}
return result, nil
}
// isZeroValue 判断值是否为零值
func isZeroValue(v interface{}) bool {
if v == nil {
return true
}
rv := reflect.ValueOf(v)
switch rv.Kind() {
case reflect.String:
return rv.String() == ""
case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64:
return rv.Int() == 0
case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64:
return rv.Uint() == 0
case reflect.Float32, reflect.Float64:
return rv.Float() == 0
case reflect.Bool:
return !rv.Bool()
default:
return false
}
}