ymt_v3_api-20260721-143604/ymt_v3_api/crypto.go

416 lines
12 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_api
import (
"bytes"
"crypto"
"crypto/aes"
"crypto/cipher"
"crypto/rand"
"crypto/rsa"
"crypto/sha256"
"crypto/x509"
"encoding/base64"
"encoding/json"
"encoding/pem"
"errors"
"fmt"
"sort"
"strings"
)
// ---------- AES-ECB 实现 ----------
// aesEcbEncrypt 使用 AES-ECB 模式加密,结果进行 Base64 编码。
func aesEcbEncrypt(plaintext []byte, key []byte) (string, error) {
block, err := aes.NewCipher(key)
if err != nil {
return "", err
}
plaintext = pkcs7Padding(plaintext, block.BlockSize())
ciphertext := make([]byte, len(plaintext))
// ECB 模式:逐块加密
for i := 0; i < len(plaintext); i += block.BlockSize() {
block.Encrypt(ciphertext[i:i+block.BlockSize()], plaintext[i:i+block.BlockSize()])
}
return base64.StdEncoding.EncodeToString(ciphertext), nil
}
// aesEcbDecrypt 解密 Base64 编码的 AES-ECB 密文。
func aesEcbDecrypt(ciphertext string, key []byte) ([]byte, error) {
data, err := base64.StdEncoding.DecodeString(ciphertext)
if err != nil {
return nil, err
}
block, err := aes.NewCipher(key)
if err != nil {
return nil, err
}
if len(data)%block.BlockSize() != 0 {
return nil, errors.New("ciphertext is not a multiple of the block size")
}
plaintext := make([]byte, len(data))
for i := 0; i < len(data); i += block.BlockSize() {
block.Decrypt(plaintext[i:i+block.BlockSize()], data[i:i+block.BlockSize()])
}
return pkcs7Unpadding(plaintext)
}
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, errors.New("invalid padding size")
}
unpadding := int(data[length-1])
if unpadding > length {
return nil, errors.New("invalid padding")
}
return data[:(length - unpadding)], nil
}
// ---------- SM4-CBC 实现 ----------
// sm4CbcEncrypt 使用 SM4-CBC 模式加密IV 为 16 字节全 0结果 Base64 编码。
func sm4CbcEncrypt(plaintext []byte, key []byte) (string, error) {
block, err := newSM4Cipher(key)
if err != nil {
return "", err
}
plaintext = pkcs7Padding(plaintext, block.BlockSize())
iv := make([]byte, block.BlockSize()) // 默认全 0 IV
ciphertext := make([]byte, len(plaintext))
mode := cipher.NewCBCEncrypter(block, iv)
mode.CryptBlocks(ciphertext, plaintext)
return base64.StdEncoding.EncodeToString(ciphertext), nil
}
// sm4CbcDecrypt 解密 Base64 编码的 SM4-CBC 密文。
func sm4CbcDecrypt(ciphertext string, key []byte) ([]byte, error) {
data, err := base64.StdEncoding.DecodeString(ciphertext)
if err != nil {
return nil, err
}
block, err := newSM4Cipher(key)
if err != nil {
return nil, err
}
if len(data)%block.BlockSize() != 0 {
return nil, errors.New("ciphertext is not a multiple of the block size")
}
iv := make([]byte, block.BlockSize())
plaintext := make([]byte, len(data))
mode := cipher.NewCBCDecrypter(block, iv)
mode.CryptBlocks(plaintext, data)
return pkcs7Unpadding(plaintext)
}
// ---------- SM4 算法实现(简化版,符合 GB/T 32907-2016 ----------
var sBox = [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 fk = [4]uint32{0xa3b1bac6, 0x56aa3350, 0x677d9197, 0xb27022dc}
var ck = [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,
}
type sm4Cipher struct {
enc []uint32
dec []uint32
}
func newSM4Cipher(key []byte) (cipher.Block, error) {
if len(key) != 16 {
return nil, errors.New("SM4 key must be 16 bytes")
}
c := &sm4Cipher{}
c.enc = make([]uint32, 32)
c.dec = make([]uint32, 32)
c.expandKey(key)
return c, nil
}
func (c *sm4Cipher) BlockSize() int { return 16 }
func (c *sm4Cipher) Encrypt(dst, src []byte) {
if len(src) < 16 || len(dst) < 16 {
panic("sm4: invalid buffer")
}
x := bytesToUint32s(src)
for i := 0; i < 32; i++ {
x = sm4Round(x, c.enc[i])
}
putUint32s(dst, x)
}
func (c *sm4Cipher) Decrypt(dst, src []byte) {
if len(src) < 16 || len(dst) < 16 {
panic("sm4: invalid buffer")
}
x := bytesToUint32s(src)
for i := 0; i < 32; i++ {
x = sm4Round(x, c.dec[i])
}
putUint32s(dst, x)
}
func (c *sm4Cipher) expandKey(key []byte) {
mk := bytesToUint32s(key)
k := make([]uint32, 36)
k[0] = mk[0] ^ fk[0]
k[1] = mk[1] ^ fk[1]
k[2] = mk[2] ^ fk[2]
k[3] = mk[3] ^ fk[3]
for i := 0; i < 32; i++ {
k[i+4] = k[i] ^ sm4T(k[i+1]^k[i+2]^k[i+3]^ck[i])
c.enc[i] = k[i+4]
c.dec[31-i] = k[i+4]
}
}
func sm4Round(x []uint32, rk uint32) []uint32 {
return []uint32{
x[1],
x[2],
x[3],
x[0] ^ sm4T(x[1]^x[2]^x[3]^rk),
}
}
func sm4T(x uint32) uint32 {
return sm4L(sm4Tau(x))
}
func sm4Tau(a uint32) uint32 {
return uint32(sBox[a>>24])<<24 |
uint32(sBox[a>>16&0xff])<<16 |
uint32(sBox[a>>8&0xff])<<8 |
uint32(sBox[a&0xff])
}
func sm4L(b uint32) uint32 {
return b ^ rotl(b, 2) ^ rotl(b, 10) ^ rotl(b, 18) ^ rotl(b, 24)
}
func rotl(x uint32, n uint) uint32 {
return (x << n) | (x >> (32 - n))
}
func bytesToUint32s(b []byte) []uint32 {
return []uint32{
uint32(b[0])<<24 | uint32(b[1])<<16 | uint32(b[2])<<8 | uint32(b[3]),
uint32(b[4])<<24 | uint32(b[5])<<16 | uint32(b[6])<<8 | uint32(b[7]),
uint32(b[8])<<24 | uint32(b[9])<<16 | uint32(b[10])<<8 | uint32(b[11]),
uint32(b[12])<<24 | uint32(b[13])<<16 | uint32(b[14])<<8 | uint32(b[15]),
}
}
func putUint32s(dst []byte, x []uint32) {
_ = dst[15]
dst[0] = byte(x[0] >> 24)
dst[1] = byte(x[0] >> 16)
dst[2] = byte(x[0] >> 8)
dst[3] = byte(x[0])
dst[4] = byte(x[1] >> 24)
dst[5] = byte(x[1] >> 16)
dst[6] = byte(x[1] >> 8)
dst[7] = byte(x[1])
dst[8] = byte(x[2] >> 24)
dst[9] = byte(x[2] >> 16)
dst[10] = byte(x[2] >> 8)
dst[11] = byte(x[2])
dst[12] = byte(x[3] >> 24)
dst[13] = byte(x[3] >> 16)
dst[14] = byte(x[3] >> 8)
dst[15] = byte(x[3])
}
// ---------- RSA 签名与验签 ----------
// rsaSign 使用 RSA 私钥对数据进行 SHA256 签名,返回 Base64 编码的签名。
func rsaSign(data []byte, priv *rsa.PrivateKey) (string, error) {
hashed := sha256.Sum256(data)
sig, err := rsa.SignPKCS1v15(rand.Reader, priv, crypto.SHA256, hashed[:])
if err != nil {
return "", err
}
return base64.StdEncoding.EncodeToString(sig), nil
}
// rsaVerify 使用 RSA 公钥验证 Base64 编码的签名。
func rsaVerify(data []byte, sign string, pub *rsa.PublicKey) error {
sig, err := base64.StdEncoding.DecodeString(sign)
if err != nil {
return err
}
hashed := sha256.Sum256(data)
return rsa.VerifyPKCS1v15(pub, crypto.SHA256, hashed[:], sig)
}
// ---------- 密钥解析 ----------
func parsePrivateKey(pemStr string) (*rsa.PrivateKey, error) {
block, _ := pem.Decode([]byte(pemStr))
if block == nil {
return nil, errors.New("failed to decode PEM block")
}
priv, err := x509.ParsePKCS8PrivateKey(block.Bytes)
if err != nil {
priv, err = x509.ParsePKCS1PrivateKey(block.Bytes)
if err != nil {
return nil, err
}
}
rsaPriv, ok := priv.(*rsa.PrivateKey)
if !ok {
return nil, errors.New("not an RSA private key")
}
return rsaPriv, nil
}
func parsePublicKey(pemStr string) (*rsa.PublicKey, error) {
block, _ := pem.Decode([]byte(pemStr))
if block == nil {
return nil, errors.New("failed to decode PEM block")
}
pub, err := x509.ParsePKIXPublicKey(block.Bytes)
if err != nil {
return nil, err
}
rsaPub, ok := pub.(*rsa.PublicKey)
if !ok {
return nil, errors.New("not an RSA public key")
}
return rsaPub, nil
}
// ---------- 业务参数排序 JSON ----------
// toSortedJSON 将结构体转为按 key 排序的 JSON 字符串,并忽略零值字段。
func toSortedJSON(v interface{}) ([]byte, error) {
// 先通过 json.Marshal 得到 map利用 omitempty 忽略零值
raw, err := json.Marshal(v)
if err != nil {
return nil, err
}
var m map[string]interface{}
if err := json.Unmarshal(raw, &m); err != nil {
return nil, err
}
// 过滤掉值为 nil 的键json.Unmarshal 会将 null 转为 nil
for k, val := range m {
if val == nil {
delete(m, k)
}
}
// 按 key 排序
keys := make([]string, 0, len(m))
for k := range m {
keys = append(keys, k)
}
sort.Strings(keys)
// 构建有序的 JSON
var buf bytes.Buffer
buf.WriteByte('{')
for i, k := range keys {
if i > 0 {
buf.WriteByte(',')
}
keyJSON, _ := json.Marshal(k)
valJSON, _ := json.Marshal(m[k])
buf.Write(keyJSON)
buf.WriteByte(':')
buf.Write(valJSON)
}
buf.WriteByte('}')
return buf.Bytes(), nil
}
// ---------- 回调验签辅助 ----------
// VerifyCallbackSign 验证回调请求的签名。
// body: 回调请求的原始 bodyJSON 字符串)
// appID: 应用 ID
// timestamp: 回调 Header 中的 Timestamp
// sign: 回调 Header 中的 Sign
func (c *Client) VerifyCallbackSign(body []byte, appID, timestamp, sign string) error {
// 1. 解析 body 中的 data 字段
var resp struct {
Data json.RawMessage `json:"data"`
}
if err := json.Unmarshal(body, &resp); err != nil {
return fmt.Errorf("unmarshal callback body: %w", err)
}
// 2. 将 data 转为排序 JSON
var dataMap map[string]interface{}
if err := json.Unmarshal(resp.Data, &dataMap); err != nil {
return fmt.Errorf("unmarshal data: %w", err)
}
// 过滤零值
for k, v := range dataMap {
if v == nil || isZeroValue(v) {
delete(dataMap, k)
}
}
plainBytes, err := json.Marshal(dataMap) // map 序列化自动排序
if err != nil {
return err
}
// 3. 加密得到 ciphertext
var ciphertext string
switch c.encryptType {
case "SM4":
ciphertext, err = sm4CbcEncrypt(plainBytes, c.encryptKey)
default:
ciphertext, err = aesEcbEncrypt(plainBytes, c.encryptKey)
}
if err != nil {
return err
}
// 4. 拼接签名字符串并验签
return c.verifySign(appID, timestamp, ciphertext, sign)
}
func isZeroValue(v interface{}) bool {
switch val := v.(type) {
case string:
return val == ""
case float64:
return val == 0
case bool:
return !val
default:
return false
}
}
// 为了兼容旧版 Go添加 strings 包的引用(实际未使用,但防止 import 被移除)
var _ = strings.TrimSpace