重构 JWT 与 Transit 签名流程

This commit is contained in:
2026-09-11 16:43:46 +00:00
parent 0532fc25bf
commit 6fb166768f
5 changed files with 131 additions and 193 deletions
+52 -83
View File
@@ -5,16 +5,15 @@ import (
"crypto/rsa"
"crypto/x509"
"encoding/base64"
"encoding/json"
"encoding/pem"
"errors"
"fmt"
"math"
"sort"
"strconv"
"strings"
"github.com/go-jose/go-jose/v4"
"github.com/go-viper/mapstructure/v2"
baoapi "github.com/openbao/openbao/api/v2"
)
@@ -24,12 +23,20 @@ type TransitSigner struct {
key string
}
type transitKeyData struct {
LatestVersion int `mapstructure:"latest_version"`
Keys map[string]transitKeyVersion `mapstructure:"keys"`
}
type transitKeyVersion struct {
PublicKey string `mapstructure:"public_key"`
}
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)
mount, key = strings.Trim(mount, "/"), strings.TrimSpace(key)
if mount == "" || key == "" || strings.Contains(key, "/") {
return nil, errors.New("valid Transit mount and key are required")
}
@@ -37,23 +44,18 @@ func NewTransitSigner(client *baoapi.Client, mount, key string) (*TransitSigner,
}
func (s *TransitSigner) ActiveKey(ctx context.Context) (SigningKey, error) {
secret, err := s.client.Logical().ReadWithContext(ctx, s.keyPath())
data, err := s.readKeyData(ctx)
if err != nil {
return SigningKey{}, fmt.Errorf("read Transit key: %w", err)
return SigningKey{}, 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 {
if data.LatestVersion < 1 {
return SigningKey{}, errors.New("read Transit key: invalid latest_version")
}
return SigningKey{ID: s.keyID(version), Version: version}, nil
return SigningKey{ID: s.keyID(data.LatestVersion), Version: data.LatestVersion}, nil
}
func (s *TransitSigner) SignRS256(ctx context.Context, key SigningKey, input []byte) ([]byte, error) {
if key.ID != s.keyID(key.Version) || key.Version < 1 {
if key.Version < 1 || key.ID != s.keyID(key.Version) {
return nil, errors.New("sign with Transit: invalid signing key")
}
secret, err := s.client.Logical().WriteWithContext(ctx, s.signPath(), map[string]any{
@@ -67,57 +69,30 @@ func (s *TransitSigner) SignRS256(ctx context.Context, key SigningKey, input []b
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
return decodeTransitSignature(encoded, key.Version)
}
func (s *TransitSigner) JWKS(ctx context.Context) (jose.JSONWebKeySet, error) {
secret, err := s.client.Logical().ReadWithContext(ctx, s.keyPath())
data, err := s.readKeyData(ctx)
if err != nil {
return jose.JSONWebKeySet{}, fmt.Errorf("read Transit keys: %w", err)
return jose.JSONWebKeySet{}, 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)
versions := make([]int, 0, len(data.Keys))
for text := range data.Keys {
version, err := strconv.Atoi(text)
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))}
set := jose.JSONWebKeySet{Keys: make([]jose.JSONWebKey, 0, len(versions))}
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)
publicKey, err := parseRSAPublicKey(data.Keys[strconv.Itoa(version)].PublicKey)
if err != nil {
return jose.JSONWebKeySet{}, fmt.Errorf("read Transit keys version %d: %w", version, err)
}
@@ -131,43 +106,37 @@ func (s *TransitSigner) JWKS(ctx context.Context) (jose.JSONWebKeySet, error) {
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 (s *TransitSigner) readKeyData(ctx context.Context) (transitKeyData, error) {
secret, err := s.client.Logical().ReadWithContext(ctx, s.keyPath())
if err != nil {
return transitKeyData{}, fmt.Errorf("read Transit key: %w", err)
}
if secret == nil || secret.Data == nil {
return transitKeyData{}, errors.New("read Transit key: empty response")
}
var data transitKeyData
if err := mapstructure.Decode(secret.Data, &data); err != nil {
return transitKeyData{}, fmt.Errorf("decode Transit key metadata: %w", err)
}
return data, nil
}
func decodeTransitSignature(encoded string, version int) ([]byte, error) {
parts := strings.SplitN(encoded, ":", 3)
if len(parts) != 3 || parts[0] != "vault" || parts[1] != "v"+strconv.Itoa(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) 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 parseRSAPublicKey(value string) (*rsa.PublicKey, error) {
block, _ := pem.Decode([]byte(value))
if block == nil {