diff --git a/crypto.go b/crypto.go new file mode 100644 index 0000000..6c66e1d --- /dev/null +++ b/crypto.go @@ -0,0 +1,142 @@ +package xyshanghai + +import ( + "crypto/cipher" + "encoding/base64" + "encoding/hex" + "fmt" + + "github.com/tjfoc/gmsm/sm3" + "github.com/tjfoc/gmsm/sm4" +) + +// SM4Encrypt 使用SM4 ECB模式加密数据,返回base64编码字符串 +func SM4Encrypt(key, plaintext []byte) (string, error) { + block, err := sm4.NewCipher(key) + if err != nil { + return "", fmt.Errorf("sm4 new cipher: %w", err) + } + + // PKCS7填充 + plaintext = pkcs7Padding(plaintext, block.BlockSize()) + + // ECB模式 + mode := newECBEncrypter(block) + ciphertext := make([]byte, len(plaintext)) + mode.CryptBlocks(ciphertext, plaintext) + + return base64.StdEncoding.EncodeToString(ciphertext), nil +} + +// SM4Decrypt 解密base64编码的SM4 ECB密文,返回明文 +func SM4Decrypt(key []byte, ciphertextBase64 string) ([]byte, error) { + ciphertext, err := base64.StdEncoding.DecodeString(ciphertextBase64) + if err != nil { + return nil, fmt.Errorf("base64 decode: %w", err) + } + + block, err := sm4.NewCipher(key) + if err != nil { + return nil, fmt.Errorf("sm4 new cipher: %w", err) + } + + mode := newECBDecrypter(block) + plaintext := make([]byte, len(ciphertext)) + mode.CryptBlocks(plaintext, ciphertext) + + // 去除PKCS7填充 + plaintext, err = pkcs7Unpadding(plaintext, block.BlockSize()) + if err != nil { + return nil, fmt.Errorf("pkcs7 unpadding: %w", err) + } + + return plaintext, nil +} + +// SM3Sign 使用SM3加盐签名,返回十六进制字符串 +func SM3Sign(salt []byte, data string) string { + h := sm3.Sm3WithSalt(salt) + h.Write([]byte(data)) + return hex.EncodeToString(h.Sum(nil)) +} + +// pkcs7Padding PKCS7填充 +func pkcs7Padding(src []byte, blockSize int) []byte { + padding := blockSize - len(src)%blockSize + padtext := make([]byte, padding) + for i := range padtext { + padtext[i] = byte(padding) + } + return append(src, padtext...) +} + +// pkcs7Unpadding 去除PKCS7填充 +func pkcs7Unpadding(src []byte, blockSize int) ([]byte, error) { + length := len(src) + if length == 0 { + return nil, fmt.Errorf("empty data") + } + unpadding := int(src[length-1]) + if unpadding > blockSize || unpadding == 0 { + return nil, fmt.Errorf("invalid padding") + } + for i := length - unpadding; i < length; i++ { + if src[i] != byte(unpadding) { + return nil, fmt.Errorf("invalid padding") + } + } + return src[:length-unpadding], nil +} + +// ecb 实现ECB模式(标准库未提供) +type ecb struct { + b cipher.Block + blockSize int +} + +func newECB(b cipher.Block) *ecb { + return &ecb{ + b: b, + blockSize: b.BlockSize(), + } +} + +type ecbEncrypter ecb + +func newECBEncrypter(b cipher.Block) cipher.BlockMode { + return (*ecbEncrypter)(newECB(b)) +} + +func (x *ecbEncrypter) BlockSize() int { return x.blockSize } + +func (x *ecbEncrypter) CryptBlocks(dst, src []byte) { + if len(src)%x.blockSize != 0 { + panic("crypto/cipher: input not full blocks") + } + if len(dst) < len(src) { + panic("crypto/cipher: output smaller than input") + } + for i := 0; i < len(src); i += x.blockSize { + x.b.Encrypt(dst[i:i+x.blockSize], src[i:i+x.blockSize]) + } +} + +type ecbDecrypter ecb + +func newECBDecrypter(b cipher.Block) cipher.BlockMode { + return (*ecbDecrypter)(newECB(b)) +} + +func (x *ecbDecrypter) BlockSize() int { return x.blockSize } + +func (x *ecbDecrypter) CryptBlocks(dst, src []byte) { + if len(src)%x.blockSize != 0 { + panic("crypto/cipher: input not full blocks") + } + if len(dst) < len(src) { + panic("crypto/cipher: output smaller than input") + } + for i := 0; i < len(src); i += x.blockSize { + x.b.Decrypt(dst[i:i+x.blockSize], src[i:i+x.blockSize]) + } +} \ No newline at end of file