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 }