From 6fb166768f822cb8715dce5d8d557b6e572bcc73 Mon Sep 17 00:00:00 2001 From: panxiao81 Date: Fri, 11 Sep 2026 16:43:46 +0000 Subject: [PATCH] =?UTF-8?q?=E9=87=8D=E6=9E=84=20JWT=20=E4=B8=8E=20Transit?= =?UTF-8?q?=20=E7=AD=BE=E5=90=8D=E6=B5=81=E7=A8=8B?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- go.mod | 2 +- internal/signing/jwt.go | 133 +++++++++--------- internal/signing/jwt_test.go | 41 +----- internal/signing/transit.go | 135 +++++++------------ internal/signing/transit_integration_test.go | 13 +- 5 files changed, 131 insertions(+), 193 deletions(-) diff --git a/go.mod b/go.mod index 852763b..e2db0ea 100644 --- a/go.mod +++ b/go.mod @@ -5,6 +5,7 @@ go 1.26.0 require ( github.com/go-chi/chi/v5 v5.3.2 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 go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.71.0 go.opentelemetry.io/otel v1.46.0 @@ -20,7 +21,6 @@ require ( github.com/felixge/httpsnoop v1.1.0 // indirect github.com/go-logr/logr v1.4.4 // 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/hashicorp/errwrap v1.1.0 // indirect github.com/hashicorp/go-cleanhttp v0.5.2 // indirect diff --git a/internal/signing/jwt.go b/internal/signing/jwt.go index 31211bc..d661613 100644 --- a/internal/signing/jwt.go +++ b/internal/signing/jwt.go @@ -2,18 +2,19 @@ package signing import ( "context" + "encoding/base64" "encoding/json" "errors" "fmt" "strings" "time" - - "github.com/go-jose/go-jose/v4" ) +const maximumTTL = 5 * time.Minute + type RS256Signer interface { - ActiveKey(ctx context.Context) (SigningKey, error) - SignRS256(ctx context.Context, key SigningKey, signingInput []byte) ([]byte, error) + ActiveKey(context.Context) (SigningKey, error) + SignRS256(context.Context, SigningKey, []byte) ([]byte, error) } type SigningKey struct { @@ -21,6 +22,14 @@ type SigningKey struct { Version int } +type IssueRequest struct { + Subject string + PrincipalName string + Audience string + JWTID string + Scope string +} + type Claims struct { Issuer string `json:"iss"` Subject string `json:"sub"` @@ -34,22 +43,29 @@ type Claims struct { } type Issuer struct { + issuer string + ttl time.Duration signer RS256Signer - now func() time.Time } -func NewIssuer(signer RS256Signer) (*Issuer, error) { - if signer == nil { - return nil, errors.New("signer is required") +func NewIssuer(issuer string, ttl time.Duration, signer RS256Signer) (*Issuer, error) { + if issuer == "" || signer == nil { + 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) { - if err := validateClaims(claims, i.now()); err != nil { +func (i *Issuer) Issue(ctx context.Context, request IssueRequest) (string, error) { + 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 } - key, err := i.signer.ActiveKey(ctx) if err != nil { 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") } + 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) if err != nil { return "", fmt.Errorf("encode claims: %w", err) } - - 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 + return base64.RawURLEncoding.EncodeToString(header) + "." + base64.RawURLEncoding.EncodeToString(payload), nil } -type contextSigner struct { - ctx context.Context - 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 +func serializeCompact(signingInput string, signature []byte) string { + return signingInput + "." + base64.RawURLEncoding.EncodeToString(signature) } diff --git a/internal/signing/jwt_test.go b/internal/signing/jwt_test.go index 5ad29b4..beb6cf9 100644 --- a/internal/signing/jwt_test.go +++ b/internal/signing/jwt_test.go @@ -21,24 +21,19 @@ func TestIssuerProducesInteroperableRS256JWT(t *testing.T) { t.Fatalf("generate key: %v", err) } 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 { t.Fatalf("NewIssuer() error = %v", err) } now := time.Unix(1_789_062_000, 0) - issuer.now = func() time.Time { return now } - token, err := issuer.Sign(context.Background(), Claims{ - Issuer: "https://identity.ad.ddupan.top", + token, err := issuer.issueAt(context.Background(), IssueRequest{ Subject: "01993f4d-5e1a-7000-8000-000000000001", PrincipalName: "ci/homelab-infra-plan", - Audience: []string{"https://bao.ad.ddupan.top:8200"}, - IssuedAt: now.Unix(), - NotBefore: now.Unix(), - ExpiresAt: now.Add(5 * time.Minute).Unix(), + Audience: "https://bao.ad.ddupan.top:8200", JWTID: "01993f4d-5e1a-7000-8000-000000000002", Scope: "bao.login", - }) + }, now) if err != nil { t.Fatalf("Sign() error = %v", err) } @@ -79,33 +74,9 @@ func TestIssuerRejectsExcessiveLifetime(t *testing.T) { if err != nil { t.Fatalf("generate key: %v", err) } - issuer, err := NewIssuer(&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", - }) + _, err = NewIssuer("https://identity.ad.ddupan.top", 5*time.Minute+time.Second, &localRS256Signer{keyID: "test", key: privateKey}) if err == nil { - t.Fatal("Sign() 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) - } + t.Fatal("NewIssuer() error = nil, want excessive lifetime error") } } diff --git a/internal/signing/transit.go b/internal/signing/transit.go index 1227ef2..607f233 100644 --- a/internal/signing/transit.go +++ b/internal/signing/transit.go @@ -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 { diff --git a/internal/signing/transit_integration_test.go b/internal/signing/transit_integration_test.go index 9e28057..939df42 100644 --- a/internal/signing/transit_integration_test.go +++ b/internal/signing/transit_integration_test.go @@ -30,24 +30,19 @@ func TestTransitSignerProducesJWTVerifiedByPublishedJWKS(t *testing.T) { if err != nil { 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 { t.Fatalf("NewIssuer() error = %v", err) } now := time.Now().UTC().Truncate(time.Second) - issuer.now = func() time.Time { return now } - compact, err := issuer.Sign(context.Background(), Claims{ - Issuer: "https://identity.ad.ddupan.top", + compact, err := issuer.issueAt(context.Background(), IssueRequest{ Subject: "01993f4d-5e1a-7000-8000-000000000001", PrincipalName: "ci/homelab-infra-plan", - Audience: []string{"https://bao.ad.ddupan.top:8200"}, - IssuedAt: now.Unix(), - NotBefore: now.Unix(), - ExpiresAt: now.Add(5 * time.Minute).Unix(), + Audience: "https://bao.ad.ddupan.top:8200", JWTID: "01993f4d-5e1a-7000-8000-000000000002", Scope: "bao.login", - }) + }, now) if err != nil { t.Fatalf("Sign() error = %v", err) }