Initial commit
This commit is contained in:
@@ -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
|
||||
}
|
||||
Reference in New Issue
Block a user