Archived
实现首轮协议与 Transit 签名 PoC
This commit is contained in:
@@ -0,0 +1,51 @@
|
||||
package httpapi
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"net/http"
|
||||
|
||||
"github.com/go-chi/chi/v5"
|
||||
"go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp"
|
||||
"go.opentelemetry.io/otel/metric"
|
||||
"go.opentelemetry.io/otel/propagation"
|
||||
"go.opentelemetry.io/otel/trace"
|
||||
)
|
||||
|
||||
type Endpoints struct {
|
||||
Metadata http.Handler
|
||||
Token http.Handler
|
||||
JWKS http.Handler
|
||||
Health http.Handler
|
||||
Ready http.Handler
|
||||
}
|
||||
|
||||
type Telemetry struct {
|
||||
TracerProvider trace.TracerProvider
|
||||
MeterProvider metric.MeterProvider
|
||||
Propagators propagation.TextMapPropagator
|
||||
}
|
||||
|
||||
func NewRouter(endpoints Endpoints, telemetry Telemetry) (http.Handler, error) {
|
||||
if endpoints.Metadata == nil || endpoints.Token == nil || endpoints.JWKS == nil || endpoints.Health == nil || endpoints.Ready == nil {
|
||||
return nil, errors.New("all HTTP endpoints are required")
|
||||
}
|
||||
|
||||
router := chi.NewRouter()
|
||||
router.Get("/.well-known/oauth-authorization-server", endpoints.Metadata.ServeHTTP)
|
||||
router.Post("/oauth2/token", endpoints.Token.ServeHTTP)
|
||||
router.Get("/oauth2/jwks", endpoints.JWKS.ServeHTTP)
|
||||
router.Get("/healthz", endpoints.Health.ServeHTTP)
|
||||
router.Get("/readyz", endpoints.Ready.ServeHTTP)
|
||||
|
||||
options := make([]otelhttp.Option, 0, 3)
|
||||
if telemetry.TracerProvider != nil {
|
||||
options = append(options, otelhttp.WithTracerProvider(telemetry.TracerProvider))
|
||||
}
|
||||
if telemetry.MeterProvider != nil {
|
||||
options = append(options, otelhttp.WithMeterProvider(telemetry.MeterProvider))
|
||||
}
|
||||
if telemetry.Propagators != nil {
|
||||
options = append(options, otelhttp.WithPropagators(telemetry.Propagators))
|
||||
}
|
||||
return otelhttp.NewHandler(router, "workload-sts.http", options...), nil
|
||||
}
|
||||
@@ -0,0 +1,95 @@
|
||||
package httpapi
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
sdktrace "go.opentelemetry.io/otel/sdk/trace"
|
||||
"go.opentelemetry.io/otel/sdk/trace/tracetest"
|
||||
)
|
||||
|
||||
func TestRouterExposesOnlySpecifiedMethods(t *testing.T) {
|
||||
endpoint := http.HandlerFunc(func(response http.ResponseWriter, _ *http.Request) {
|
||||
response.WriteHeader(http.StatusNoContent)
|
||||
})
|
||||
router, err := NewRouter(Endpoints{
|
||||
Metadata: endpoint,
|
||||
Token: endpoint,
|
||||
JWKS: endpoint,
|
||||
Health: endpoint,
|
||||
Ready: endpoint,
|
||||
}, Telemetry{})
|
||||
if err != nil {
|
||||
t.Fatalf("NewRouter() error = %v", err)
|
||||
}
|
||||
|
||||
tests := []struct {
|
||||
method string
|
||||
path string
|
||||
status int
|
||||
}{
|
||||
{http.MethodGet, "/.well-known/oauth-authorization-server", http.StatusNoContent},
|
||||
{http.MethodPost, "/oauth2/token", http.StatusNoContent},
|
||||
{http.MethodGet, "/oauth2/jwks", http.StatusNoContent},
|
||||
{http.MethodGet, "/healthz", http.StatusNoContent},
|
||||
{http.MethodGet, "/readyz", http.StatusNoContent},
|
||||
{http.MethodGet, "/oauth2/token", http.StatusMethodNotAllowed},
|
||||
{http.MethodPost, "/oauth2/jwks", http.StatusMethodNotAllowed},
|
||||
{http.MethodGet, "/unknown", http.StatusNotFound},
|
||||
}
|
||||
|
||||
for _, test := range tests {
|
||||
t.Run(test.method+" "+test.path, func(t *testing.T) {
|
||||
request := httptest.NewRequest(test.method, test.path, nil)
|
||||
response := httptest.NewRecorder()
|
||||
router.ServeHTTP(response, request)
|
||||
if response.Code != test.status {
|
||||
t.Fatalf("status = %d, want %d", response.Code, test.status)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestRouterRequiresEveryEndpoint(t *testing.T) {
|
||||
if _, err := NewRouter(Endpoints{}, Telemetry{}); err == nil {
|
||||
t.Fatal("NewRouter() error = nil, want missing endpoint error")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRouterTelemetryDoesNotCaptureCredentialValues(t *testing.T) {
|
||||
const canary = "canary-secret-token"
|
||||
recorder := tracetest.NewSpanRecorder()
|
||||
tracerProvider := sdktrace.NewTracerProvider(sdktrace.WithSpanProcessor(recorder))
|
||||
t.Cleanup(func() { _ = tracerProvider.Shutdown(t.Context()) })
|
||||
|
||||
endpoint := http.HandlerFunc(func(response http.ResponseWriter, _ *http.Request) {
|
||||
response.WriteHeader(http.StatusNoContent)
|
||||
})
|
||||
router, err := NewRouter(Endpoints{
|
||||
Metadata: endpoint,
|
||||
Token: endpoint,
|
||||
JWKS: endpoint,
|
||||
Health: endpoint,
|
||||
Ready: endpoint,
|
||||
}, Telemetry{TracerProvider: tracerProvider})
|
||||
if err != nil {
|
||||
t.Fatalf("NewRouter() error = %v", err)
|
||||
}
|
||||
|
||||
request := httptest.NewRequest(http.MethodGet, "/oauth2/jwks?subject_token="+canary, nil)
|
||||
request.Header.Set("Authorization", "Bearer "+canary)
|
||||
response := httptest.NewRecorder()
|
||||
router.ServeHTTP(response, request)
|
||||
|
||||
spans := recorder.Ended()
|
||||
if len(spans) != 1 {
|
||||
t.Fatalf("ended spans = %d, want 1", len(spans))
|
||||
}
|
||||
for _, attribute := range spans[0].Attributes() {
|
||||
if strings.Contains(attribute.Value.Emit(), canary) {
|
||||
t.Fatalf("span attribute %q captured credential canary", attribute.Key)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,78 @@
|
||||
package protocol
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net/url"
|
||||
"regexp"
|
||||
"strings"
|
||||
)
|
||||
|
||||
const (
|
||||
TokenExchangeGrantType = "urn:ietf:params:oauth:grant-type:token-exchange"
|
||||
JWTTokenType = "urn:ietf:params:oauth:token-type:jwt"
|
||||
AccessTokenType = "urn:ietf:params:oauth:token-type:access_token"
|
||||
)
|
||||
|
||||
var clientIDPattern = regexp.MustCompile(`^[a-z0-9](?:[a-z0-9-]{1,61}[a-z0-9])$`)
|
||||
|
||||
type ExchangeRequest struct {
|
||||
ClientID string
|
||||
SubjectToken string
|
||||
SubjectTokenType string
|
||||
RequestedTokenType string
|
||||
Audience string
|
||||
Scopes []string
|
||||
}
|
||||
|
||||
type RequestError struct {
|
||||
Code string
|
||||
Description string
|
||||
}
|
||||
|
||||
func (e *RequestError) Error() string {
|
||||
return fmt.Sprintf("%s: %s", e.Code, e.Description)
|
||||
}
|
||||
|
||||
func ParseExchangeRequest(form url.Values) (ExchangeRequest, error) {
|
||||
request := ExchangeRequest{
|
||||
ClientID: form.Get("client_id"),
|
||||
SubjectToken: form.Get("subject_token"),
|
||||
SubjectTokenType: form.Get("subject_token_type"),
|
||||
RequestedTokenType: form.Get("requested_token_type"),
|
||||
Audience: form.Get("audience"),
|
||||
Scopes: strings.Fields(form.Get("scope")),
|
||||
}
|
||||
|
||||
if form.Get("grant_type") != TokenExchangeGrantType {
|
||||
return ExchangeRequest{}, requestError("unsupported_grant_type", "unsupported grant_type")
|
||||
}
|
||||
if !clientIDPattern.MatchString(request.ClientID) {
|
||||
return ExchangeRequest{}, requestError("invalid_request", "client_id must use 3-63 lowercase letters, digits, or interior hyphens")
|
||||
}
|
||||
if request.SubjectToken == "" {
|
||||
return ExchangeRequest{}, requestError("invalid_request", "subject_token is required")
|
||||
}
|
||||
if request.SubjectTokenType != JWTTokenType {
|
||||
return ExchangeRequest{}, requestError("invalid_request", "unsupported subject_token_type")
|
||||
}
|
||||
if request.RequestedTokenType != AccessTokenType {
|
||||
return ExchangeRequest{}, requestError("invalid_request", "unsupported requested_token_type")
|
||||
}
|
||||
if request.Audience == "" || len(form["audience"]) != 1 {
|
||||
return ExchangeRequest{}, requestError("invalid_target", "exactly one audience is required")
|
||||
}
|
||||
if len(request.Scopes) == 0 {
|
||||
return ExchangeRequest{}, requestError("invalid_scope", "scope is required")
|
||||
}
|
||||
for _, unsupported := range []string{"resource", "actor_token", "actor_token_type"} {
|
||||
if form.Has(unsupported) {
|
||||
return ExchangeRequest{}, requestError("invalid_request", unsupported+" is not supported")
|
||||
}
|
||||
}
|
||||
|
||||
return request, nil
|
||||
}
|
||||
|
||||
func requestError(code, description string) *RequestError {
|
||||
return &RequestError{Code: code, Description: description}
|
||||
}
|
||||
@@ -0,0 +1,69 @@
|
||||
package protocol
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"net/url"
|
||||
"reflect"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestParseExchangeRequest(t *testing.T) {
|
||||
request, err := ParseExchangeRequest(validForm())
|
||||
if err != nil {
|
||||
t.Fatalf("ParseExchangeRequest() error = %v", err)
|
||||
}
|
||||
|
||||
if request.ClientID != "homelab-infra-ci" {
|
||||
t.Fatalf("ClientID = %q", request.ClientID)
|
||||
}
|
||||
if !reflect.DeepEqual(request.Scopes, []string{"bao.login", "ssh.certificate"}) {
|
||||
t.Fatalf("Scopes = %#v", request.Scopes)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseExchangeRequestRejectsInvalidInput(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
mutate func(url.Values)
|
||||
code string
|
||||
}{
|
||||
{name: "missing client id", mutate: func(v url.Values) { v.Del("client_id") }, code: "invalid_request"},
|
||||
{name: "short client id", mutate: func(v url.Values) { v.Set("client_id", "ci") }, code: "invalid_request"},
|
||||
{name: "uppercase client id", mutate: func(v url.Values) { v.Set("client_id", "CI-worker") }, code: "invalid_request"},
|
||||
{name: "leading hyphen", mutate: func(v url.Values) { v.Set("client_id", "-ci-worker") }, code: "invalid_request"},
|
||||
{name: "missing token", mutate: func(v url.Values) { v.Del("subject_token") }, code: "invalid_request"},
|
||||
{name: "wrong grant", mutate: func(v url.Values) { v.Set("grant_type", "client_credentials") }, code: "unsupported_grant_type"},
|
||||
{name: "missing audience", mutate: func(v url.Values) { v.Del("audience") }, code: "invalid_target"},
|
||||
{name: "multiple audiences", mutate: func(v url.Values) { v.Add("audience", "second") }, code: "invalid_target"},
|
||||
{name: "missing scope", mutate: func(v url.Values) { v.Del("scope") }, code: "invalid_scope"},
|
||||
{name: "actor token", mutate: func(v url.Values) { v.Set("actor_token", "not-supported") }, code: "invalid_request"},
|
||||
}
|
||||
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
form := validForm()
|
||||
test.mutate(form)
|
||||
|
||||
_, err := ParseExchangeRequest(form)
|
||||
var requestErr *RequestError
|
||||
if !errors.As(err, &requestErr) {
|
||||
t.Fatalf("error = %v, want RequestError", err)
|
||||
}
|
||||
if requestErr.Code != test.code {
|
||||
t.Fatalf("error code = %q, want %q", requestErr.Code, test.code)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func validForm() url.Values {
|
||||
return url.Values{
|
||||
"grant_type": {TokenExchangeGrantType},
|
||||
"client_id": {"homelab-infra-ci"},
|
||||
"subject_token": {"test-subject-token"},
|
||||
"subject_token_type": {JWTTokenType},
|
||||
"requested_token_type": {AccessTokenType},
|
||||
"audience": {"https://bao.ad.ddupan.top:8200"},
|
||||
"scope": {"bao.login ssh.certificate"},
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,97 @@
|
||||
package signing
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
var rawURLEncoding = base64.RawURLEncoding
|
||||
|
||||
type RS256Signer interface {
|
||||
ActiveKey(ctx context.Context) (SigningKey, error)
|
||||
SignRS256(ctx context.Context, key SigningKey, signingInput []byte) ([]byte, error)
|
||||
}
|
||||
|
||||
type SigningKey struct {
|
||||
ID string
|
||||
Version int
|
||||
}
|
||||
|
||||
type Claims struct {
|
||||
Issuer string `json:"iss"`
|
||||
Subject string `json:"sub"`
|
||||
PrincipalName string `json:"principal_name,omitempty"`
|
||||
Audience []string `json:"aud"`
|
||||
IssuedAt int64 `json:"iat"`
|
||||
NotBefore int64 `json:"nbf"`
|
||||
ExpiresAt int64 `json:"exp"`
|
||||
JWTID string `json:"jti"`
|
||||
Scope string `json:"scope"`
|
||||
}
|
||||
|
||||
type Issuer struct {
|
||||
signer RS256Signer
|
||||
now func() time.Time
|
||||
}
|
||||
|
||||
func NewIssuer(signer RS256Signer) (*Issuer, error) {
|
||||
if signer == nil {
|
||||
return nil, errors.New("signer is required")
|
||||
}
|
||||
return &Issuer{signer: signer, now: time.Now}, nil
|
||||
}
|
||||
|
||||
func (i *Issuer) Sign(ctx context.Context, claims Claims) (string, error) {
|
||||
if err := validateClaims(claims, i.now()); err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
key, err := i.signer.ActiveKey(ctx)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("select signing key: %w", err)
|
||||
}
|
||||
if key.ID == "" || key.Version < 1 {
|
||||
return "", errors.New("signer returned an invalid signing key")
|
||||
}
|
||||
|
||||
header, err := json.Marshal(map[string]string{
|
||||
"alg": "RS256",
|
||||
"kid": key.ID,
|
||||
"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)
|
||||
}
|
||||
|
||||
signingInput := rawURLEncoding.EncodeToString(header) + "." + rawURLEncoding.EncodeToString(payload)
|
||||
signature, err := i.signer.SignRS256(ctx, key, []byte(signingInput))
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("sign JWT: %w", err)
|
||||
}
|
||||
return signingInput + "." + rawURLEncoding.EncodeToString(signature), nil
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
@@ -0,0 +1,124 @@
|
||||
package signing
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto"
|
||||
"crypto/rand"
|
||||
"crypto/rsa"
|
||||
"crypto/sha256"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/go-jose/go-jose/v4"
|
||||
)
|
||||
|
||||
func TestIssuerProducesInteroperableRS256JWT(t *testing.T) {
|
||||
privateKey, err := rsa.GenerateKey(rand.Reader, 2048)
|
||||
if err != nil {
|
||||
t.Fatalf("generate key: %v", err)
|
||||
}
|
||||
signer := &localRS256Signer{keyID: "workload-sts-v1", key: privateKey}
|
||||
issuer, err := NewIssuer(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",
|
||||
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(),
|
||||
JWTID: "01993f4d-5e1a-7000-8000-000000000002",
|
||||
Scope: "bao.login",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Sign() error = %v", err)
|
||||
}
|
||||
|
||||
jws, err := jose.ParseSigned(token, []jose.SignatureAlgorithm{jose.RS256})
|
||||
if err != nil {
|
||||
t.Fatalf("go-jose rejected compact JWT: %v", err)
|
||||
}
|
||||
payload, err := jws.Verify(&privateKey.PublicKey)
|
||||
if err != nil {
|
||||
t.Fatalf("go-jose rejected signature: %v", err)
|
||||
}
|
||||
|
||||
var got Claims
|
||||
if err := json.Unmarshal(payload, &got); err != nil {
|
||||
t.Fatalf("decode claims: %v", err)
|
||||
}
|
||||
if got.Subject != "01993f4d-5e1a-7000-8000-000000000001" {
|
||||
t.Fatalf("subject = %q", got.Subject)
|
||||
}
|
||||
|
||||
parts := strings.Split(token, ".")
|
||||
headerJSON, err := base64.RawURLEncoding.DecodeString(parts[0])
|
||||
if err != nil {
|
||||
t.Fatalf("decode header: %v", err)
|
||||
}
|
||||
var header map[string]string
|
||||
if err := json.Unmarshal(headerJSON, &header); err != nil {
|
||||
t.Fatalf("unmarshal header: %v", err)
|
||||
}
|
||||
if header["typ"] != "at+jwt" || header["kid"] != "workload-sts-v1" {
|
||||
t.Fatalf("header = %#v", header)
|
||||
}
|
||||
}
|
||||
|
||||
func TestIssuerRejectsExcessiveLifetime(t *testing.T) {
|
||||
privateKey, err := rsa.GenerateKey(rand.Reader, 2048)
|
||||
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",
|
||||
})
|
||||
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)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
type localRS256Signer struct {
|
||||
keyID string
|
||||
key *rsa.PrivateKey
|
||||
}
|
||||
|
||||
func (s *localRS256Signer) ActiveKey(context.Context) (SigningKey, error) {
|
||||
return SigningKey{ID: s.keyID, Version: 1}, nil
|
||||
}
|
||||
|
||||
func (s *localRS256Signer) SignRS256(_ context.Context, _ SigningKey, input []byte) ([]byte, error) {
|
||||
digest := sha256.Sum256(input)
|
||||
return rsa.SignPKCS1v15(rand.Reader, s.key, crypto.SHA256, digest[:])
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -0,0 +1,80 @@
|
||||
package signing
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"os"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/go-jose/go-jose/v4"
|
||||
baoapi "github.com/openbao/openbao/api/v2"
|
||||
)
|
||||
|
||||
func TestTransitSignerProducesJWTVerifiedByPublishedJWKS(t *testing.T) {
|
||||
address := os.Getenv("WORKLOAD_STS_TEST_BAO_ADDR")
|
||||
token := os.Getenv("WORKLOAD_STS_TEST_BAO_TOKEN")
|
||||
if address == "" || token == "" {
|
||||
t.Skip("set WORKLOAD_STS_TEST_BAO_ADDR and WORKLOAD_STS_TEST_BAO_TOKEN")
|
||||
}
|
||||
|
||||
config := baoapi.DefaultConfig()
|
||||
config.Address = address
|
||||
client, err := baoapi.NewClient(config)
|
||||
if err != nil {
|
||||
t.Fatalf("create OpenBao client: %v", err)
|
||||
}
|
||||
client.SetToken(token)
|
||||
|
||||
signer, err := NewTransitSigner(client, "transit", "workload-sts")
|
||||
if err != nil {
|
||||
t.Fatalf("NewTransitSigner() error = %v", err)
|
||||
}
|
||||
issuer, err := NewIssuer(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",
|
||||
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(),
|
||||
JWTID: "01993f4d-5e1a-7000-8000-000000000002",
|
||||
Scope: "bao.login",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Sign() error = %v", err)
|
||||
}
|
||||
|
||||
set, err := signer.JWKS(context.Background())
|
||||
if err != nil {
|
||||
t.Fatalf("JWKS() error = %v", err)
|
||||
}
|
||||
encodedSet, err := json.Marshal(set)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal JWKS: %v", err)
|
||||
}
|
||||
var decodedSet jose.JSONWebKeySet
|
||||
if err := json.Unmarshal(encodedSet, &decodedSet); err != nil {
|
||||
t.Fatalf("decode JWKS with go-jose: %v", err)
|
||||
}
|
||||
|
||||
jws, err := jose.ParseSigned(compact, []jose.SignatureAlgorithm{jose.RS256})
|
||||
if err != nil {
|
||||
t.Fatalf("parse compact JWT: %v", err)
|
||||
}
|
||||
keyID := jws.Signatures[0].Header.KeyID
|
||||
keys := decodedSet.Key(keyID)
|
||||
if len(keys) != 1 {
|
||||
t.Fatalf("JWKS keys for %q = %d, want 1", keyID, len(keys))
|
||||
}
|
||||
if _, err := jws.Verify(keys[0].Key); err != nil {
|
||||
t.Fatalf("published JWKS rejected Transit signature: %v", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,76 @@
|
||||
package telemetry
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
|
||||
"go.opentelemetry.io/otel/metric"
|
||||
metricnoop "go.opentelemetry.io/otel/metric/noop"
|
||||
sdkmetric "go.opentelemetry.io/otel/sdk/metric"
|
||||
sdkresource "go.opentelemetry.io/otel/sdk/resource"
|
||||
sdktrace "go.opentelemetry.io/otel/sdk/trace"
|
||||
"go.opentelemetry.io/otel/trace"
|
||||
tracenoop "go.opentelemetry.io/otel/trace/noop"
|
||||
)
|
||||
|
||||
type Providers struct {
|
||||
Tracer trace.TracerProvider
|
||||
Meter metric.MeterProvider
|
||||
|
||||
shutdown func(context.Context) error
|
||||
}
|
||||
|
||||
type Config struct {
|
||||
Resource *sdkresource.Resource
|
||||
TraceExporter sdktrace.SpanExporter
|
||||
MetricReader sdkmetric.Reader
|
||||
}
|
||||
|
||||
func New(config Config) (*Providers, error) {
|
||||
if config.Resource == nil {
|
||||
return nil, errors.New("OpenTelemetry resource is required")
|
||||
}
|
||||
|
||||
providers := &Providers{
|
||||
Tracer: tracenoop.NewTracerProvider(),
|
||||
Meter: metricnoop.NewMeterProvider(),
|
||||
}
|
||||
|
||||
var tracerProvider *sdktrace.TracerProvider
|
||||
if config.TraceExporter != nil {
|
||||
tracerProvider = sdktrace.NewTracerProvider(
|
||||
sdktrace.WithBatcher(config.TraceExporter),
|
||||
sdktrace.WithResource(config.Resource),
|
||||
)
|
||||
providers.Tracer = tracerProvider
|
||||
}
|
||||
|
||||
var meterProvider *sdkmetric.MeterProvider
|
||||
if config.MetricReader != nil {
|
||||
meterProvider = sdkmetric.NewMeterProvider(
|
||||
sdkmetric.WithReader(config.MetricReader),
|
||||
sdkmetric.WithResource(config.Resource),
|
||||
)
|
||||
providers.Meter = meterProvider
|
||||
}
|
||||
|
||||
providers.shutdown = func(ctx context.Context) error {
|
||||
var shutdownErrors []error
|
||||
if meterProvider != nil {
|
||||
shutdownErrors = append(shutdownErrors, meterProvider.Shutdown(ctx))
|
||||
}
|
||||
if tracerProvider != nil {
|
||||
shutdownErrors = append(shutdownErrors, tracerProvider.Shutdown(ctx))
|
||||
}
|
||||
if err := errors.Join(shutdownErrors...); err != nil {
|
||||
return fmt.Errorf("shutdown OpenTelemetry providers: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
return providers, nil
|
||||
}
|
||||
|
||||
func (p *Providers) Shutdown(ctx context.Context) error {
|
||||
return p.shutdown(ctx)
|
||||
}
|
||||
@@ -0,0 +1,27 @@
|
||||
package telemetry
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
sdkresource "go.opentelemetry.io/otel/sdk/resource"
|
||||
)
|
||||
|
||||
func TestNewWithoutExportersUsesNoopProviders(t *testing.T) {
|
||||
providers, err := New(Config{Resource: sdkresource.Empty()})
|
||||
if err != nil {
|
||||
t.Fatalf("New() error = %v", err)
|
||||
}
|
||||
if providers.Tracer == nil || providers.Meter == nil {
|
||||
t.Fatal("New() returned nil provider")
|
||||
}
|
||||
if err := providers.Shutdown(context.Background()); err != nil {
|
||||
t.Fatalf("Shutdown() error = %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewRequiresResource(t *testing.T) {
|
||||
if _, err := New(Config{}); err == nil {
|
||||
t.Fatal("New() error = nil, want missing resource error")
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user