Initial commit

This commit is contained in:
ywlmac
2026-05-13 10:30:19 +08:00
commit a2a693885b
38 changed files with 5757 additions and 0 deletions
@@ -0,0 +1,433 @@
package hdsignature
import (
"crypto/elliptic"
"crypto/rand"
"encoding/asn1"
"encoding/base64"
"encoding/hex"
"encoding/pem"
"fmt"
"math/big"
"github.com/tjfoc/gmsm/sm2"
"github.com/tjfoc/gmsm/sm3"
sm2x509 "github.com/tjfoc/gmsm/x509"
)
var (
curve = sm2.P256Sm2()
)
type PublicKey struct {
Key *sm2.PublicKey
Code []byte
}
type PrivateKey struct {
Key *sm2.PrivateKey
Code []byte
}
func CreateMasterKey(mnemonic string) (*PrivateKey, *PublicKey, error) {
keyStr, err := deriveKeyString([]byte(mnemonic), []byte(ZXRootSeed))
if err != nil {
return nil, nil, err
}
skStr := keyStr[:HDWKeyLength]
code := keyStr[HDWKeyLength:]
d := new(big.Int).SetBytes(skStr)
sk, err := composePrivateKey(d, code)
if err != nil {
return nil, nil, err
}
pk := sk.GetPublicKey()
return sk, pk, nil
}
// Private key methods
func (sk *PrivateKey) GetPublicKey() *PublicKey {
return &PublicKey{
Key: &sk.Key.PublicKey,
Code: sk.Code,
}
}
func (sk *PrivateKey) GetPEM(password []byte) (string, error) {
var skDER []byte
var err error
if password == nil {
skDER, err = sm2x509.MarshalSm2UnecryptedPrivateKey(sk.Key)
if err != nil {
return "", err
}
} else {
skDER, err = sm2x509.MarshalSm2EcryptedPrivateKey(sk.Key, password)
if err != nil {
return "", err
}
}
keyASN1 := HDKeyASN1{
Key: skDER,
Code: sk.Code,
}
keyBytes, err := asn1.Marshal(keyASN1)
if err != nil {
return "", err
}
keyBlock := &pem.Block{
Type: ZXPrivateKey,
Bytes: keyBytes,
}
return string(pem.EncodeToMemory(keyBlock)), nil
}
func (sk *PrivateKey) GetPrunedPEM(password []byte) (string, error) {
var skDER []byte
var err error
if password == nil {
skDER, err = sm2x509.MarshalSm2UnecryptedPrivateKey(sk.Key)
if err != nil {
return "", err
}
} else {
skDER, err = sm2x509.MarshalSm2EcryptedPrivateKey(sk.Key, password)
if err != nil {
return "", err
}
}
keyBlock := &pem.Block{
Type: PEMPrivateKey,
Bytes: skDER,
}
return string(pem.EncodeToMemory(keyBlock)), nil
}
func (sk *PrivateKey) GetBase64String(password []byte) (string, error) {
var skDER []byte
var err error
if password == nil {
skDER, err = sm2x509.MarshalSm2UnecryptedPrivateKey(sk.Key)
if err != nil {
return "", err
}
} else {
skDER, err = sm2x509.MarshalSm2EcryptedPrivateKey(sk.Key, password)
if err != nil {
return "", err
}
}
return base64.StdEncoding.EncodeToString(skDER), nil
}
func (sk *PrivateKey) Sign(data string) ([]byte, error) {
r, s, err := sm2.Sm2Sign(sk.Key, []byte(data), []byte(SM2DefaultUID), rand.Reader)
if err != nil {
return nil, err
}
return asn1.Marshal(ECSignature{R: r, S: s,})
}
func (sk *PrivateKey) GetAddress() (string, error) {
return sk.GetPublicKey().GetAddress()
}
func (sk *PrivateKey) DeriveChildPrivateKey(index uint32) (*PrivateKey, error) {
if sk.Code == nil {
return nil, fmt.Errorf("parent key is pruned")
}
var keySeed []byte
indexBytes := uint32Bytes(index)
if index >= FirstHardenedKey {
skBytes := sk.Key.D.Bytes()
if len(skBytes) < HDWKeyLength {
skBytes = append([]byte{0}, skBytes...)
}
keySeed = append(skBytes, indexBytes...)
} else {
pkBytes := sk.GetPublicKey().Bytes()
keySeed = append(pkBytes, indexBytes...)
}
keyString, err := deriveKeyString(keySeed, sk.Code)
if err != nil {
return nil, err
}
skStr := keyString[:HDWKeyLength]
code := keyString[HDWKeyLength:]
d := new(big.Int).Mod(new(big.Int).Add(new(big.Int).SetBytes(skStr), sk.Key.D), curve.Params().N)
skChild, err := composePrivateKey(d, code)
if err != nil {
return nil, err
}
return skChild, nil
}
func (sk *PrivateKey) DeriveChildPublicKey(index uint32) (*PublicKey, error) {
skChild, err := sk.DeriveChildPrivateKey(index)
if err != nil {
return nil, err
}
return skChild.GetPublicKey(), nil
}
func (sk *PrivateKey) DeriveKeyPairObjFromFullPath(indexPath []uint32) (*PrivateKey, error) {
var err error
if indexPath == nil || len(indexPath) == 0 {
return nil, fmt.Errorf("empty index path")
}
skChild := sk
for _, index := range indexPath {
skChild, err = skChild.DeriveChildPrivateKey(index)
if err != nil {
return nil, err
}
}
return skChild, nil
}
func (sk *PrivateKey) CreateCSRFromTemplatePEM(templateStr string) ([]byte, error) {
templateBlock, _ := pem.Decode([]byte(templateStr))
if templateBlock == nil {
return nil, fmt.Errorf("invalid CSR template")
}
template, err := sm2x509.ParseCertificateRequest(templateBlock.Bytes)
if err != nil {
return nil, err
}
return sk.CreateCSRFromTemplateObj(template)
}
func (sk *PrivateKey) CreateCSRFromTemplateObj(template *sm2x509.CertificateRequest) ([]byte, error) {
csrDER, err := sm2x509.CreateCertificateRequest(rand.Reader, template, sk.Key)
if err != nil {
return nil, err
}
return pem.EncodeToMemory(&pem.Block{
Type: PEMCSR,
Bytes: csrDER,
}), nil
}
// Public key methods
func (pk *PublicKey) GetPEM() (string, error) {
pkDER, err := sm2x509.MarshalSm2PublicKey(pk.Key)
if err != nil {
return "", err
}
keyASN1 := HDKeyASN1{
Key: pkDER,
Code: pk.Code,
}
keyBytes, err := asn1.Marshal(keyASN1)
if err != nil {
return "", err
}
keyBlock := &pem.Block{
Type: ZXPublicKey,
Bytes: keyBytes,
}
return string(pem.EncodeToMemory(keyBlock)), nil
}
func (pk *PublicKey) GetPrunedPEM() (string, error) {
pkDER, err := sm2x509.MarshalSm2PublicKey(pk.Key)
if err != nil {
return "", err
}
keyBlock := &pem.Block{
Type: PEMPublicKey,
Bytes: pkDER,
}
return string(pem.EncodeToMemory(keyBlock)), nil
}
func (pk *PublicKey) GetBase64String() (string, error) {
pkDER, err := sm2x509.MarshalSm2PublicKey(pk.Key)
if err != nil {
return "", nil
}
return base64.StdEncoding.EncodeToString(pkDER), nil
}
func (pk *PublicKey) Verify(data string, sig []byte) (bool, error) {
sigEC := &ECSignature{}
_, err := asn1.Unmarshal(sig, sigEC)
if err != nil {
return false, err
}
return sm2.Sm2Verify(pk.Key, []byte(data), []byte(SM2DefaultUID), sigEC.R, sigEC.S), nil
}
func (pk *PublicKey) Bytes() []byte {
return elliptic.Marshal(curve, pk.Key.X, pk.Key.Y)
}
func (pk *PublicKey) GetAddress() (string, error) {
pkBytes := pk.Bytes()
sm3Hash := sm3.New()
_, err := sm3Hash.Write(pkBytes)
if err != nil {
return "", err
}
pkDgst := sm3Hash.Sum(nil)
if len(pkDgst) <= HDWAddressLength {
return "", fmt.Errorf("invalid public key")
}
addrBytes := pkDgst[:HDWAddressLength]
addrHex := hex.EncodeToString(addrBytes)
return ZXAddrPrefix + addrHex, nil
}
func (pk *PublicKey) DeriveChildPublicKey(index uint32) (*PublicKey, error) {
if pk.Code == nil {
return nil, fmt.Errorf("parent key is pruned")
}
if index >= FirstHardenedKey {
return nil, fmt.Errorf("only private key can derive keys for hardened child %v", index)
}
indexBytes := uint32Bytes(index)
keySeed := append(pk.Bytes(), indexBytes...)
keyString, err := deriveKeyString(keySeed, pk.Code)
if err != nil {
return nil, err
}
skStr := keyString[:HDWKeyLength]
code := keyString[HDWKeyLength:]
x, y := curve.ScalarBaseMult(new(big.Int).SetBytes(skStr).Bytes())
xChild, yChild := curve.Add(x, y, pk.Key.X, pk.Key.Y)
return &PublicKey{
Key: &sm2.PublicKey{
Curve: curve,
X: xChild,
Y: yChild,
},
Code: code,
}, nil
}
func ParsePrivateKey(skPEM []byte, password []byte) (*PrivateKey, error) {
skBlock, _ := pem.Decode(skPEM)
if skBlock == nil {
return nil, fmt.Errorf("invalid private key PEM")
}
keyASN1 := &HDKeyASN1{}
_, err := asn1.Unmarshal(skBlock.Bytes, keyASN1)
if err == nil {
skSM2, err := sm2x509.ParsePKCS8PrivateKey(keyASN1.Key, password)
if err != nil {
return nil, err
}
return &PrivateKey{
Key: skSM2,
Code: keyASN1.Code,
}, nil
}
skSM2, err := sm2x509.ParsePKCS8PrivateKey(skBlock.Bytes, password)
if err != nil {
return nil, err
}
return &PrivateKey{
Key: skSM2,
Code: nil,
}, nil
}
func ParsePublicKey(pkPEM []byte) (*PublicKey, error) {
pkBlock, _ := pem.Decode(pkPEM)
if pkBlock == nil {
return nil, fmt.Errorf("invalid public key PEM")
}
keyASN1 := &HDKeyASN1{}
_, err := asn1.Unmarshal(pkBlock.Bytes, keyASN1)
if err == nil {
pkSM2, err := sm2x509.ParseSm2PublicKey(keyASN1.Key)
if err != nil {
return nil, err
}
return &PublicKey{
Key: pkSM2,
Code: keyASN1.Code,
}, nil
}
pkSM2, err := sm2x509.ParseSm2PublicKey(pkBlock.Bytes)
if err != nil {
return nil, err
}
return &PublicKey{
Key: pkSM2,
Code: nil,
}, nil
}
func ParseBase64PrivateKey(skBase64 string, password []byte) (*PrivateKey, error) {
skDER, err := base64.StdEncoding.DecodeString(skBase64)
if err != nil {
return nil, err
}
key, err := sm2x509.ParsePKCS8PrivateKey(skDER, password)
if err != nil {
return nil, err
}
return &PrivateKey{
Key: key,
Code: nil,
}, nil
}
func ParseBase64PublicKey(pkBase64 string) (*PublicKey, error) {
pkDER, err := base64.StdEncoding.DecodeString(pkBase64)
if err != nil {
return nil, err
}
key, err := sm2x509.ParseSm2PublicKey(pkDER)
if err != nil {
return nil, err
}
return &PublicKey{
Key: key,
Code: nil,
}, nil
}
func (pk *PublicKey) DeriveKeyPairObjFromFullPath(indexPath []uint32) (*PublicKey, error) {
var err error
if indexPath == nil || len(indexPath) == 0 {
return nil, fmt.Errorf("empty index path")
}
pkChild := pk
for _, index := range indexPath {
pkChild, err = pkChild.DeriveChildPublicKey(index)
if err != nil {
return nil, err
}
}
return pkChild, nil
}
@@ -0,0 +1,134 @@
package hdsignature
import (
"bytes"
"crypto/hmac"
"encoding/binary"
"fmt"
"math/big"
"github.com/tjfoc/gmsm/sm2"
"github.com/tjfoc/gmsm/sm3"
)
const (
FirstHardenedKey uint32 = 1 << 31
HDWKeyLength = 32
HDWKeyStringLength = HDWKeyLength * 2
HDWAddressLength = 20
SM2DefaultUID = "1234567812345678"
ZXAddrPrefix = "ZX"
ZXRootSeed = "ZX seed"
ZXPublicKey = "ZX PUBLIC KEY"
ZXPrivateKey = "ZX PRIVATE KEY"
PEMPublicKey = "PUBLIC KEY"
PEMPrivateKey = "PRIVATE KEY"
PEMCSR = "CERTIFICATE REQUEST"
SigMaxLength = 73
SigMinLength = 70
PubKeyMaxLength = 150
PubKeyMinLength = 80
ZXAddrPrefixV1 = "zx"
ZXAddrPrefixV2 = "Zx"
ZXAddrPrefixV3 = "zX"
)
type HDKeyASN1 struct {
Key []byte
Code []byte
}
type ECSignature struct {
R *big.Int `json:"r"`
S *big.Int `json:"s"`
}
func SM3HMAC(key []byte, data []byte) ([]byte, error) {
hmacSM3 := hmac.New(sm3.New, key)
_, err := hmacSM3.Write(data)
if err != nil {
return nil, err
}
return hmacSM3.Sum(nil), nil
}
func SM3Hash(data []byte) ([]byte, error) {
hashSM3 := sm3.New()
_, err := hashSM3.Write(data)
if err != nil {
return nil, err
}
return hashSM3.Sum(nil), nil
}
func uint32Bytes(i uint32) []byte {
bytes := make([]byte, 4)
binary.BigEndian.PutUint32(bytes, i)
return bytes
}
func validatePrivateKey(key []byte) error {
if fmt.Sprintf("%x", key) == "0000000000000000000000000000000000000000000000000000000000000000" || //if the key is zero
bytes.Compare(key, curve.Params().N.Bytes()) >= 0 || //or is outside of the curve
len(key) != 32 { //or is too short
return fmt.Errorf("invalid integer")
}
return nil
}
func deriveKeyString(seed, code []byte) ([]byte, error) {
skL, err := SM3HMAC(code, seed)
if err != nil {
return nil, err
}
err = validatePrivateKey(skL)
for err != nil {
skL, err = SM3HMAC(code, skL)
if err != nil {
return nil, err
}
err = validatePrivateKey(skL)
}
skR, err := SM3HMAC(code, skL)
if err != nil {
return nil, err
}
return append(skL, skR...), nil
}
func composePrivateKey(d *big.Int, code []byte) (*PrivateKey, error) {
x, y := curve.ScalarBaseMult(d.Bytes())
pkSM2 := sm2.PublicKey{
Curve: curve,
X: x,
Y: y,
}
skSM2 := sm2.PrivateKey{
PublicKey: pkSM2,
D: d,
}
sk := &PrivateKey{
Key: &skSM2,
Code: code,
}
return sk, nil
}