Archived
重构 JWT 与 Transit 签名流程
This commit is contained in:
@@ -5,6 +5,7 @@ go 1.26.0
|
|||||||
require (
|
require (
|
||||||
github.com/go-chi/chi/v5 v5.3.2
|
github.com/go-chi/chi/v5 v5.3.2
|
||||||
github.com/go-jose/go-jose/v4 v4.1.5
|
github.com/go-jose/go-jose/v4 v4.1.5
|
||||||
|
github.com/go-viper/mapstructure/v2 v2.5.0
|
||||||
github.com/openbao/openbao/api/v2 v2.7.0
|
github.com/openbao/openbao/api/v2 v2.7.0
|
||||||
go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.71.0
|
go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.71.0
|
||||||
go.opentelemetry.io/otel v1.46.0
|
go.opentelemetry.io/otel v1.46.0
|
||||||
@@ -20,7 +21,6 @@ require (
|
|||||||
github.com/felixge/httpsnoop v1.1.0 // indirect
|
github.com/felixge/httpsnoop v1.1.0 // indirect
|
||||||
github.com/go-logr/logr v1.4.4 // indirect
|
github.com/go-logr/logr v1.4.4 // indirect
|
||||||
github.com/go-logr/stdr v1.2.2 // indirect
|
github.com/go-logr/stdr v1.2.2 // indirect
|
||||||
github.com/go-viper/mapstructure/v2 v2.5.0 // indirect
|
|
||||||
github.com/google/uuid v1.6.0 // indirect
|
github.com/google/uuid v1.6.0 // indirect
|
||||||
github.com/hashicorp/errwrap v1.1.0 // indirect
|
github.com/hashicorp/errwrap v1.1.0 // indirect
|
||||||
github.com/hashicorp/go-cleanhttp v0.5.2 // indirect
|
github.com/hashicorp/go-cleanhttp v0.5.2 // indirect
|
||||||
|
|||||||
+68
-65
@@ -2,18 +2,19 @@ package signing
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"encoding/base64"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/go-jose/go-jose/v4"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
const maximumTTL = 5 * time.Minute
|
||||||
|
|
||||||
type RS256Signer interface {
|
type RS256Signer interface {
|
||||||
ActiveKey(ctx context.Context) (SigningKey, error)
|
ActiveKey(context.Context) (SigningKey, error)
|
||||||
SignRS256(ctx context.Context, key SigningKey, signingInput []byte) ([]byte, error)
|
SignRS256(context.Context, SigningKey, []byte) ([]byte, error)
|
||||||
}
|
}
|
||||||
|
|
||||||
type SigningKey struct {
|
type SigningKey struct {
|
||||||
@@ -21,6 +22,14 @@ type SigningKey struct {
|
|||||||
Version int
|
Version int
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type IssueRequest struct {
|
||||||
|
Subject string
|
||||||
|
PrincipalName string
|
||||||
|
Audience string
|
||||||
|
JWTID string
|
||||||
|
Scope string
|
||||||
|
}
|
||||||
|
|
||||||
type Claims struct {
|
type Claims struct {
|
||||||
Issuer string `json:"iss"`
|
Issuer string `json:"iss"`
|
||||||
Subject string `json:"sub"`
|
Subject string `json:"sub"`
|
||||||
@@ -34,22 +43,29 @@ type Claims struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
type Issuer struct {
|
type Issuer struct {
|
||||||
|
issuer string
|
||||||
|
ttl time.Duration
|
||||||
signer RS256Signer
|
signer RS256Signer
|
||||||
now func() time.Time
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewIssuer(signer RS256Signer) (*Issuer, error) {
|
func NewIssuer(issuer string, ttl time.Duration, signer RS256Signer) (*Issuer, error) {
|
||||||
if signer == nil {
|
if issuer == "" || signer == nil {
|
||||||
return nil, errors.New("signer is required")
|
return nil, errors.New("issuer and signer are required")
|
||||||
}
|
}
|
||||||
return &Issuer{signer: signer, now: time.Now}, nil
|
if ttl <= 0 || ttl > maximumTTL {
|
||||||
|
return nil, errors.New("TTL must be between zero and five minutes")
|
||||||
|
}
|
||||||
|
return &Issuer{issuer: issuer, ttl: ttl, signer: signer}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (i *Issuer) Sign(ctx context.Context, claims Claims) (string, error) {
|
func (i *Issuer) Issue(ctx context.Context, request IssueRequest) (string, error) {
|
||||||
if err := validateClaims(claims, i.now()); err != nil {
|
return i.issueAt(ctx, request, time.Now())
|
||||||
|
}
|
||||||
|
|
||||||
|
func (i *Issuer) issueAt(ctx context.Context, request IssueRequest, now time.Time) (string, error) {
|
||||||
|
if err := validateIssueRequest(request); err != nil {
|
||||||
return "", err
|
return "", err
|
||||||
}
|
}
|
||||||
|
|
||||||
key, err := i.signer.ActiveKey(ctx)
|
key, err := i.signer.ActiveKey(ctx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return "", fmt.Errorf("select signing key: %w", err)
|
return "", fmt.Errorf("select signing key: %w", err)
|
||||||
@@ -58,63 +74,50 @@ func (i *Issuer) Sign(ctx context.Context, claims Claims) (string, error) {
|
|||||||
return "", errors.New("signer returned an invalid signing key")
|
return "", errors.New("signer returned an invalid signing key")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
claims := Claims{
|
||||||
|
Issuer: i.issuer,
|
||||||
|
Subject: request.Subject,
|
||||||
|
PrincipalName: request.PrincipalName,
|
||||||
|
Audience: []string{request.Audience},
|
||||||
|
IssuedAt: now.Unix(),
|
||||||
|
NotBefore: now.Unix(),
|
||||||
|
ExpiresAt: now.Add(i.ttl).Unix(),
|
||||||
|
JWTID: request.JWTID,
|
||||||
|
Scope: request.Scope,
|
||||||
|
}
|
||||||
|
signingInput, err := encodeSigningInput(key.ID, claims)
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
signature, err := i.signer.SignRS256(ctx, key, []byte(signingInput))
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("sign JWT: %w", err)
|
||||||
|
}
|
||||||
|
return serializeCompact(signingInput, signature), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func validateIssueRequest(request IssueRequest) error {
|
||||||
|
if request.Subject == "" || request.Audience == "" || request.JWTID == "" || request.Scope == "" {
|
||||||
|
return errors.New("required JWT input is missing")
|
||||||
|
}
|
||||||
|
if strings.ContainsAny(request.Subject, "\r\n") || strings.ContainsAny(request.JWTID, "\r\n") {
|
||||||
|
return errors.New("invalid JWT input")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func encodeSigningInput(keyID string, claims Claims) (string, error) {
|
||||||
|
header, err := json.Marshal(map[string]string{"alg": "RS256", "kid": keyID, "typ": "at+jwt"})
|
||||||
|
if err != nil {
|
||||||
|
return "", fmt.Errorf("encode protected header: %w", err)
|
||||||
|
}
|
||||||
payload, err := json.Marshal(claims)
|
payload, err := json.Marshal(claims)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return "", fmt.Errorf("encode claims: %w", err)
|
return "", fmt.Errorf("encode claims: %w", err)
|
||||||
}
|
}
|
||||||
|
return base64.RawURLEncoding.EncodeToString(header) + "." + base64.RawURLEncoding.EncodeToString(payload), nil
|
||||||
opaque := &contextSigner{ctx: ctx, signer: i.signer, key: key}
|
|
||||||
options := (&jose.SignerOptions{}).
|
|
||||||
WithType(jose.ContentType("at+jwt")).
|
|
||||||
WithHeader(jose.HeaderKey("kid"), key.ID)
|
|
||||||
signer, err := jose.NewSigner(jose.SigningKey{Algorithm: jose.RS256, Key: opaque}, options)
|
|
||||||
if err != nil {
|
|
||||||
return "", fmt.Errorf("create JWT signer: %w", err)
|
|
||||||
}
|
|
||||||
jws, err := signer.Sign(payload)
|
|
||||||
if err != nil {
|
|
||||||
return "", fmt.Errorf("sign JWT: %w", err)
|
|
||||||
}
|
|
||||||
compact, err := jws.CompactSerialize()
|
|
||||||
if err != nil {
|
|
||||||
return "", fmt.Errorf("serialize JWT: %w", err)
|
|
||||||
}
|
|
||||||
return compact, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
type contextSigner struct {
|
func serializeCompact(signingInput string, signature []byte) string {
|
||||||
ctx context.Context
|
return signingInput + "." + base64.RawURLEncoding.EncodeToString(signature)
|
||||||
signer RS256Signer
|
|
||||||
key SigningKey
|
|
||||||
}
|
|
||||||
|
|
||||||
func (*contextSigner) Public() *jose.JSONWebKey {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func (*contextSigner) Algs() []jose.SignatureAlgorithm {
|
|
||||||
return []jose.SignatureAlgorithm{jose.RS256}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (s *contextSigner) SignPayload(payload []byte, algorithm jose.SignatureAlgorithm) ([]byte, error) {
|
|
||||||
if algorithm != jose.RS256 {
|
|
||||||
return nil, errors.New("unsupported signing algorithm")
|
|
||||||
}
|
|
||||||
return s.signer.SignRS256(s.ctx, s.key, payload)
|
|
||||||
}
|
|
||||||
|
|
||||||
func validateClaims(claims Claims, now time.Time) error {
|
|
||||||
if claims.Issuer == "" || claims.Subject == "" || len(claims.Audience) != 1 || claims.JWTID == "" || claims.Scope == "" {
|
|
||||||
return errors.New("required JWT claim is missing")
|
|
||||||
}
|
|
||||||
if strings.ContainsAny(claims.Subject, "\r\n") || strings.ContainsAny(claims.JWTID, "\r\n") {
|
|
||||||
return errors.New("invalid JWT claim")
|
|
||||||
}
|
|
||||||
if claims.ExpiresAt <= claims.NotBefore || claims.NotBefore < claims.IssuedAt {
|
|
||||||
return errors.New("invalid JWT time window")
|
|
||||||
}
|
|
||||||
if time.Unix(claims.ExpiresAt, 0).After(now.Add(5 * time.Minute)) {
|
|
||||||
return errors.New("JWT lifetime exceeds five minutes")
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -21,24 +21,19 @@ func TestIssuerProducesInteroperableRS256JWT(t *testing.T) {
|
|||||||
t.Fatalf("generate key: %v", err)
|
t.Fatalf("generate key: %v", err)
|
||||||
}
|
}
|
||||||
signer := &localRS256Signer{keyID: "workload-sts-v1", key: privateKey}
|
signer := &localRS256Signer{keyID: "workload-sts-v1", key: privateKey}
|
||||||
issuer, err := NewIssuer(signer)
|
issuer, err := NewIssuer("https://identity.ad.ddupan.top", 5*time.Minute, signer)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("NewIssuer() error = %v", err)
|
t.Fatalf("NewIssuer() error = %v", err)
|
||||||
}
|
}
|
||||||
now := time.Unix(1_789_062_000, 0)
|
now := time.Unix(1_789_062_000, 0)
|
||||||
issuer.now = func() time.Time { return now }
|
|
||||||
|
|
||||||
token, err := issuer.Sign(context.Background(), Claims{
|
token, err := issuer.issueAt(context.Background(), IssueRequest{
|
||||||
Issuer: "https://identity.ad.ddupan.top",
|
|
||||||
Subject: "01993f4d-5e1a-7000-8000-000000000001",
|
Subject: "01993f4d-5e1a-7000-8000-000000000001",
|
||||||
PrincipalName: "ci/homelab-infra-plan",
|
PrincipalName: "ci/homelab-infra-plan",
|
||||||
Audience: []string{"https://bao.ad.ddupan.top:8200"},
|
Audience: "https://bao.ad.ddupan.top:8200",
|
||||||
IssuedAt: now.Unix(),
|
|
||||||
NotBefore: now.Unix(),
|
|
||||||
ExpiresAt: now.Add(5 * time.Minute).Unix(),
|
|
||||||
JWTID: "01993f4d-5e1a-7000-8000-000000000002",
|
JWTID: "01993f4d-5e1a-7000-8000-000000000002",
|
||||||
Scope: "bao.login",
|
Scope: "bao.login",
|
||||||
})
|
}, now)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Sign() error = %v", err)
|
t.Fatalf("Sign() error = %v", err)
|
||||||
}
|
}
|
||||||
@@ -79,33 +74,9 @@ func TestIssuerRejectsExcessiveLifetime(t *testing.T) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("generate key: %v", err)
|
t.Fatalf("generate key: %v", err)
|
||||||
}
|
}
|
||||||
issuer, err := NewIssuer(&localRS256Signer{keyID: "test", key: privateKey})
|
_, err = NewIssuer("https://identity.ad.ddupan.top", 5*time.Minute+time.Second, &localRS256Signer{keyID: "test", key: privateKey})
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("NewIssuer() error = %v", err)
|
|
||||||
}
|
|
||||||
now := time.Unix(1_789_062_000, 0)
|
|
||||||
issuer.now = func() time.Time { return now }
|
|
||||||
|
|
||||||
_, err = issuer.Sign(context.Background(), Claims{
|
|
||||||
Issuer: "https://identity.ad.ddupan.top",
|
|
||||||
Subject: "principal",
|
|
||||||
Audience: []string{"audience"},
|
|
||||||
IssuedAt: now.Unix(),
|
|
||||||
NotBefore: now.Unix(),
|
|
||||||
ExpiresAt: now.Add(5*time.Minute + time.Second).Unix(),
|
|
||||||
JWTID: "jti",
|
|
||||||
Scope: "scope",
|
|
||||||
})
|
|
||||||
if err == nil {
|
if err == nil {
|
||||||
t.Fatal("Sign() error = nil, want excessive lifetime error")
|
t.Fatal("NewIssuer() error = nil, want excessive lifetime error")
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestIntegerRejectsNonIntegralJSONNumber(t *testing.T) {
|
|
||||||
for _, value := range []any{1.5, json.Number("1.5")} {
|
|
||||||
if _, err := integer(value); err == nil {
|
|
||||||
t.Fatalf("integer(%v) error = nil, want invalid integer error", value)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+50
-81
@@ -5,16 +5,15 @@ import (
|
|||||||
"crypto/rsa"
|
"crypto/rsa"
|
||||||
"crypto/x509"
|
"crypto/x509"
|
||||||
"encoding/base64"
|
"encoding/base64"
|
||||||
"encoding/json"
|
|
||||||
"encoding/pem"
|
"encoding/pem"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
"math"
|
|
||||||
"sort"
|
"sort"
|
||||||
"strconv"
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
"github.com/go-jose/go-jose/v4"
|
"github.com/go-jose/go-jose/v4"
|
||||||
|
"github.com/go-viper/mapstructure/v2"
|
||||||
baoapi "github.com/openbao/openbao/api/v2"
|
baoapi "github.com/openbao/openbao/api/v2"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -24,12 +23,20 @@ type TransitSigner struct {
|
|||||||
key string
|
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) {
|
func NewTransitSigner(client *baoapi.Client, mount, key string) (*TransitSigner, error) {
|
||||||
if client == nil {
|
if client == nil {
|
||||||
return nil, errors.New("OpenBao client is required")
|
return nil, errors.New("OpenBao client is required")
|
||||||
}
|
}
|
||||||
mount = strings.Trim(mount, "/")
|
mount, key = strings.Trim(mount, "/"), strings.TrimSpace(key)
|
||||||
key = strings.TrimSpace(key)
|
|
||||||
if mount == "" || key == "" || strings.Contains(key, "/") {
|
if mount == "" || key == "" || strings.Contains(key, "/") {
|
||||||
return nil, errors.New("valid Transit mount and key are required")
|
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) {
|
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 {
|
if err != nil {
|
||||||
return SigningKey{}, fmt.Errorf("read Transit key: %w", err)
|
return SigningKey{}, err
|
||||||
}
|
}
|
||||||
if secret == nil || secret.Data == nil {
|
if data.LatestVersion < 1 {
|
||||||
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{}, 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) {
|
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")
|
return nil, errors.New("sign with Transit: invalid signing key")
|
||||||
}
|
}
|
||||||
secret, err := s.client.Logical().WriteWithContext(ctx, s.signPath(), map[string]any{
|
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 {
|
if secret == nil || secret.Data == nil {
|
||||||
return nil, errors.New("sign with Transit: empty response")
|
return nil, errors.New("sign with Transit: empty response")
|
||||||
}
|
}
|
||||||
|
|
||||||
encoded, ok := secret.Data["signature"].(string)
|
encoded, ok := secret.Data["signature"].(string)
|
||||||
if !ok {
|
if !ok {
|
||||||
return nil, errors.New("sign with Transit: missing signature")
|
return nil, errors.New("sign with Transit: missing signature")
|
||||||
}
|
}
|
||||||
parts := strings.SplitN(encoded, ":", 3)
|
return decodeTransitSignature(encoded, key.Version)
|
||||||
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) {
|
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 {
|
if err != nil {
|
||||||
return jose.JSONWebKeySet{}, fmt.Errorf("read Transit keys: %w", err)
|
return jose.JSONWebKeySet{}, err
|
||||||
}
|
}
|
||||||
if secret == nil || secret.Data == nil {
|
versions := make([]int, 0, len(data.Keys))
|
||||||
return jose.JSONWebKeySet{}, errors.New("read Transit keys: empty response")
|
for text := range data.Keys {
|
||||||
}
|
version, err := strconv.Atoi(text)
|
||||||
|
|
||||||
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 {
|
if err != nil || version < 1 {
|
||||||
return jose.JSONWebKeySet{}, errors.New("read Transit keys: invalid key version")
|
return jose.JSONWebKeySet{}, errors.New("read Transit keys: invalid key version")
|
||||||
}
|
}
|
||||||
versions = append(versions, version)
|
versions = append(versions, version)
|
||||||
}
|
}
|
||||||
sort.Ints(versions)
|
sort.Ints(versions)
|
||||||
|
set := jose.JSONWebKeySet{Keys: make([]jose.JSONWebKey, 0, len(versions))}
|
||||||
set := jose.JSONWebKeySet{Keys: make([]jose.JSONWebKey, 0, len(keys))}
|
|
||||||
for _, version := range versions {
|
for _, version := range versions {
|
||||||
value := keys[strconv.Itoa(version)]
|
publicKey, err := parseRSAPublicKey(data.Keys[strconv.Itoa(version)].PublicKey)
|
||||||
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 {
|
if err != nil {
|
||||||
return jose.JSONWebKeySet{}, fmt.Errorf("read Transit keys version %d: %w", version, err)
|
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
|
return set, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *TransitSigner) keyPath() string {
|
func (s *TransitSigner) readKeyData(ctx context.Context) (transitKeyData, error) {
|
||||||
return s.mount + "/keys/" + s.key
|
secret, err := s.client.Logical().ReadWithContext(ctx, s.keyPath())
|
||||||
}
|
|
||||||
|
|
||||||
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 {
|
if err != nil {
|
||||||
return 0, fmt.Errorf("invalid integer: %w", err)
|
return transitKeyData{}, fmt.Errorf("read Transit key: %w", err)
|
||||||
}
|
}
|
||||||
return integer(parsed)
|
if secret == nil || secret.Data == nil {
|
||||||
default:
|
return transitKeyData{}, errors.New("read Transit key: empty response")
|
||||||
return 0, fmt.Errorf("not an integer: %T", value)
|
|
||||||
}
|
}
|
||||||
|
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) {
|
func parseRSAPublicKey(value string) (*rsa.PublicKey, error) {
|
||||||
block, _ := pem.Decode([]byte(value))
|
block, _ := pem.Decode([]byte(value))
|
||||||
if block == nil {
|
if block == nil {
|
||||||
|
|||||||
@@ -30,24 +30,19 @@ func TestTransitSignerProducesJWTVerifiedByPublishedJWKS(t *testing.T) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("NewTransitSigner() error = %v", err)
|
t.Fatalf("NewTransitSigner() error = %v", err)
|
||||||
}
|
}
|
||||||
issuer, err := NewIssuer(signer)
|
issuer, err := NewIssuer("https://identity.ad.ddupan.top", 5*time.Minute, signer)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("NewIssuer() error = %v", err)
|
t.Fatalf("NewIssuer() error = %v", err)
|
||||||
}
|
}
|
||||||
now := time.Now().UTC().Truncate(time.Second)
|
now := time.Now().UTC().Truncate(time.Second)
|
||||||
issuer.now = func() time.Time { return now }
|
|
||||||
|
|
||||||
compact, err := issuer.Sign(context.Background(), Claims{
|
compact, err := issuer.issueAt(context.Background(), IssueRequest{
|
||||||
Issuer: "https://identity.ad.ddupan.top",
|
|
||||||
Subject: "01993f4d-5e1a-7000-8000-000000000001",
|
Subject: "01993f4d-5e1a-7000-8000-000000000001",
|
||||||
PrincipalName: "ci/homelab-infra-plan",
|
PrincipalName: "ci/homelab-infra-plan",
|
||||||
Audience: []string{"https://bao.ad.ddupan.top:8200"},
|
Audience: "https://bao.ad.ddupan.top:8200",
|
||||||
IssuedAt: now.Unix(),
|
|
||||||
NotBefore: now.Unix(),
|
|
||||||
ExpiresAt: now.Add(5 * time.Minute).Unix(),
|
|
||||||
JWTID: "01993f4d-5e1a-7000-8000-000000000002",
|
JWTID: "01993f4d-5e1a-7000-8000-000000000002",
|
||||||
Scope: "bao.login",
|
Scope: "bao.login",
|
||||||
})
|
}, now)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Sign() error = %v", err)
|
t.Fatalf("Sign() error = %v", err)
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user