Archived
实现首轮协议与 Transit 签名 PoC
This commit is contained in:
@@ -0,0 +1,185 @@
|
||||
package signing
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rsa"
|
||||
"crypto/x509"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"encoding/pem"
|
||||
"errors"
|
||||
"fmt"
|
||||
"math"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/go-jose/go-jose/v4"
|
||||
baoapi "github.com/openbao/openbao/api/v2"
|
||||
)
|
||||
|
||||
type TransitSigner struct {
|
||||
client *baoapi.Client
|
||||
mount string
|
||||
key string
|
||||
}
|
||||
|
||||
func NewTransitSigner(client *baoapi.Client, mount, key string) (*TransitSigner, error) {
|
||||
if client == nil {
|
||||
return nil, errors.New("OpenBao client is required")
|
||||
}
|
||||
mount = strings.Trim(mount, "/")
|
||||
key = strings.TrimSpace(key)
|
||||
if mount == "" || key == "" || strings.Contains(key, "/") {
|
||||
return nil, errors.New("valid Transit mount and key are required")
|
||||
}
|
||||
return &TransitSigner{client: client, mount: mount, key: key}, nil
|
||||
}
|
||||
|
||||
func (s *TransitSigner) ActiveKey(ctx context.Context) (SigningKey, error) {
|
||||
secret, err := s.client.Logical().ReadWithContext(ctx, s.keyPath())
|
||||
if err != nil {
|
||||
return SigningKey{}, fmt.Errorf("read Transit key: %w", err)
|
||||
}
|
||||
if secret == nil || secret.Data == nil {
|
||||
return SigningKey{}, errors.New("read Transit key: empty response")
|
||||
}
|
||||
|
||||
version, err := integer(secret.Data["latest_version"])
|
||||
if err != nil || version < 1 {
|
||||
return SigningKey{}, errors.New("read Transit key: invalid latest_version")
|
||||
}
|
||||
return SigningKey{ID: s.keyID(version), Version: version}, nil
|
||||
}
|
||||
|
||||
func (s *TransitSigner) SignRS256(ctx context.Context, key SigningKey, input []byte) ([]byte, error) {
|
||||
if key.ID != s.keyID(key.Version) || key.Version < 1 {
|
||||
return nil, errors.New("sign with Transit: invalid signing key")
|
||||
}
|
||||
secret, err := s.client.Logical().WriteWithContext(ctx, s.signPath(), map[string]any{
|
||||
"input": base64.StdEncoding.EncodeToString(input),
|
||||
"key_version": key.Version,
|
||||
"signature_algorithm": "pkcs1v15",
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("sign with Transit: %w", err)
|
||||
}
|
||||
if secret == nil || secret.Data == nil {
|
||||
return nil, errors.New("sign with Transit: empty response")
|
||||
}
|
||||
|
||||
encoded, ok := secret.Data["signature"].(string)
|
||||
if !ok {
|
||||
return nil, errors.New("sign with Transit: missing signature")
|
||||
}
|
||||
parts := strings.SplitN(encoded, ":", 3)
|
||||
if len(parts) != 3 || parts[0] != "vault" || parts[1] != "v"+strconv.Itoa(key.Version) {
|
||||
return nil, errors.New("sign with Transit: unexpected signature version")
|
||||
}
|
||||
signature, err := base64.StdEncoding.DecodeString(parts[2])
|
||||
if err != nil {
|
||||
return nil, errors.New("sign with Transit: invalid signature encoding")
|
||||
}
|
||||
return signature, nil
|
||||
}
|
||||
|
||||
func (s *TransitSigner) JWKS(ctx context.Context) (jose.JSONWebKeySet, error) {
|
||||
secret, err := s.client.Logical().ReadWithContext(ctx, s.keyPath())
|
||||
if err != nil {
|
||||
return jose.JSONWebKeySet{}, fmt.Errorf("read Transit keys: %w", err)
|
||||
}
|
||||
if secret == nil || secret.Data == nil {
|
||||
return jose.JSONWebKeySet{}, errors.New("read Transit keys: empty response")
|
||||
}
|
||||
|
||||
keys, ok := secret.Data["keys"].(map[string]any)
|
||||
if !ok {
|
||||
return jose.JSONWebKeySet{}, errors.New("read Transit keys: invalid keys")
|
||||
}
|
||||
versions := make([]int, 0, len(keys))
|
||||
for versionText := range keys {
|
||||
version, err := strconv.Atoi(versionText)
|
||||
if err != nil || version < 1 {
|
||||
return jose.JSONWebKeySet{}, errors.New("read Transit keys: invalid key version")
|
||||
}
|
||||
versions = append(versions, version)
|
||||
}
|
||||
sort.Ints(versions)
|
||||
|
||||
set := jose.JSONWebKeySet{Keys: make([]jose.JSONWebKey, 0, len(keys))}
|
||||
for _, version := range versions {
|
||||
value := keys[strconv.Itoa(version)]
|
||||
metadata, ok := value.(map[string]any)
|
||||
if !ok {
|
||||
return jose.JSONWebKeySet{}, errors.New("read Transit keys: invalid key metadata")
|
||||
}
|
||||
publicPEM, ok := metadata["public_key"].(string)
|
||||
if !ok {
|
||||
return jose.JSONWebKeySet{}, errors.New("read Transit keys: missing public key")
|
||||
}
|
||||
publicKey, err := parseRSAPublicKey(publicPEM)
|
||||
if err != nil {
|
||||
return jose.JSONWebKeySet{}, fmt.Errorf("read Transit keys version %d: %w", version, err)
|
||||
}
|
||||
set.Keys = append(set.Keys, jose.JSONWebKey{
|
||||
Key: publicKey,
|
||||
KeyID: s.keyID(version),
|
||||
Algorithm: string(jose.RS256),
|
||||
Use: "sig",
|
||||
})
|
||||
}
|
||||
return set, nil
|
||||
}
|
||||
|
||||
func (s *TransitSigner) keyPath() string {
|
||||
return s.mount + "/keys/" + s.key
|
||||
}
|
||||
|
||||
func (s *TransitSigner) signPath() string {
|
||||
return s.mount + "/sign/" + s.key + "/sha2-256"
|
||||
}
|
||||
|
||||
func (s *TransitSigner) keyID(version int) string {
|
||||
return s.key + "-v" + strconv.Itoa(version)
|
||||
}
|
||||
|
||||
func integer(value any) (int, error) {
|
||||
switch value := value.(type) {
|
||||
case int:
|
||||
return value, nil
|
||||
case int64:
|
||||
if value > int64(^uint(0)>>1) || value < -int64(^uint(0)>>1)-1 {
|
||||
return 0, errors.New("integer out of range")
|
||||
}
|
||||
return int(value), nil
|
||||
case float64:
|
||||
if math.IsNaN(value) || math.IsInf(value, 0) || math.Trunc(value) != value || value > float64(^uint(0)>>1) || value < -float64(^uint(0)>>1)-1 {
|
||||
return 0, errors.New("invalid integer value")
|
||||
}
|
||||
return int(value), nil
|
||||
case json.Number:
|
||||
parsed, err := value.Int64()
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("invalid integer: %w", err)
|
||||
}
|
||||
return integer(parsed)
|
||||
default:
|
||||
return 0, fmt.Errorf("not an integer: %T", value)
|
||||
}
|
||||
}
|
||||
|
||||
func parseRSAPublicKey(value string) (*rsa.PublicKey, error) {
|
||||
block, _ := pem.Decode([]byte(value))
|
||||
if block == nil {
|
||||
return nil, errors.New("invalid PEM public key")
|
||||
}
|
||||
parsed, err := x509.ParsePKIXPublicKey(block.Bytes)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("parse public key: %w", err)
|
||||
}
|
||||
publicKey, ok := parsed.(*rsa.PublicKey)
|
||||
if !ok {
|
||||
return nil, errors.New("Transit key is not RSA")
|
||||
}
|
||||
return publicKey, nil
|
||||
}
|
||||
Reference in New Issue
Block a user