mirror of
https://github.com/seaweedfs/seaweedfs.git
synced 2026-10-01 12:16:07 +00:00
Compare commits
7
Commits
4.26
...
iam-phase2a
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
ef1acaa98c | ||
|
|
12688c249e | ||
|
|
641bea825d | ||
|
|
91fe0a5162 | ||
|
|
64a60607c6 | ||
|
|
d4365e2f37 | ||
|
|
f9dfc0ea37 |
@@ -7,8 +7,10 @@ import (
|
||||
"fmt"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/golang-jwt/jwt/v5"
|
||||
"github.com/seaweedfs/seaweedfs/weed/glog"
|
||||
"github.com/seaweedfs/seaweedfs/weed/iam/policy"
|
||||
"github.com/seaweedfs/seaweedfs/weed/iam/providers"
|
||||
"github.com/seaweedfs/seaweedfs/weed/iam/sts"
|
||||
@@ -27,12 +29,44 @@ type IAMManager struct {
|
||||
policyEngine *policy.PolicyEngine
|
||||
roleStore RoleStore
|
||||
userStore UserStore
|
||||
oidcProviderStore OIDCProviderStore
|
||||
filerAddressProvider func() string // Function to get current filer address
|
||||
initialized bool
|
||||
runtimePolicyMu sync.Mutex
|
||||
runtimePolicyNames map[string]struct{}
|
||||
}
|
||||
|
||||
// SetOIDCProviderStore configures the IAM-managed OIDC provider store. When
|
||||
// nil, OIDC provider IAM actions return ServiceNotReady. The store is the
|
||||
// source of truth for AssumeRoleWithWebIdentity provider resolution once
|
||||
// Phase 2b lands; in Phase 2a it is read-only and populated from static
|
||||
// configuration at boot.
|
||||
func (m *IAMManager) SetOIDCProviderStore(store OIDCProviderStore) {
|
||||
m.oidcProviderStore = store
|
||||
}
|
||||
|
||||
// GetOIDCProviderStore returns the configured store (may be nil).
|
||||
func (m *IAMManager) GetOIDCProviderStore() OIDCProviderStore {
|
||||
return m.oidcProviderStore
|
||||
}
|
||||
|
||||
// GetOIDCProvider returns the record for the given ARN, or an error if the
|
||||
// store is not configured or the record is missing.
|
||||
func (m *IAMManager) GetOIDCProvider(ctx context.Context, arn string) (*OIDCProviderRecord, error) {
|
||||
if m.oidcProviderStore == nil {
|
||||
return nil, fmt.Errorf("OIDC provider store not configured")
|
||||
}
|
||||
return m.oidcProviderStore.GetProviderByARN(ctx, m.getFilerAddress(), arn)
|
||||
}
|
||||
|
||||
// ListOIDCProviders enumerates all configured OIDC providers.
|
||||
func (m *IAMManager) ListOIDCProviders(ctx context.Context) ([]*OIDCProviderRecord, error) {
|
||||
if m.oidcProviderStore == nil {
|
||||
return nil, fmt.Errorf("OIDC provider store not configured")
|
||||
}
|
||||
return m.oidcProviderStore.ListProviders(ctx, m.getFilerAddress())
|
||||
}
|
||||
|
||||
// IAMConfig holds configuration for all IAM components
|
||||
type IAMConfig struct {
|
||||
// STS service configuration
|
||||
@@ -43,6 +77,17 @@ type IAMConfig struct {
|
||||
|
||||
// Role store configuration
|
||||
Roles *RoleStoreConfig `json:"roleStore"`
|
||||
|
||||
// OIDCProviders configures the IAM-managed OIDC provider store. Optional;
|
||||
// if absent the manager defaults to an in-memory store hydrated from
|
||||
// STS.Providers at boot.
|
||||
OIDCProviders *OIDCProviderStoreConfig `json:"oidcProviderStore,omitempty"`
|
||||
}
|
||||
|
||||
// OIDCProviderStoreConfig holds OIDC provider store configuration.
|
||||
type OIDCProviderStoreConfig struct {
|
||||
StoreType string `json:"storeType"` // memory, filer
|
||||
StoreConfig map[string]interface{} `json:"storeConfig,omitempty"`
|
||||
}
|
||||
|
||||
// RoleStoreConfig holds role store configuration
|
||||
@@ -75,6 +120,12 @@ type RoleDefinition struct {
|
||||
|
||||
// Description is an optional description of the role
|
||||
Description string `json:"description,omitempty"`
|
||||
|
||||
// MaxSessionDuration is the upper bound (in seconds) on session length when
|
||||
// callers assume this role. Zero means "use the global STS default". When
|
||||
// set it must satisfy AWS bounds: 3600 ≤ MaxSessionDuration ≤ 43200.
|
||||
// Honoured by AssumeRole, AssumeRoleWithWebIdentity, AssumeRoleWithCredentials.
|
||||
MaxSessionDuration int64 `json:"maxSessionDuration,omitempty"`
|
||||
}
|
||||
|
||||
// ActionRequest represents a request to perform an action
|
||||
@@ -190,10 +241,110 @@ func (m *IAMManager) Initialize(config *IAMConfig, filerAddressProvider func() s
|
||||
}
|
||||
m.roleStore = roleStore
|
||||
|
||||
// Initialize OIDC provider store and hydrate from static configuration so
|
||||
// the read-only IAM API can return the same providers the STS service
|
||||
// already accepts. Mutations will land in Phase 2b.
|
||||
if err := m.initOIDCProviderStore(config); err != nil {
|
||||
return fmt.Errorf("failed to initialize OIDC provider store: %w", err)
|
||||
}
|
||||
|
||||
m.initialized = true
|
||||
return nil
|
||||
}
|
||||
|
||||
// initOIDCProviderStore creates the OIDC provider store and seeds it from the
|
||||
// static STS provider configuration. The static path remains the bootstrap
|
||||
// source: each enabled OIDC entry under STS.Providers is mirrored as an
|
||||
// OIDCProviderRecord so the IAM API surfaces the same set the STS service
|
||||
// validates against.
|
||||
func (m *IAMManager) initOIDCProviderStore(config *IAMConfig) error {
|
||||
store, err := m.createOIDCProviderStore(config.OIDCProviders)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
m.oidcProviderStore = store
|
||||
|
||||
if config.STS == nil {
|
||||
return nil
|
||||
}
|
||||
for _, pc := range config.STS.Providers {
|
||||
if pc == nil || !pc.Enabled || pc.Type != sts.ProviderTypeOIDC {
|
||||
continue
|
||||
}
|
||||
issuer, _ := pc.Config["issuer"].(string)
|
||||
if issuer == "" {
|
||||
glog.Warningf("OIDC provider %s in static config has empty issuer; skipping mirror to store", pc.Name)
|
||||
continue
|
||||
}
|
||||
accountID := ""
|
||||
if config.STS != nil {
|
||||
accountID = config.STS.AccountId
|
||||
}
|
||||
arn, err := DeriveOIDCProviderARN(accountID, issuer)
|
||||
if err != nil {
|
||||
glog.Warningf("derive ARN for static OIDC provider %s: %v", pc.Name, err)
|
||||
continue
|
||||
}
|
||||
clientIDs := extractClientIDs(pc.Config)
|
||||
ctx := context.Background()
|
||||
// Preserve CreatedAt across reboots when a persistent store already
|
||||
// has this provider — IAM's GetOpenIDConnectProvider response
|
||||
// shouldn't shift its CreateDate every time the server restarts.
|
||||
now := time.Now().UTC()
|
||||
createdAt := now
|
||||
if existing, err := store.GetProviderByARN(ctx, m.getFilerAddress(), arn); err == nil && existing != nil && !existing.CreatedAt.IsZero() {
|
||||
createdAt = existing.CreatedAt
|
||||
}
|
||||
rec := &OIDCProviderRecord{
|
||||
AccountID: accountID,
|
||||
ARN: arn,
|
||||
URL: issuer,
|
||||
ClientIDs: clientIDs,
|
||||
CreatedAt: createdAt,
|
||||
UpdatedAt: now,
|
||||
}
|
||||
if err := store.StoreProvider(ctx, m.getFilerAddress(), rec); err != nil {
|
||||
glog.Warningf("mirror static OIDC provider %s into store: %v", pc.Name, err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// createOIDCProviderStore selects an OIDCProviderStore implementation. Defaults
|
||||
// to memory; "filer" requires a filerAddressProvider to be configured.
|
||||
func (m *IAMManager) createOIDCProviderStore(cfg *OIDCProviderStoreConfig) (OIDCProviderStore, error) {
|
||||
if cfg == nil || cfg.StoreType == "" || cfg.StoreType == "memory" {
|
||||
return NewMemoryOIDCProviderStore(), nil
|
||||
}
|
||||
if cfg.StoreType == "filer" {
|
||||
return NewFilerOIDCProviderStore(cfg.StoreConfig, m.filerAddressProvider), nil
|
||||
}
|
||||
return nil, fmt.Errorf("unsupported OIDC provider store type: %s", cfg.StoreType)
|
||||
}
|
||||
|
||||
// extractClientIDs reads a single clientId or a clientIds list from the
|
||||
// provider's static config map. Mirrors the OIDCConfig schema.
|
||||
func extractClientIDs(cfg map[string]interface{}) []string {
|
||||
if cfg == nil {
|
||||
return nil
|
||||
}
|
||||
if list, ok := cfg["clientIds"].([]interface{}); ok {
|
||||
out := make([]string, 0, len(list))
|
||||
for _, v := range list {
|
||||
if s, ok := v.(string); ok && s != "" {
|
||||
out = append(out, s)
|
||||
}
|
||||
}
|
||||
if len(out) > 0 {
|
||||
return out
|
||||
}
|
||||
}
|
||||
if id, ok := cfg["clientId"].(string); ok && id != "" {
|
||||
return []string{id}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// getFilerAddress returns the current filer address using the provider function
|
||||
func (m *IAMManager) getFilerAddress() string {
|
||||
if m.filerAddressProvider != nil {
|
||||
@@ -272,6 +423,13 @@ func (m *IAMManager) CreateRole(ctx context.Context, filerAddress string, roleNa
|
||||
}
|
||||
}
|
||||
|
||||
// Validate per-role MaxSessionDuration if specified. AWS bounds: 1h..12h.
|
||||
if roleDef.MaxSessionDuration != 0 {
|
||||
if roleDef.MaxSessionDuration < 3600 || roleDef.MaxSessionDuration > 43200 {
|
||||
return fmt.Errorf("MaxSessionDuration must be between 3600 and 43200 seconds, got %d", roleDef.MaxSessionDuration)
|
||||
}
|
||||
}
|
||||
|
||||
// Store role definition
|
||||
return m.roleStore.StoreRole(ctx, "", roleName, roleDef)
|
||||
}
|
||||
@@ -329,10 +487,33 @@ func (m *IAMManager) AssumeRoleWithWebIdentity(ctx context.Context, request *sts
|
||||
return nil, fmt.Errorf("trust policy validation failed: %w", err)
|
||||
}
|
||||
|
||||
// Apply role-level MaxSessionDuration cap. The STS service still applies
|
||||
// the global MaxSessionLength and the source-token-expiry cap on top of
|
||||
// this; per-role takes precedence whenever it is the tightest bound.
|
||||
request.DurationSeconds = capDurationByRole(request.DurationSeconds, roleDef.MaxSessionDuration)
|
||||
|
||||
// Use STS service to assume the role
|
||||
return m.stsService.AssumeRoleWithWebIdentity(ctx, request)
|
||||
}
|
||||
|
||||
// capDurationByRole returns the requested duration clamped to the role's
|
||||
// MaxSessionDuration. A nil requested duration is left nil so the STS
|
||||
// service's calculateSessionDuration applies the global default (typically
|
||||
// 1 hour) — substituting the role's max here would silently mint a 12h
|
||||
// session for any caller who omitted DurationSeconds, which AWS does not
|
||||
// do. The role-max upper bound still applies in the downstream cap chain
|
||||
// once the request has a concrete duration.
|
||||
func capDurationByRole(requested *int64, roleMax int64) *int64 {
|
||||
if roleMax <= 0 || requested == nil {
|
||||
return requested
|
||||
}
|
||||
if *requested > roleMax {
|
||||
v := roleMax
|
||||
return &v
|
||||
}
|
||||
return requested
|
||||
}
|
||||
|
||||
// AssumeRoleWithCredentials assumes a role using credentials (LDAP)
|
||||
func (m *IAMManager) AssumeRoleWithCredentials(ctx context.Context, request *sts.AssumeRoleWithCredentialsRequest) (*sts.AssumeRoleResponse, error) {
|
||||
if !m.initialized {
|
||||
@@ -353,6 +534,9 @@ func (m *IAMManager) AssumeRoleWithCredentials(ctx context.Context, request *sts
|
||||
return nil, fmt.Errorf("trust policy validation failed: %w", err)
|
||||
}
|
||||
|
||||
// Apply role-level MaxSessionDuration cap.
|
||||
request.DurationSeconds = capDurationByRole(request.DurationSeconds, roleDef.MaxSessionDuration)
|
||||
|
||||
// Use STS service to assume the role
|
||||
return m.stsService.AssumeRoleWithCredentials(ctx, request)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,113 @@
|
||||
package integration
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/seaweedfs/seaweedfs/weed/iam/policy"
|
||||
"github.com/seaweedfs/seaweedfs/weed/iam/sts"
|
||||
)
|
||||
|
||||
func TestStaticConfigSeedsProviderStore(t *testing.T) {
|
||||
mgr := NewIAMManager()
|
||||
cfg := &IAMConfig{
|
||||
STS: &sts.STSConfig{
|
||||
TokenDuration: sts.FlexibleDuration{Duration: time.Hour},
|
||||
MaxSessionLength: sts.FlexibleDuration{Duration: 12 * time.Hour},
|
||||
Issuer: "test-sts",
|
||||
SigningKey: []byte("test-signing-key-32-characters-long"),
|
||||
AccountId: "111122223333",
|
||||
Providers: []*sts.ProviderConfig{
|
||||
{
|
||||
Name: "google",
|
||||
Type: sts.ProviderTypeOIDC,
|
||||
Enabled: true,
|
||||
Config: map[string]interface{}{
|
||||
"issuer": "https://accounts.google.com",
|
||||
"clientId": "1234.apps.googleusercontent.com",
|
||||
},
|
||||
},
|
||||
{
|
||||
Name: "github-actions",
|
||||
Type: sts.ProviderTypeOIDC,
|
||||
Enabled: true,
|
||||
Config: map[string]interface{}{
|
||||
"issuer": "https://token.actions.githubusercontent.com",
|
||||
"clientId": "sts.amazonaws.com",
|
||||
},
|
||||
},
|
||||
{
|
||||
Name: "disabled-google",
|
||||
Type: sts.ProviderTypeOIDC,
|
||||
Enabled: false,
|
||||
Config: map[string]interface{}{
|
||||
"issuer": "https://accounts.disabled.example",
|
||||
"clientId": "x",
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
Policy: &policy.PolicyEngineConfig{DefaultEffect: "Deny", StoreType: "memory"},
|
||||
Roles: &RoleStoreConfig{StoreType: "memory"},
|
||||
}
|
||||
if err := mgr.Initialize(cfg, func() string { return "localhost:8888" }); err != nil {
|
||||
t.Fatalf("Initialize: %v", err)
|
||||
}
|
||||
|
||||
store := mgr.GetOIDCProviderStore()
|
||||
if store == nil {
|
||||
t.Fatal("expected store to be initialized")
|
||||
}
|
||||
|
||||
got, err := mgr.ListOIDCProviders(context.Background())
|
||||
if err != nil {
|
||||
t.Fatalf("ListOIDCProviders: %v", err)
|
||||
}
|
||||
// Disabled providers must not be mirrored.
|
||||
if len(got) != 2 {
|
||||
t.Fatalf("expected 2 providers, got %d (%v)", len(got), got)
|
||||
}
|
||||
|
||||
byArn := map[string]*OIDCProviderRecord{}
|
||||
for _, r := range got {
|
||||
byArn[r.ARN] = r
|
||||
}
|
||||
|
||||
googleARN := "arn:aws:iam::111122223333:oidc-provider/accounts.google.com"
|
||||
g, ok := byArn[googleARN]
|
||||
if !ok {
|
||||
t.Fatalf("missing google ARN; have %v", byArn)
|
||||
}
|
||||
if len(g.ClientIDs) != 1 || g.ClientIDs[0] != "1234.apps.googleusercontent.com" {
|
||||
t.Fatalf("google clientIds wrong: %v", g.ClientIDs)
|
||||
}
|
||||
|
||||
ghARN := "arn:aws:iam::111122223333:oidc-provider/token.actions.githubusercontent.com"
|
||||
gh, ok := byArn[ghARN]
|
||||
if !ok {
|
||||
t.Fatalf("missing github-actions ARN; have %v", byArn)
|
||||
}
|
||||
if len(gh.ClientIDs) != 1 || gh.ClientIDs[0] != "sts.amazonaws.com" {
|
||||
t.Fatalf("github clientIds wrong: %v", gh.ClientIDs)
|
||||
}
|
||||
|
||||
// Spot-check the IAM read path returns the same record.
|
||||
rec, err := mgr.GetOIDCProvider(context.Background(), googleARN)
|
||||
if err != nil {
|
||||
t.Fatalf("GetOIDCProvider: %v", err)
|
||||
}
|
||||
if rec.URL != "https://accounts.google.com" {
|
||||
t.Fatalf("URL mismatch: %s", rec.URL)
|
||||
}
|
||||
}
|
||||
|
||||
func TestStoreNotConfiguredReturnsClearError(t *testing.T) {
|
||||
mgr := NewIAMManager()
|
||||
if _, err := mgr.GetOIDCProvider(context.Background(), "arn:..."); err == nil {
|
||||
t.Fatal("expected error when store not configured")
|
||||
}
|
||||
if _, err := mgr.ListOIDCProviders(context.Background()); err == nil {
|
||||
t.Fatal("expected error when store not configured")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,421 @@
|
||||
package integration
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha1"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/url"
|
||||
"sort"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/seaweedfs/seaweedfs/weed/glog"
|
||||
"github.com/seaweedfs/seaweedfs/weed/pb"
|
||||
"github.com/seaweedfs/seaweedfs/weed/pb/filer_pb"
|
||||
"google.golang.org/grpc"
|
||||
)
|
||||
|
||||
// OIDCProviderRecord is the persisted, IAM-managed view of an OIDC identity
|
||||
// provider. It is the source of truth consulted at AssumeRoleWithWebIdentity
|
||||
// time. Static configuration entries are loaded into this same store at
|
||||
// startup so the resolution path is uniform.
|
||||
type OIDCProviderRecord struct {
|
||||
// AccountID scopes the record. Empty means "global" — usable from any
|
||||
// account. Non-empty records are only resolvable when the assuming
|
||||
// caller's RoleArn lives in the same account. Phase 2a leaves this
|
||||
// field unenforced; Phase 3c lights up the cross-account check.
|
||||
AccountID string `json:"accountId,omitempty"`
|
||||
|
||||
// ARN is the canonical IAM identifier for the provider:
|
||||
// arn:aws:iam::<account>:oidc-provider/<host>/<path>.
|
||||
ARN string `json:"arn"`
|
||||
|
||||
// URL is the issuer URL (no trailing slash), e.g.
|
||||
// https://token.actions.githubusercontent.com.
|
||||
URL string `json:"url"`
|
||||
|
||||
// ClientIDs are the audiences AWS calls "client IDs". A token's `aud` or
|
||||
// `azp` must match one of these for the token to be accepted.
|
||||
ClientIDs []string `json:"clientIds,omitempty"`
|
||||
|
||||
// Thumbprints are SHA-1 hex digests of the IDP's TLS certificate, matching
|
||||
// the AWS-compatible thumbprint algorithm. Used at JWKS-fetch time when
|
||||
// non-empty; an empty list means "trust the system root store".
|
||||
Thumbprints []string `json:"thumbprints,omitempty"`
|
||||
|
||||
// AllowedPrincipalTagKeys is the per-provider ABAC allowlist for the
|
||||
// `https://aws.amazon.com/tags/principal_tags` claim namespace. Empty
|
||||
// means no tags are surfaced from this provider.
|
||||
AllowedPrincipalTagKeys []string `json:"allowedPrincipalTagKeys,omitempty"`
|
||||
|
||||
// PolicyClaim, when non-empty, enables claim-based policy mode for this
|
||||
// provider: tokens whose RoleArn matches the configured sentinel will
|
||||
// derive their effective policies from this JWT claim.
|
||||
PolicyClaim string `json:"policyClaim,omitempty"`
|
||||
|
||||
// Tags is the AWS-style tag set attached to the IAM resource itself
|
||||
// (audit/inventory metadata, not propagated into sessions).
|
||||
Tags map[string]string `json:"tags,omitempty"`
|
||||
|
||||
CreatedAt time.Time `json:"createdAt"`
|
||||
UpdatedAt time.Time `json:"updatedAt"`
|
||||
}
|
||||
|
||||
// OIDCProviderStore stores OIDCProviderRecord entries. Implementations are
|
||||
// expected to be safe for concurrent use.
|
||||
type OIDCProviderStore interface {
|
||||
StoreProvider(ctx context.Context, filerAddress string, record *OIDCProviderRecord) error
|
||||
GetProviderByARN(ctx context.Context, filerAddress string, arn string) (*OIDCProviderRecord, error)
|
||||
GetProviderByIssuer(ctx context.Context, filerAddress string, issuer string) (*OIDCProviderRecord, error)
|
||||
ListProviders(ctx context.Context, filerAddress string) ([]*OIDCProviderRecord, error)
|
||||
DeleteProvider(ctx context.Context, filerAddress string, arn string) error
|
||||
}
|
||||
|
||||
// MemoryOIDCProviderStore is a process-local store, suitable for tests and
|
||||
// single-node deployments. It also acts as the in-memory cache hydrated from
|
||||
// static config at boot.
|
||||
type MemoryOIDCProviderStore struct {
|
||||
mu sync.RWMutex
|
||||
providers map[string]*OIDCProviderRecord // keyed by ARN
|
||||
}
|
||||
|
||||
// NewMemoryOIDCProviderStore creates an empty in-memory store.
|
||||
func NewMemoryOIDCProviderStore() *MemoryOIDCProviderStore {
|
||||
return &MemoryOIDCProviderStore{providers: make(map[string]*OIDCProviderRecord)}
|
||||
}
|
||||
|
||||
// StoreProvider replaces any existing record with the same ARN.
|
||||
func (m *MemoryOIDCProviderStore) StoreProvider(ctx context.Context, _ string, record *OIDCProviderRecord) error {
|
||||
if record == nil {
|
||||
return fmt.Errorf("record cannot be nil")
|
||||
}
|
||||
if record.ARN == "" {
|
||||
return fmt.Errorf("record.ARN is required")
|
||||
}
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
m.providers[record.ARN] = copyOIDCProviderRecord(record)
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetProviderByARN returns the record with the given ARN, or an error if absent.
|
||||
func (m *MemoryOIDCProviderStore) GetProviderByARN(ctx context.Context, _ string, arn string) (*OIDCProviderRecord, error) {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
rec, ok := m.providers[arn]
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("OIDC provider not found: %s", arn)
|
||||
}
|
||||
return copyOIDCProviderRecord(rec), nil
|
||||
}
|
||||
|
||||
// GetProviderByIssuer scans for the first record whose URL matches `issuer`.
|
||||
// Comparison strips trailing slashes and is case-insensitive on host.
|
||||
func (m *MemoryOIDCProviderStore) GetProviderByIssuer(ctx context.Context, _ string, issuer string) (*OIDCProviderRecord, error) {
|
||||
want := normalizeIssuer(issuer)
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
for _, rec := range m.providers {
|
||||
if normalizeIssuer(rec.URL) == want {
|
||||
return copyOIDCProviderRecord(rec), nil
|
||||
}
|
||||
}
|
||||
return nil, fmt.Errorf("no OIDC provider registered for issuer: %s", issuer)
|
||||
}
|
||||
|
||||
// ListProviders returns every record sorted by ARN for stable output.
|
||||
func (m *MemoryOIDCProviderStore) ListProviders(ctx context.Context, _ string) ([]*OIDCProviderRecord, error) {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
out := make([]*OIDCProviderRecord, 0, len(m.providers))
|
||||
for _, rec := range m.providers {
|
||||
out = append(out, copyOIDCProviderRecord(rec))
|
||||
}
|
||||
sort.Slice(out, func(i, j int) bool { return out[i].ARN < out[j].ARN })
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// DeleteProvider removes the record. Idempotent — deleting a missing ARN is success.
|
||||
func (m *MemoryOIDCProviderStore) DeleteProvider(ctx context.Context, _ string, arn string) error {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
delete(m.providers, arn)
|
||||
return nil
|
||||
}
|
||||
|
||||
// FilerOIDCProviderStore persists records as JSON files in a filer directory,
|
||||
// mirroring FilerRoleStore.
|
||||
type FilerOIDCProviderStore struct {
|
||||
grpcDialOption grpc.DialOption
|
||||
basePath string
|
||||
filerAddressProvider func() string
|
||||
}
|
||||
|
||||
// NewFilerOIDCProviderStore returns a filer-backed store. Default path
|
||||
// `/etc/iam/oidc-providers` aligns with the existing roles directory.
|
||||
func NewFilerOIDCProviderStore(config map[string]interface{}, filerAddressProvider func() string) *FilerOIDCProviderStore {
|
||||
store := &FilerOIDCProviderStore{
|
||||
basePath: "/etc/iam/oidc-providers",
|
||||
filerAddressProvider: filerAddressProvider,
|
||||
}
|
||||
if config != nil {
|
||||
if bp, ok := config["basePath"].(string); ok && bp != "" {
|
||||
store.basePath = strings.TrimSuffix(bp, "/")
|
||||
}
|
||||
}
|
||||
glog.V(2).Infof("Initialized FilerOIDCProviderStore with basePath %s", store.basePath)
|
||||
return store
|
||||
}
|
||||
|
||||
func (f *FilerOIDCProviderStore) resolveFilerAddress(filerAddress string) string {
|
||||
if filerAddress != "" {
|
||||
return filerAddress
|
||||
}
|
||||
if f.filerAddressProvider != nil {
|
||||
return f.filerAddressProvider()
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (f *FilerOIDCProviderStore) fileName(arn string) string {
|
||||
// Hash the ARN to a fixed-width filename so we can store any character set
|
||||
// (including the URL fragments AWS lets through) without filer-path issues.
|
||||
h := sha1.Sum([]byte(arn))
|
||||
return hex.EncodeToString(h[:]) + ".json"
|
||||
}
|
||||
|
||||
// StoreProvider persists `record` as JSON.
|
||||
func (f *FilerOIDCProviderStore) StoreProvider(ctx context.Context, filerAddress string, record *OIDCProviderRecord) error {
|
||||
filerAddress = f.resolveFilerAddress(filerAddress)
|
||||
if filerAddress == "" {
|
||||
return fmt.Errorf("filer address is required")
|
||||
}
|
||||
if record == nil || record.ARN == "" {
|
||||
return fmt.Errorf("record with ARN is required")
|
||||
}
|
||||
|
||||
data, err := json.MarshalIndent(record, "", " ")
|
||||
if err != nil {
|
||||
return fmt.Errorf("marshal OIDC provider: %v", err)
|
||||
}
|
||||
|
||||
return f.withFilerClient(filerAddress, func(client filer_pb.SeaweedFilerClient) error {
|
||||
_, err := client.CreateEntry(ctx, &filer_pb.CreateEntryRequest{
|
||||
Directory: f.basePath,
|
||||
Entry: &filer_pb.Entry{
|
||||
Name: f.fileName(record.ARN),
|
||||
IsDirectory: false,
|
||||
Attributes: &filer_pb.FuseAttributes{
|
||||
Mtime: time.Now().Unix(),
|
||||
Crtime: time.Now().Unix(),
|
||||
FileMode: uint32(0o600),
|
||||
},
|
||||
Content: data,
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("store OIDC provider %s: %v", record.ARN, err)
|
||||
}
|
||||
return nil
|
||||
})
|
||||
}
|
||||
|
||||
// GetProviderByARN looks up the record file by its ARN-derived filename.
|
||||
func (f *FilerOIDCProviderStore) GetProviderByARN(ctx context.Context, filerAddress string, arn string) (*OIDCProviderRecord, error) {
|
||||
filerAddress = f.resolveFilerAddress(filerAddress)
|
||||
if filerAddress == "" {
|
||||
return nil, fmt.Errorf("filer address is required")
|
||||
}
|
||||
|
||||
var data []byte
|
||||
err := f.withFilerClient(filerAddress, func(client filer_pb.SeaweedFilerClient) error {
|
||||
resp, err := client.LookupDirectoryEntry(ctx, &filer_pb.LookupDirectoryEntryRequest{
|
||||
Directory: f.basePath,
|
||||
Name: f.fileName(arn),
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("OIDC provider not found: %v", err)
|
||||
}
|
||||
if resp.Entry == nil {
|
||||
return fmt.Errorf("OIDC provider not found: %s", arn)
|
||||
}
|
||||
data = resp.Entry.Content
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var rec OIDCProviderRecord
|
||||
if err := json.Unmarshal(data, &rec); err != nil {
|
||||
return nil, fmt.Errorf("unmarshal OIDC provider: %v", err)
|
||||
}
|
||||
return &rec, nil
|
||||
}
|
||||
|
||||
// GetProviderByIssuer linearly scans the directory. Phase 2b adds an issuer
|
||||
// index when write throughput grows; for read-only this is good enough.
|
||||
func (f *FilerOIDCProviderStore) GetProviderByIssuer(ctx context.Context, filerAddress string, issuer string) (*OIDCProviderRecord, error) {
|
||||
records, err := f.ListProviders(ctx, filerAddress)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
want := normalizeIssuer(issuer)
|
||||
for _, rec := range records {
|
||||
if normalizeIssuer(rec.URL) == want {
|
||||
return rec, nil
|
||||
}
|
||||
}
|
||||
return nil, fmt.Errorf("no OIDC provider registered for issuer: %s", issuer)
|
||||
}
|
||||
|
||||
// ListProviders enumerates every record in the directory.
|
||||
func (f *FilerOIDCProviderStore) ListProviders(ctx context.Context, filerAddress string) ([]*OIDCProviderRecord, error) {
|
||||
filerAddress = f.resolveFilerAddress(filerAddress)
|
||||
if filerAddress == "" {
|
||||
return nil, fmt.Errorf("filer address is required")
|
||||
}
|
||||
|
||||
var out []*OIDCProviderRecord
|
||||
err := f.withFilerClient(filerAddress, func(client filer_pb.SeaweedFilerClient) error {
|
||||
// Stream-paginate via StartFromFileName so deployments with more
|
||||
// than 1000 providers don't get a silently truncated list.
|
||||
const pageSize = 1000
|
||||
startFrom := ""
|
||||
for {
|
||||
stream, err := client.ListEntries(ctx, &filer_pb.ListEntriesRequest{
|
||||
Directory: f.basePath,
|
||||
Limit: pageSize,
|
||||
StartFromFileName: startFrom,
|
||||
InclusiveStartFrom: false,
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("list OIDC providers: %v", err)
|
||||
}
|
||||
lastName := ""
|
||||
pageCount := 0
|
||||
for {
|
||||
resp, recvErr := stream.Recv()
|
||||
if recvErr != nil {
|
||||
if errors.Is(recvErr, io.EOF) {
|
||||
break
|
||||
}
|
||||
return fmt.Errorf("recv OIDC provider entry: %w", recvErr)
|
||||
}
|
||||
if resp.Entry == nil || resp.Entry.IsDirectory {
|
||||
continue
|
||||
}
|
||||
lastName = resp.Entry.Name
|
||||
pageCount++
|
||||
if !strings.HasSuffix(resp.Entry.Name, ".json") {
|
||||
continue
|
||||
}
|
||||
var rec OIDCProviderRecord
|
||||
if err := json.Unmarshal(resp.Entry.Content, &rec); err != nil {
|
||||
glog.Warningf("skipping malformed OIDC provider record %s: %v", resp.Entry.Name, err)
|
||||
continue
|
||||
}
|
||||
out = append(out, &rec)
|
||||
}
|
||||
if pageCount < pageSize {
|
||||
break
|
||||
}
|
||||
startFrom = lastName
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
sort.Slice(out, func(i, j int) bool { return out[i].ARN < out[j].ARN })
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// DeleteProvider removes the record. Missing ARN is treated as success.
|
||||
func (f *FilerOIDCProviderStore) DeleteProvider(ctx context.Context, filerAddress string, arn string) error {
|
||||
filerAddress = f.resolveFilerAddress(filerAddress)
|
||||
if filerAddress == "" {
|
||||
return fmt.Errorf("filer address is required")
|
||||
}
|
||||
return f.withFilerClient(filerAddress, func(client filer_pb.SeaweedFilerClient) error {
|
||||
resp, err := client.DeleteEntry(ctx, &filer_pb.DeleteEntryRequest{
|
||||
Directory: f.basePath,
|
||||
Name: f.fileName(arn),
|
||||
IsDeleteData: true,
|
||||
})
|
||||
if err != nil {
|
||||
if strings.Contains(err.Error(), "not found") {
|
||||
return nil
|
||||
}
|
||||
return fmt.Errorf("delete OIDC provider %s: %v", arn, err)
|
||||
}
|
||||
if resp.Error != "" && !strings.Contains(resp.Error, "not found") {
|
||||
return fmt.Errorf("delete OIDC provider %s: %s", arn, resp.Error)
|
||||
}
|
||||
return nil
|
||||
})
|
||||
}
|
||||
|
||||
func (f *FilerOIDCProviderStore) withFilerClient(filerAddress string, fn func(filer_pb.SeaweedFilerClient) error) error {
|
||||
return pb.WithGrpcFilerClient(false, 0, pb.ServerAddress(filerAddress), f.grpcDialOption, fn)
|
||||
}
|
||||
|
||||
// DeriveOIDCProviderARN turns an issuer URL into the canonical IAM ARN AWS uses.
|
||||
// Empty `accountID` produces a global-style ARN.
|
||||
//
|
||||
// Examples:
|
||||
//
|
||||
// https://accounts.google.com -> arn:aws:iam::<acct>:oidc-provider/accounts.google.com
|
||||
// https://oidc.eks.us-west-2.amazonaws.com/id/EXAMPLED ->
|
||||
// arn:aws:iam::<acct>:oidc-provider/oidc.eks.us-west-2.amazonaws.com/id/EXAMPLED
|
||||
func DeriveOIDCProviderARN(accountID, issuerURL string) (string, error) {
|
||||
if issuerURL == "" {
|
||||
return "", fmt.Errorf("issuer URL is required")
|
||||
}
|
||||
u, err := url.Parse(issuerURL)
|
||||
if err != nil || u.Host == "" {
|
||||
return "", fmt.Errorf("invalid issuer URL: %s", issuerURL)
|
||||
}
|
||||
host := strings.ToLower(u.Host)
|
||||
resource := host + strings.TrimSuffix(u.Path, "/")
|
||||
return fmt.Sprintf("arn:aws:iam::%s:oidc-provider/%s", accountID, resource), nil
|
||||
}
|
||||
|
||||
// normalizeIssuer compares-friendly form: lowercased host, no trailing slash.
|
||||
func normalizeIssuer(issuer string) string {
|
||||
u, err := url.Parse(issuer)
|
||||
if err != nil || u.Host == "" {
|
||||
return strings.TrimSuffix(strings.ToLower(issuer), "/")
|
||||
}
|
||||
u.Host = strings.ToLower(u.Host)
|
||||
u.Path = strings.TrimSuffix(u.Path, "/")
|
||||
return u.String()
|
||||
}
|
||||
|
||||
func copyOIDCProviderRecord(rec *OIDCProviderRecord) *OIDCProviderRecord {
|
||||
if rec == nil {
|
||||
return nil
|
||||
}
|
||||
cp := *rec
|
||||
if rec.ClientIDs != nil {
|
||||
cp.ClientIDs = append([]string(nil), rec.ClientIDs...)
|
||||
}
|
||||
if rec.Thumbprints != nil {
|
||||
cp.Thumbprints = append([]string(nil), rec.Thumbprints...)
|
||||
}
|
||||
if rec.AllowedPrincipalTagKeys != nil {
|
||||
cp.AllowedPrincipalTagKeys = append([]string(nil), rec.AllowedPrincipalTagKeys...)
|
||||
}
|
||||
if rec.Tags != nil {
|
||||
cp.Tags = make(map[string]string, len(rec.Tags))
|
||||
for k, v := range rec.Tags {
|
||||
cp.Tags[k] = v
|
||||
}
|
||||
}
|
||||
return &cp
|
||||
}
|
||||
@@ -0,0 +1,199 @@
|
||||
package integration
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func newRecord(arn, url string) *OIDCProviderRecord {
|
||||
now := time.Now()
|
||||
return &OIDCProviderRecord{
|
||||
ARN: arn,
|
||||
URL: url,
|
||||
ClientIDs: []string{"sts.amazonaws.com"},
|
||||
CreatedAt: now,
|
||||
UpdatedAt: now,
|
||||
}
|
||||
}
|
||||
|
||||
func TestMemoryStoreCRUD(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
store := NewMemoryOIDCProviderStore()
|
||||
|
||||
rec := newRecord(
|
||||
"arn:aws:iam::123:oidc-provider/token.actions.githubusercontent.com",
|
||||
"https://token.actions.githubusercontent.com",
|
||||
)
|
||||
|
||||
// Store + Get round-trip preserves the record.
|
||||
if err := store.StoreProvider(ctx, "", rec); err != nil {
|
||||
t.Fatalf("StoreProvider: %v", err)
|
||||
}
|
||||
got, err := store.GetProviderByARN(ctx, "", rec.ARN)
|
||||
if err != nil {
|
||||
t.Fatalf("GetProviderByARN: %v", err)
|
||||
}
|
||||
if got.URL != rec.URL {
|
||||
t.Fatalf("URL mismatch: got=%s want=%s", got.URL, rec.URL)
|
||||
}
|
||||
|
||||
// Mutate the returned copy and verify the store wasn't affected.
|
||||
got.ClientIDs[0] = "tampered"
|
||||
again, _ := store.GetProviderByARN(ctx, "", rec.ARN)
|
||||
if again.ClientIDs[0] == "tampered" {
|
||||
t.Fatal("store handed out a shared slice; mutations leaked back")
|
||||
}
|
||||
|
||||
// List returns the entry.
|
||||
all, err := store.ListProviders(ctx, "")
|
||||
if err != nil {
|
||||
t.Fatalf("ListProviders: %v", err)
|
||||
}
|
||||
if len(all) != 1 || all[0].ARN != rec.ARN {
|
||||
t.Fatalf("ListProviders unexpected result: %+v", all)
|
||||
}
|
||||
|
||||
// Delete -> Get returns not found.
|
||||
if err := store.DeleteProvider(ctx, "", rec.ARN); err != nil {
|
||||
t.Fatalf("DeleteProvider: %v", err)
|
||||
}
|
||||
if _, err := store.GetProviderByARN(ctx, "", rec.ARN); err == nil {
|
||||
t.Fatal("expected not-found after delete")
|
||||
}
|
||||
// Delete is idempotent.
|
||||
if err := store.DeleteProvider(ctx, "", rec.ARN); err != nil {
|
||||
t.Fatalf("idempotent delete should succeed: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMemoryStoreGetByIssuerNormalizesHost(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
store := NewMemoryOIDCProviderStore()
|
||||
|
||||
rec := newRecord(
|
||||
"arn:aws:iam::123:oidc-provider/token.actions.githubusercontent.com",
|
||||
"https://Token.Actions.GithubUserContent.com/", // mixed case + trailing slash
|
||||
)
|
||||
if err := store.StoreProvider(ctx, "", rec); err != nil {
|
||||
t.Fatalf("StoreProvider: %v", err)
|
||||
}
|
||||
|
||||
cases := []string{
|
||||
"https://token.actions.githubusercontent.com",
|
||||
"https://token.actions.githubusercontent.com/",
|
||||
"https://TOKEN.actions.GITHUBUSERCONTENT.com",
|
||||
}
|
||||
for _, want := range cases {
|
||||
got, err := store.GetProviderByIssuer(ctx, "", want)
|
||||
if err != nil {
|
||||
t.Errorf("issuer %q: GetProviderByIssuer: %v", want, err)
|
||||
continue
|
||||
}
|
||||
if got.ARN != rec.ARN {
|
||||
t.Errorf("issuer %q: ARN mismatch: got=%s", want, got.ARN)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestMemoryStoreGetByIssuerMissing(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
store := NewMemoryOIDCProviderStore()
|
||||
_, err := store.GetProviderByIssuer(ctx, "", "https://other.example/")
|
||||
if err == nil {
|
||||
t.Fatal("expected error for unregistered issuer")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDeriveOIDCProviderARN(t *testing.T) {
|
||||
cases := []struct {
|
||||
name string
|
||||
accountID string
|
||||
issuer string
|
||||
want string
|
||||
}{
|
||||
{
|
||||
name: "google",
|
||||
accountID: "111122223333",
|
||||
issuer: "https://accounts.google.com",
|
||||
want: "arn:aws:iam::111122223333:oidc-provider/accounts.google.com",
|
||||
},
|
||||
{
|
||||
name: "EKS with path",
|
||||
accountID: "999999999999",
|
||||
issuer: "https://oidc.eks.us-west-2.amazonaws.com/id/EXAMPLED",
|
||||
want: "arn:aws:iam::999999999999:oidc-provider/oidc.eks.us-west-2.amazonaws.com/id/EXAMPLED",
|
||||
},
|
||||
{
|
||||
name: "uppercase host normalized",
|
||||
accountID: "111122223333",
|
||||
issuer: "https://Accounts.Google.com",
|
||||
want: "arn:aws:iam::111122223333:oidc-provider/accounts.google.com",
|
||||
},
|
||||
{
|
||||
name: "trailing slash trimmed",
|
||||
accountID: "111122223333",
|
||||
issuer: "https://accounts.google.com/",
|
||||
want: "arn:aws:iam::111122223333:oidc-provider/accounts.google.com",
|
||||
},
|
||||
{
|
||||
name: "empty account allowed",
|
||||
accountID: "",
|
||||
issuer: "https://accounts.google.com",
|
||||
want: "arn:aws:iam:::oidc-provider/accounts.google.com",
|
||||
},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
got, err := DeriveOIDCProviderARN(tc.accountID, tc.issuer)
|
||||
if err != nil {
|
||||
t.Fatalf("err: %v", err)
|
||||
}
|
||||
if got != tc.want {
|
||||
t.Fatalf("got=%s want=%s", got, tc.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestDeriveOIDCProviderARNRejectsBadInput(t *testing.T) {
|
||||
cases := []string{"", "not a url", "http://"}
|
||||
for _, in := range cases {
|
||||
if _, err := DeriveOIDCProviderARN("123", in); err == nil {
|
||||
t.Errorf("expected error for input %q", in)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestStoreRejectsEmptyARN(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
store := NewMemoryOIDCProviderStore()
|
||||
rec := newRecord("", "https://issuer/")
|
||||
if err := store.StoreProvider(ctx, "", rec); err == nil {
|
||||
t.Fatal("expected error storing record with empty ARN")
|
||||
}
|
||||
}
|
||||
|
||||
func TestStoreRejectsNil(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
store := NewMemoryOIDCProviderStore()
|
||||
if err := store.StoreProvider(ctx, "", nil); err == nil {
|
||||
t.Fatal("expected error storing nil record")
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeIssuerRejectsEmpty(t *testing.T) {
|
||||
if got := normalizeIssuer(""); got != "" {
|
||||
t.Fatalf("empty issuer should normalize to empty, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeIssuerHandlesNonURL(t *testing.T) {
|
||||
// Defense in depth: even when issuer is junk, normalize doesn't panic and
|
||||
// at least lowercases the input.
|
||||
got := normalizeIssuer("Some Random String/")
|
||||
if !strings.Contains(got, "some random") {
|
||||
t.Fatalf("normalize should lowercase non-URL input, got %q", got)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,34 @@
|
||||
package integration
|
||||
|
||||
import "testing"
|
||||
|
||||
func intPtr(v int64) *int64 { return &v }
|
||||
|
||||
func TestCapDurationByRole(t *testing.T) {
|
||||
cases := []struct {
|
||||
name string
|
||||
requested *int64
|
||||
roleMax int64
|
||||
want *int64
|
||||
}{
|
||||
{"no cap, no request", nil, 0, nil},
|
||||
{"no cap, with request", intPtr(7200), 0, intPtr(7200)},
|
||||
{"cap only, no request -> nil so STS default applies", nil, 3600, nil},
|
||||
{"request below cap -> request", intPtr(1800), 3600, intPtr(1800)},
|
||||
{"request equal cap -> request", intPtr(3600), 3600, intPtr(3600)},
|
||||
{"request above cap -> cap", intPtr(43200), 3600, intPtr(3600)},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
got := capDurationByRole(tc.requested, tc.roleMax)
|
||||
switch {
|
||||
case got == nil && tc.want == nil:
|
||||
return
|
||||
case got == nil || tc.want == nil:
|
||||
t.Fatalf("nilness mismatch: got=%v want=%v", got, tc.want)
|
||||
case *got != *tc.want:
|
||||
t.Fatalf("got=%d want=%d", *got, *tc.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,197 @@
|
||||
package oidc
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// fakeIDP wraps an httptest.Server and counts how many times each well-known
|
||||
// endpoint is hit. Tests use it to assert discovery vs. fallback behaviour.
|
||||
type fakeIDP struct {
|
||||
server *httptest.Server
|
||||
discoveryHits atomic.Int32
|
||||
jwksHits atomic.Int32
|
||||
customJWKSHits atomic.Int32
|
||||
disableDiscovery bool
|
||||
discoveryStatusCode int
|
||||
discoveryIssuer string
|
||||
omitDiscoveryIssuer bool // when true, the discovery doc omits the "issuer" field entirely
|
||||
customJWKSPathSuffix string // optional suffix that fakeIDP serves at /custom/<suffix>
|
||||
jwks JWKS
|
||||
}
|
||||
|
||||
func newFakeIDP(t *testing.T) *fakeIDP {
|
||||
t.Helper()
|
||||
idp := &fakeIDP{
|
||||
discoveryStatusCode: http.StatusOK,
|
||||
jwks: JWKS{Keys: []JWK{{Kty: "RSA", Kid: "k1", Use: "sig", Alg: "RS256", N: "AQAB", E: "AQAB"}}},
|
||||
}
|
||||
mux := http.NewServeMux()
|
||||
mux.HandleFunc("/.well-known/openid-configuration", func(w http.ResponseWriter, r *http.Request) {
|
||||
idp.discoveryHits.Add(1)
|
||||
if idp.disableDiscovery {
|
||||
http.NotFound(w, r)
|
||||
return
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.WriteHeader(idp.discoveryStatusCode)
|
||||
issuer := idp.discoveryIssuer
|
||||
if issuer == "" {
|
||||
issuer = idp.server.URL
|
||||
}
|
||||
jwksURI := idp.server.URL + "/discovered/jwks"
|
||||
if idp.customJWKSPathSuffix != "" {
|
||||
jwksURI = idp.server.URL + "/custom/" + idp.customJWKSPathSuffix
|
||||
}
|
||||
body := map[string]string{"jwks_uri": jwksURI}
|
||||
if !idp.omitDiscoveryIssuer {
|
||||
body["issuer"] = issuer
|
||||
}
|
||||
_ = json.NewEncoder(w).Encode(body)
|
||||
})
|
||||
mux.HandleFunc("/discovered/jwks", func(w http.ResponseWriter, r *http.Request) {
|
||||
idp.jwksHits.Add(1)
|
||||
_ = json.NewEncoder(w).Encode(idp.jwks)
|
||||
})
|
||||
mux.HandleFunc("/.well-known/jwks.json", func(w http.ResponseWriter, r *http.Request) {
|
||||
idp.jwksHits.Add(1)
|
||||
_ = json.NewEncoder(w).Encode(idp.jwks)
|
||||
})
|
||||
mux.HandleFunc("/custom/", func(w http.ResponseWriter, r *http.Request) {
|
||||
idp.customJWKSHits.Add(1)
|
||||
_ = json.NewEncoder(w).Encode(idp.jwks)
|
||||
})
|
||||
idp.server = httptest.NewServer(mux)
|
||||
t.Cleanup(idp.server.Close)
|
||||
return idp
|
||||
}
|
||||
|
||||
func newProviderForIDP(t *testing.T, idp *fakeIDP, jwksURIOverride string) *OIDCProvider {
|
||||
t.Helper()
|
||||
p := NewOIDCProvider("test")
|
||||
cfg := &OIDCConfig{
|
||||
Issuer: idp.server.URL,
|
||||
ClientID: "test-client",
|
||||
JWKSUri: jwksURIOverride,
|
||||
}
|
||||
if err := p.Initialize(cfg); err != nil {
|
||||
t.Fatalf("Initialize: %v", err)
|
||||
}
|
||||
return p
|
||||
}
|
||||
|
||||
func TestDiscoveryHappyPath(t *testing.T) {
|
||||
idp := newFakeIDP(t)
|
||||
p := newProviderForIDP(t, idp, "")
|
||||
|
||||
if err := p.fetchJWKS(context.Background()); err != nil {
|
||||
t.Fatalf("fetchJWKS: %v", err)
|
||||
}
|
||||
if got := idp.discoveryHits.Load(); got != 1 {
|
||||
t.Fatalf("expected 1 discovery hit, got %d", got)
|
||||
}
|
||||
if got := idp.jwksHits.Load(); got != 1 {
|
||||
t.Fatalf("expected 1 JWKS hit at discovered uri, got %d", got)
|
||||
}
|
||||
|
||||
// A second fetch reuses the cached jwks_uri without re-discovering.
|
||||
if err := p.fetchJWKS(context.Background()); err != nil {
|
||||
t.Fatalf("fetchJWKS second: %v", err)
|
||||
}
|
||||
if got := idp.discoveryHits.Load(); got != 1 {
|
||||
t.Fatalf("discovery should be cached, got %d hits", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiscoveryFallback404(t *testing.T) {
|
||||
idp := newFakeIDP(t)
|
||||
idp.disableDiscovery = true
|
||||
p := newProviderForIDP(t, idp, "")
|
||||
|
||||
if err := p.fetchJWKS(context.Background()); err != nil {
|
||||
t.Fatalf("fetchJWKS: %v", err)
|
||||
}
|
||||
if got := idp.discoveryHits.Load(); got != 1 {
|
||||
t.Fatalf("expected 1 discovery probe, got %d", got)
|
||||
}
|
||||
if got := idp.jwksHits.Load(); got != 1 {
|
||||
t.Fatalf("expected 1 JWKS hit at fallback uri, got %d", got)
|
||||
}
|
||||
|
||||
// Subsequent fetches retry discovery — discoveryFailed resets at the top
|
||||
// of fetchJWKSLocked when no URI was cached, so a transient 5xx at
|
||||
// startup doesn't lock the provider into the fallback path forever.
|
||||
// Retry rate is bounded by the JWKS TTL (one retry per refresh cycle).
|
||||
if err := p.fetchJWKS(context.Background()); err != nil {
|
||||
t.Fatalf("fetchJWKS second: %v", err)
|
||||
}
|
||||
if got := idp.discoveryHits.Load(); got != 2 {
|
||||
t.Fatalf("discovery probe should retry while no URI is cached, got %d hits", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiscoveryDisabledByExplicitJWKSUri(t *testing.T) {
|
||||
idp := newFakeIDP(t)
|
||||
override := idp.server.URL + "/custom/explicit"
|
||||
idp.customJWKSPathSuffix = "explicit"
|
||||
p := newProviderForIDP(t, idp, override)
|
||||
|
||||
if err := p.fetchJWKS(context.Background()); err != nil {
|
||||
t.Fatalf("fetchJWKS: %v", err)
|
||||
}
|
||||
if got := idp.discoveryHits.Load(); got != 0 {
|
||||
t.Fatalf("explicit JWKSUri should bypass discovery, got %d hits", got)
|
||||
}
|
||||
if got := idp.customJWKSHits.Load(); got != 1 {
|
||||
t.Fatalf("expected 1 custom JWKS hit, got %d", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiscoveryRejectsIssuerMismatch(t *testing.T) {
|
||||
idp := newFakeIDP(t)
|
||||
idp.discoveryIssuer = "https://attacker.example/"
|
||||
p := newProviderForIDP(t, idp, "")
|
||||
|
||||
if err := p.fetchJWKS(context.Background()); err != nil {
|
||||
t.Fatalf("fetchJWKS should fall back to /.well-known/jwks.json, got error: %v", err)
|
||||
}
|
||||
// Discovery probe was tried once, rejected, then fell through to fallback path.
|
||||
if got := idp.discoveryHits.Load(); got != 1 {
|
||||
t.Fatalf("expected 1 discovery probe, got %d", got)
|
||||
}
|
||||
if got := idp.jwksHits.Load(); got != 1 {
|
||||
t.Fatalf("expected 1 fallback JWKS hit, got %d", got)
|
||||
}
|
||||
}
|
||||
|
||||
// TestDiscoveryRejectsMissingIssuer: a discovery document that omits the
|
||||
// issuer field entirely must be treated the same as one that supplies a
|
||||
// mismatched issuer. Otherwise an attacker who can intercept the discovery
|
||||
// response can strip the issuer field and the comparison silently passes,
|
||||
// letting the document point fetchJWKS at any URL it pleases.
|
||||
func TestDiscoveryRejectsMissingIssuer(t *testing.T) {
|
||||
idp := newFakeIDP(t)
|
||||
idp.omitDiscoveryIssuer = true
|
||||
p := newProviderForIDP(t, idp, "")
|
||||
|
||||
if err := p.fetchJWKS(context.Background()); err != nil {
|
||||
t.Fatalf("fetchJWKS should fall back to /.well-known/jwks.json on issuer-missing discovery: %v", err)
|
||||
}
|
||||
if got := idp.discoveryHits.Load(); got != 1 {
|
||||
t.Fatalf("expected 1 discovery probe, got %d", got)
|
||||
}
|
||||
// The discovery document was rejected; the JWKS that ultimately served
|
||||
// us must be the fallback one, not the discovered URI. The fakeIDP
|
||||
// counts both hits under jwksHits since they share a counter; what
|
||||
// matters is that customJWKSHits stayed zero.
|
||||
if got := idp.customJWKSHits.Load(); got != 0 {
|
||||
t.Fatalf("custom JWKS endpoint must not have been used, got %d hits", got)
|
||||
}
|
||||
if got := idp.jwksHits.Load(); got != 1 {
|
||||
t.Fatalf("expected 1 fallback JWKS hit, got %d", got)
|
||||
}
|
||||
}
|
||||
+161
-21
@@ -16,6 +16,7 @@ import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/golang-jwt/jwt/v5"
|
||||
@@ -25,13 +26,21 @@ import (
|
||||
|
||||
// OIDCProvider implements OpenID Connect authentication
|
||||
type OIDCProvider struct {
|
||||
name string
|
||||
config *OIDCConfig
|
||||
initialized bool
|
||||
jwksCache *JWKS
|
||||
httpClient *http.Client
|
||||
jwksFetchedAt time.Time
|
||||
jwksTTL time.Duration
|
||||
name string
|
||||
config *OIDCConfig
|
||||
initialized bool
|
||||
httpClient *http.Client
|
||||
jwksTTL time.Duration
|
||||
|
||||
// mu guards the lazily-mutated cache fields below: jwksCache, jwksFetchedAt,
|
||||
// resolvedJWKSUri, and discoveryFailed are all populated on the first
|
||||
// validate-token call and refreshed when the cache expires. Multiple S3
|
||||
// requests can land here in parallel, so they need synchronization.
|
||||
mu sync.RWMutex
|
||||
jwksCache *JWKS
|
||||
jwksFetchedAt time.Time
|
||||
resolvedJWKSUri string
|
||||
discoveryFailed bool
|
||||
}
|
||||
|
||||
// OIDCConfig holds OIDC provider configuration
|
||||
@@ -272,6 +281,7 @@ func (p *OIDCProvider) Authenticate(ctx context.Context, token string) (*provide
|
||||
Groups: groups,
|
||||
Attributes: attributes,
|
||||
Provider: p.name,
|
||||
Issuer: claims.Issuer,
|
||||
}
|
||||
|
||||
// Pass the token expiration to limit session duration
|
||||
@@ -550,39 +560,169 @@ func (p *OIDCProvider) mapClaimsToRolesWithConfig(claims *providers.TokenClaims)
|
||||
return roles
|
||||
}
|
||||
|
||||
// getPublicKey retrieves the public key for the given key ID from JWKS
|
||||
// getPublicKey retrieves the public key for the given key ID from JWKS.
|
||||
// Cache hits use the read lock so concurrent token validations don't
|
||||
// serialize on JWKS lookup. Misses and expirations promote to the write
|
||||
// lock so the JWKS fetch + cache write happens once per refresh cycle.
|
||||
func (p *OIDCProvider) getPublicKey(ctx context.Context, kid string) (interface{}, error) {
|
||||
// Fetch JWKS if not cached or refresh if expired
|
||||
if p.jwksCache == nil || (!p.jwksFetchedAt.IsZero() && time.Since(p.jwksFetchedAt) > p.jwksTTL) {
|
||||
if err := p.fetchJWKS(ctx); err != nil {
|
||||
// Fast path: read lock and look in cache.
|
||||
p.mu.RLock()
|
||||
if p.jwksCache != nil && (p.jwksFetchedAt.IsZero() || time.Since(p.jwksFetchedAt) <= p.jwksTTL) {
|
||||
for _, key := range p.jwksCache.Keys {
|
||||
if key.Kid == kid {
|
||||
k := key
|
||||
p.mu.RUnlock()
|
||||
return p.parseJWK(&k)
|
||||
}
|
||||
}
|
||||
}
|
||||
p.mu.RUnlock()
|
||||
|
||||
// Slow path: take the write lock for the (re)fetch + retry. Re-check the
|
||||
// cache under the write lock in case another goroutine already refreshed.
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
|
||||
cacheValid := p.jwksCache != nil && (p.jwksFetchedAt.IsZero() || time.Since(p.jwksFetchedAt) <= p.jwksTTL)
|
||||
if !cacheValid {
|
||||
if err := p.fetchJWKSLocked(ctx); err != nil {
|
||||
return nil, fmt.Errorf("failed to fetch JWKS: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// Find the key with matching kid
|
||||
for _, key := range p.jwksCache.Keys {
|
||||
if key.Kid == kid {
|
||||
return p.parseJWK(&key)
|
||||
k := key
|
||||
return p.parseJWK(&k)
|
||||
}
|
||||
}
|
||||
|
||||
// Key not found in cache. Refresh JWKS once to handle key rotation and retry.
|
||||
if err := p.fetchJWKS(ctx); err != nil {
|
||||
// Key not found in cache. Refresh JWKS once to handle key rotation.
|
||||
if err := p.fetchJWKSLocked(ctx); err != nil {
|
||||
return nil, fmt.Errorf("failed to refresh JWKS after key miss: %v", err)
|
||||
}
|
||||
for _, key := range p.jwksCache.Keys {
|
||||
if key.Kid == kid {
|
||||
return p.parseJWK(&key)
|
||||
k := key
|
||||
return p.parseJWK(&k)
|
||||
}
|
||||
}
|
||||
return nil, fmt.Errorf("key with ID %s not found in JWKS after refresh", kid)
|
||||
}
|
||||
|
||||
// fetchJWKS fetches the JWKS from the provider
|
||||
// discoveryDocument is the subset of the OpenID Provider Configuration we need.
|
||||
// See https://openid.net/specs/openid-connect-discovery-1_0.html#ProviderMetadata.
|
||||
type discoveryDocument struct {
|
||||
Issuer string `json:"issuer"`
|
||||
JWKSUri string `json:"jwks_uri"`
|
||||
}
|
||||
|
||||
// resolveJWKSUriLocked determines the JWKS URI for the provider. The caller
|
||||
// must hold p.mu (write lock); the function reads/writes p.resolvedJWKSUri
|
||||
// and p.discoveryFailed without taking the lock itself.
|
||||
//
|
||||
// Order of resolution:
|
||||
// 1. explicit config.JWKSUri (operator override; never overridden by discovery).
|
||||
// 2. cached resolvedJWKSUri from a prior discovery (refreshed when JWKS cache expires).
|
||||
// 3. .well-known/openid-configuration discovery (per OIDC Discovery 1.0).
|
||||
// 4. fallback to {issuer}/.well-known/jwks.json (compat path for IDPs that
|
||||
// don't publish discovery).
|
||||
func (p *OIDCProvider) resolveJWKSUriLocked(ctx context.Context) (string, error) {
|
||||
if p.config.JWKSUri != "" {
|
||||
return p.config.JWKSUri, nil
|
||||
}
|
||||
if p.resolvedJWKSUri != "" {
|
||||
return p.resolvedJWKSUri, nil
|
||||
}
|
||||
|
||||
issuer := strings.TrimSuffix(p.config.Issuer, "/")
|
||||
|
||||
if !p.discoveryFailed {
|
||||
discoveryURL := issuer + "/.well-known/openid-configuration"
|
||||
uri, err := p.fetchDiscoveryJWKSUri(ctx, discoveryURL)
|
||||
switch {
|
||||
case err == nil:
|
||||
p.resolvedJWKSUri = uri
|
||||
return uri, nil
|
||||
default:
|
||||
// Cache the failure so we don't pay the discovery RTT on every refresh.
|
||||
// Operators with non-discovery IDPs see one failed lookup at startup.
|
||||
glog.V(3).Infof("OIDC discovery at %s failed (%v); falling back to /.well-known/jwks.json", discoveryURL, err)
|
||||
p.discoveryFailed = true
|
||||
}
|
||||
}
|
||||
|
||||
return issuer + "/.well-known/jwks.json", nil
|
||||
}
|
||||
|
||||
// fetchDiscoveryJWKSUri retrieves the OIDC discovery document and returns
|
||||
// the jwks_uri field. The issuer claim in the document must match config.Issuer
|
||||
// to defend against issuer-substitution attacks during discovery.
|
||||
func (p *OIDCProvider) fetchDiscoveryJWKSUri(ctx context.Context, discoveryURL string) (string, error) {
|
||||
req, err := http.NewRequestWithContext(ctx, "GET", discoveryURL, nil)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("create discovery request: %v", err)
|
||||
}
|
||||
req.Header.Set("Accept", "application/json")
|
||||
|
||||
resp, err := p.httpClient.Do(req)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("fetch discovery document: %v", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return "", fmt.Errorf("discovery endpoint returned status %d", resp.StatusCode)
|
||||
}
|
||||
|
||||
var doc discoveryDocument
|
||||
if err := json.NewDecoder(resp.Body).Decode(&doc); err != nil {
|
||||
return "", fmt.Errorf("decode discovery document: %v", err)
|
||||
}
|
||||
|
||||
if doc.JWKSUri == "" {
|
||||
return "", fmt.Errorf("discovery document missing jwks_uri")
|
||||
}
|
||||
|
||||
// Issuer must be present and match: a discovery doc that points to a
|
||||
// different issuer is either a misconfiguration or an attack against
|
||||
// issuer-confusion, and a doc that omits the issuer field entirely
|
||||
// would have bypassed the previous check (doc.Issuer != "" guard) and
|
||||
// silently accepted whatever JWKS URI the document supplied. OIDC
|
||||
// Discovery 1.0 §3 mandates the issuer field, so treat missing as a
|
||||
// hard failure. Compare after trimming a single trailing slash on each
|
||||
// side because real IdPs disagree on whether the configured issuer
|
||||
// has one.
|
||||
if strings.TrimSuffix(doc.Issuer, "/") != strings.TrimSuffix(p.config.Issuer, "/") {
|
||||
return "", fmt.Errorf("discovery issuer %q does not match configured issuer %q", doc.Issuer, p.config.Issuer)
|
||||
}
|
||||
|
||||
return doc.JWKSUri, nil
|
||||
}
|
||||
|
||||
// fetchJWKS is a thin wrapper around fetchJWKSLocked that acquires the
|
||||
// write lock. Used by tests; production callers in getPublicKey already
|
||||
// hold the lock and call fetchJWKSLocked directly.
|
||||
func (p *OIDCProvider) fetchJWKS(ctx context.Context) error {
|
||||
jwksURL := p.config.JWKSUri
|
||||
if jwksURL == "" {
|
||||
jwksURL = strings.TrimSuffix(p.config.Issuer, "/") + "/.well-known/jwks.json"
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
return p.fetchJWKSLocked(ctx)
|
||||
}
|
||||
|
||||
// fetchJWKSLocked fetches the JWKS from the provider. The caller must hold
|
||||
// p.mu (write lock); the function writes p.jwksCache and p.jwksFetchedAt
|
||||
// without taking the lock itself.
|
||||
//
|
||||
// Each fetch reattempts discovery if the previous attempt failed: a
|
||||
// transient 5xx that flipped discoveryFailed at startup shouldn't lock the
|
||||
// provider into the fallback path forever. The retry rate is bounded by
|
||||
// the JWKS TTL (typically 1h), so the discovery RTT cost is amortized.
|
||||
func (p *OIDCProvider) fetchJWKSLocked(ctx context.Context) error {
|
||||
if p.config.JWKSUri == "" && p.resolvedJWKSUri == "" {
|
||||
p.discoveryFailed = false
|
||||
}
|
||||
jwksURL, err := p.resolveJWKSUriLocked(ctx)
|
||||
if err != nil {
|
||||
return fmt.Errorf("resolve JWKS URI: %v", err)
|
||||
}
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, "GET", jwksURL, nil)
|
||||
|
||||
@@ -49,6 +49,11 @@ type ExternalIdentity struct {
|
||||
// Provider is the name of the identity provider
|
||||
Provider string `json:"provider"`
|
||||
|
||||
// Issuer is the OIDC `iss` claim (or equivalent) from the source token.
|
||||
// Stable per (provider, identity) and used together with UserID to derive
|
||||
// a stable parent-user hash that survives token rotation.
|
||||
Issuer string `json:"issuer,omitempty"`
|
||||
|
||||
// TokenExpiration is the expiration time of the source identity token
|
||||
// This is used to limit session duration to not exceed the token's exp claim
|
||||
TokenExpiration *time.Time `json:"tokenExpiration,omitempty"`
|
||||
|
||||
@@ -13,6 +13,39 @@ type CommonResponse struct {
|
||||
} `xml:"ResponseMetadata"`
|
||||
}
|
||||
|
||||
// IAMTag mirrors AWS IAM's Tag list element with the Key/Value pair shape.
|
||||
type IAMTag struct {
|
||||
Key string `xml:"Key"`
|
||||
Value string `xml:"Value"`
|
||||
}
|
||||
|
||||
// OpenIDConnectProviderListEntry is one element of ListOpenIDConnectProviders.
|
||||
type OpenIDConnectProviderListEntry struct {
|
||||
Arn string `xml:"Arn"`
|
||||
}
|
||||
|
||||
// ListOpenIDConnectProvidersResponse is the response for ListOpenIDConnectProviders.
|
||||
type ListOpenIDConnectProvidersResponse struct {
|
||||
XMLName xml.Name `xml:"https://iam.amazonaws.com/doc/2010-05-08/ ListOpenIDConnectProvidersResponse"`
|
||||
ListOpenIDConnectProvidersResult struct {
|
||||
OpenIDConnectProviderList []*OpenIDConnectProviderListEntry `xml:"OpenIDConnectProviderList>member"`
|
||||
} `xml:"ListOpenIDConnectProvidersResult"`
|
||||
CommonResponse
|
||||
}
|
||||
|
||||
// GetOpenIDConnectProviderResponse is the response for GetOpenIDConnectProvider.
|
||||
type GetOpenIDConnectProviderResponse struct {
|
||||
XMLName xml.Name `xml:"https://iam.amazonaws.com/doc/2010-05-08/ GetOpenIDConnectProviderResponse"`
|
||||
GetOpenIDConnectProviderResult struct {
|
||||
Url string `xml:"Url"`
|
||||
ClientIDList []string `xml:"ClientIDList>member,omitempty"`
|
||||
ThumbprintList []string `xml:"ThumbprintList>member,omitempty"`
|
||||
Tags []*IAMTag `xml:"Tags>member,omitempty"`
|
||||
CreateDate string `xml:"CreateDate,omitempty"`
|
||||
} `xml:"GetOpenIDConnectProviderResult"`
|
||||
CommonResponse
|
||||
}
|
||||
|
||||
// SetRequestId stores the request ID generated for the current HTTP request.
|
||||
func (r *CommonResponse) SetRequestId(requestID string) {
|
||||
r.ResponseMetadata.RequestId = requestID
|
||||
|
||||
@@ -0,0 +1,65 @@
|
||||
package sts
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestComputeParentUserStability(t *testing.T) {
|
||||
// Same (sub, iss) must produce the same hash, regardless of order or
|
||||
// whitespace mutations callers should never apply.
|
||||
a := ComputeParentUser("alice", "https://idp.example/")
|
||||
b := ComputeParentUser("alice", "https://idp.example/")
|
||||
if a == "" {
|
||||
t.Fatal("parent user should not be empty for non-empty sub")
|
||||
}
|
||||
if a != b {
|
||||
t.Fatalf("expected stable hash, got %q vs %q", a, b)
|
||||
}
|
||||
}
|
||||
|
||||
func TestComputeParentUserDistinguishesIssuer(t *testing.T) {
|
||||
// The point of incorporating iss is that the same `sub` from two providers
|
||||
// must not collide. If this assertion ever fails, the hash input is wrong.
|
||||
a := ComputeParentUser("alice", "https://idp-a.example/")
|
||||
b := ComputeParentUser("alice", "https://idp-b.example/")
|
||||
if a == b {
|
||||
t.Fatalf("hashes for different issuers must differ, both = %q", a)
|
||||
}
|
||||
}
|
||||
|
||||
func TestComputeParentUserDistinguishesSubject(t *testing.T) {
|
||||
a := ComputeParentUser("alice", "https://idp.example/")
|
||||
b := ComputeParentUser("bob", "https://idp.example/")
|
||||
if a == b {
|
||||
t.Fatalf("hashes for different subjects must differ, both = %q", a)
|
||||
}
|
||||
}
|
||||
|
||||
func TestComputeParentUserEmptySub(t *testing.T) {
|
||||
if got := ComputeParentUser("", "https://idp.example/"); got != "" {
|
||||
t.Fatalf("empty sub should produce empty parent user, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestComputeParentUserEncoding(t *testing.T) {
|
||||
got := ComputeParentUser("alice", "https://idp.example/")
|
||||
// Base64 RawURL has no padding and uses URL-safe alphabet — important
|
||||
// because parent_user shows up in filer paths and audit log fields.
|
||||
if strings.ContainsAny(got, "=+/") {
|
||||
t.Fatalf("parent user should be base64 raw url, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSessionClaimsRoundTripParentUser(t *testing.T) {
|
||||
parent := ComputeParentUser("alice", "https://idp.example/")
|
||||
claims := NewSTSSessionClaims("sid-1", "issuer", time.Now().Add(time.Hour)).
|
||||
WithRoleInfo("arn:aws:iam::123:role/r", "arn:aws:sts::123:assumed-role/r/s", "arn:aws:sts::123:assumed-role/r/s").
|
||||
WithParentUser(parent)
|
||||
|
||||
info := claims.ToSessionInfo()
|
||||
if info.ParentUser != parent {
|
||||
t.Fatalf("ParentUser lost on round-trip: got %q want %q", info.ParentUser, parent)
|
||||
}
|
||||
}
|
||||
@@ -1,6 +1,8 @@
|
||||
package sts
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"encoding/base64"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
@@ -8,6 +10,20 @@ import (
|
||||
"github.com/seaweedfs/seaweedfs/weed/glog"
|
||||
)
|
||||
|
||||
// ComputeParentUser returns a stable per-identity hash derived from the OIDC
|
||||
// (sub, iss) tuple. Only the (sub, iss) pair is guaranteed stable across token
|
||||
// refreshes per OpenID Connect Core 1.0 §5.7, so any per-user state (audit
|
||||
// logs, quotas) must key off this value rather than the access-key or session
|
||||
// id. The hash is base64-rawurl-encoded SHA-256 over "openid:<sub>:<iss>" so
|
||||
// it stays filesystem-safe and bounded in length for storage in audit paths.
|
||||
func ComputeParentUser(sub, iss string) string {
|
||||
if sub == "" {
|
||||
return ""
|
||||
}
|
||||
h := sha256.Sum256([]byte("openid:" + sub + ":" + iss))
|
||||
return base64.RawURLEncoding.EncodeToString(h[:])
|
||||
}
|
||||
|
||||
// defaultCredentialGenerator is a reusable instance for generating temporary credentials
|
||||
// Reusing a single instance across all calls to ToSessionInfo() reduces allocation overhead
|
||||
// since this method may be called frequently during signature verification
|
||||
@@ -45,6 +61,12 @@ type STSSessionClaims struct {
|
||||
// Session metadata
|
||||
AssumedAt time.Time `json:"assumed_at"` // when role was assumed
|
||||
MaxDuration int64 `json:"max_dur,omitempty"` // maximum session duration in seconds
|
||||
|
||||
// ParentUser is a stable hash of (sub, iss) for tokens minted from an OIDC
|
||||
// identity. It survives token rotation since only the (sub, iss) tuple is
|
||||
// guaranteed stable per OpenID Connect Core 1.0. Empty for non-federated
|
||||
// session types.
|
||||
ParentUser string `json:"puid,omitempty"`
|
||||
}
|
||||
|
||||
// NewSTSSessionClaims creates new STS session claims with all required information
|
||||
@@ -96,6 +118,7 @@ func (c *STSSessionClaims) ToSessionInfo() *SessionInfo {
|
||||
ExternalUserId: c.ExternalUserId,
|
||||
ProviderIssuer: c.ProviderIssuer,
|
||||
RequestContext: c.RequestContext,
|
||||
ParentUser: c.ParentUser,
|
||||
// Provide the Subject (sub) from registered claims
|
||||
Subject: c.Subject,
|
||||
Credentials: credentials,
|
||||
@@ -182,3 +205,10 @@ func (c *STSSessionClaims) WithSessionName(sessionName string) *STSSessionClaims
|
||||
c.SessionName = sessionName
|
||||
return c
|
||||
}
|
||||
|
||||
// WithParentUser sets the stable per-identity hash for the session. See
|
||||
// ComputeParentUser for the derivation rule.
|
||||
func (c *STSSessionClaims) WithParentUser(parentUser string) *STSSessionClaims {
|
||||
c.ParentUser = parentUser
|
||||
return c
|
||||
}
|
||||
|
||||
@@ -254,6 +254,9 @@ type SessionInfo struct {
|
||||
|
||||
// Credentials are the temporary credentials for this session
|
||||
Credentials *Credentials `json:"credentials"`
|
||||
|
||||
// ParentUser is the stable hashed identity (sub+iss) derived at federation time.
|
||||
ParentUser string `json:"parentUser,omitempty"`
|
||||
}
|
||||
|
||||
// NewSTSService creates a new STS service
|
||||
@@ -506,13 +509,26 @@ func (s *STSService) AssumeRoleWithWebIdentity(ctx context.Context, request *Ass
|
||||
// Add sub as well since it's commonly used
|
||||
requestContext["sub"] = externalIdentity.UserID
|
||||
|
||||
// Compute a stable parent-user hash from (sub, iss). Only this tuple is
|
||||
// guaranteed stable across token refresh per OIDC Core 1.0, so this is the
|
||||
// right key for any per-identity state (audit trail, future quotas).
|
||||
parentUser := ComputeParentUser(externalIdentity.UserID, externalIdentity.Issuer)
|
||||
if parentUser != "" {
|
||||
// Surface as aws:userid so policies can reference it directly without
|
||||
// caring about token-rotation churn.
|
||||
requestContext["aws:userid"] = parentUser
|
||||
}
|
||||
|
||||
// Create rich JWT claims with all session information
|
||||
sessionClaims := NewSTSSessionClaims(sessionId, s.Config.Issuer, expiresAt).
|
||||
WithSessionName(request.RoleSessionName).
|
||||
WithRoleInfo(request.RoleArn, assumedRoleUser.Arn, assumedRoleUser.Arn).
|
||||
WithIdentityProvider(provider.Name(), externalIdentity.UserID, "").
|
||||
WithIdentityProvider(provider.Name(), externalIdentity.UserID, externalIdentity.Issuer).
|
||||
WithMaxDuration(sessionDuration).
|
||||
WithRequestContext(requestContext)
|
||||
if parentUser != "" {
|
||||
sessionClaims.WithParentUser(parentUser)
|
||||
}
|
||||
if sessionPolicy != "" {
|
||||
sessionClaims.WithSessionPolicy(sessionPolicy)
|
||||
}
|
||||
|
||||
@@ -2165,13 +2165,24 @@ func (e *EmbeddedIamApi) ExecuteAction(ctx context.Context, values url.Values, s
|
||||
if e.readOnly {
|
||||
switch action {
|
||||
case "ListUsers", "ListAccessKeys", "GetUser", "GetUserPolicy", "ListUserPolicies", "ListAttachedUserPolicies", "ListPolicies", "GetPolicy", "ListPolicyVersions", "GetPolicyVersion", "ListServiceAccounts", "GetServiceAccount",
|
||||
"GetGroup", "ListGroups", "ListAttachedGroupPolicies", "GetGroupPolicy", "ListGroupPolicies", "ListGroupsForUser":
|
||||
"GetGroup", "ListGroups", "ListAttachedGroupPolicies", "GetGroupPolicy", "ListGroupPolicies", "ListGroupsForUser",
|
||||
actionListOpenIDConnectProviders, actionGetOpenIDConnectProvider:
|
||||
// Allowed read-only actions
|
||||
default:
|
||||
return nil, &iamError{Code: s3err.GetAPIError(s3err.ErrAccessDenied).Code, Error: fmt.Errorf("IAM write operations are disabled on this server")}
|
||||
}
|
||||
}
|
||||
|
||||
// OIDC provider actions don't operate on S3ApiConfiguration; dispatch
|
||||
// before the unrelated config load + reload churn.
|
||||
if response, iamErr, ok := e.dispatchOIDCProviderAction(ctx, values); ok {
|
||||
if iamErr != nil {
|
||||
return nil, iamErr
|
||||
}
|
||||
response.SetRequestId(reqID)
|
||||
return response, nil
|
||||
}
|
||||
|
||||
s3cfg := &iam_pb.S3ApiConfiguration{}
|
||||
if err := e.GetS3ApiConfiguration(s3cfg); err != nil && !errors.Is(err, filer_pb.ErrNotFound) {
|
||||
return nil, &iamError{Code: s3err.GetAPIError(s3err.ErrInternalError).Code, Error: fmt.Errorf("failed to get s3 api configuration: %v", err)}
|
||||
|
||||
@@ -0,0 +1,118 @@
|
||||
package s3api
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/url"
|
||||
"strings"
|
||||
|
||||
"github.com/aws/aws-sdk-go/service/iam"
|
||||
iamlib "github.com/seaweedfs/seaweedfs/weed/iam"
|
||||
"github.com/seaweedfs/seaweedfs/weed/iam/integration"
|
||||
)
|
||||
|
||||
// OIDC provider IAM actions handled by this file. Mutating actions are
|
||||
// reserved for Phase 2b; the read-only set lands now so existing static-config
|
||||
// providers become discoverable through the AWS IAM API.
|
||||
const (
|
||||
actionGetOpenIDConnectProvider = "GetOpenIDConnectProvider"
|
||||
actionListOpenIDConnectProviders = "ListOpenIDConnectProviders"
|
||||
)
|
||||
|
||||
// isOIDCProviderAction reports whether an action belongs to the OIDC provider
|
||||
// family. Used by ExecuteAction to short-circuit the S3ApiConfiguration code
|
||||
// path for actions that don't operate on it.
|
||||
func isOIDCProviderAction(action string) bool {
|
||||
switch action {
|
||||
case actionGetOpenIDConnectProvider,
|
||||
actionListOpenIDConnectProviders:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
// dispatchOIDCProviderAction handles the OIDC provider IAM actions. Returns a
|
||||
// response, the IAM error if any, and a boolean indicating whether the action
|
||||
// was recognised (so the caller can fall through when false).
|
||||
func (e *EmbeddedIamApi) dispatchOIDCProviderAction(ctx context.Context, values url.Values) (iamlib.RequestIDSetter, *iamError, bool) {
|
||||
if !isOIDCProviderAction(values.Get("Action")) {
|
||||
return nil, nil, false
|
||||
}
|
||||
|
||||
mgr := e.oidcIAMManager()
|
||||
if mgr == nil {
|
||||
return nil, &iamError{
|
||||
Code: iam.ErrCodeServiceFailureException,
|
||||
Error: errors.New("OIDC provider store not configured"),
|
||||
}, true
|
||||
}
|
||||
|
||||
switch values.Get("Action") {
|
||||
case actionListOpenIDConnectProviders:
|
||||
resp, err := e.listOpenIDConnectProviders(ctx, mgr)
|
||||
return resp, err, true
|
||||
case actionGetOpenIDConnectProvider:
|
||||
resp, err := e.getOpenIDConnectProvider(ctx, mgr, values)
|
||||
return resp, err, true
|
||||
}
|
||||
return nil, nil, false
|
||||
}
|
||||
|
||||
func (e *EmbeddedIamApi) oidcIAMManager() *integration.IAMManager {
|
||||
if e.iam == nil || e.iam.iamIntegration == nil {
|
||||
return nil
|
||||
}
|
||||
provider, ok := e.iam.iamIntegration.(IAMManagerProvider)
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
return provider.GetIAMManager()
|
||||
}
|
||||
|
||||
func (e *EmbeddedIamApi) listOpenIDConnectProviders(ctx context.Context, mgr *integration.IAMManager) (*iamlib.ListOpenIDConnectProvidersResponse, *iamError) {
|
||||
records, err := mgr.ListOIDCProviders(ctx)
|
||||
if err != nil {
|
||||
return nil, &iamError{Code: iam.ErrCodeServiceFailureException, Error: err}
|
||||
}
|
||||
resp := &iamlib.ListOpenIDConnectProvidersResponse{}
|
||||
resp.ListOpenIDConnectProvidersResult.OpenIDConnectProviderList = make([]*iamlib.OpenIDConnectProviderListEntry, 0, len(records))
|
||||
for _, rec := range records {
|
||||
resp.ListOpenIDConnectProvidersResult.OpenIDConnectProviderList = append(
|
||||
resp.ListOpenIDConnectProvidersResult.OpenIDConnectProviderList,
|
||||
&iamlib.OpenIDConnectProviderListEntry{Arn: rec.ARN},
|
||||
)
|
||||
}
|
||||
return resp, nil
|
||||
}
|
||||
|
||||
func (e *EmbeddedIamApi) getOpenIDConnectProvider(ctx context.Context, mgr *integration.IAMManager, values url.Values) (*iamlib.GetOpenIDConnectProviderResponse, *iamError) {
|
||||
arn := strings.TrimSpace(values.Get("OpenIDConnectProviderArn"))
|
||||
if arn == "" {
|
||||
return nil, &iamError{
|
||||
Code: iam.ErrCodeInvalidInputException,
|
||||
Error: fmt.Errorf("OpenIDConnectProviderArn is required"),
|
||||
}
|
||||
}
|
||||
rec, err := mgr.GetOIDCProvider(ctx, arn)
|
||||
if err != nil {
|
||||
return nil, &iamError{Code: iam.ErrCodeNoSuchEntityException, Error: err}
|
||||
}
|
||||
resp := &iamlib.GetOpenIDConnectProviderResponse{}
|
||||
resp.GetOpenIDConnectProviderResult.Url = rec.URL
|
||||
resp.GetOpenIDConnectProviderResult.ClientIDList = append([]string(nil), rec.ClientIDs...)
|
||||
resp.GetOpenIDConnectProviderResult.ThumbprintList = append([]string(nil), rec.Thumbprints...)
|
||||
if !rec.CreatedAt.IsZero() {
|
||||
// AWS uses ISO-8601; the IAM XML format accepts time.Time-string output.
|
||||
resp.GetOpenIDConnectProviderResult.CreateDate = rec.CreatedAt.UTC().Format("2006-01-02T15:04:05Z")
|
||||
}
|
||||
if len(rec.Tags) > 0 {
|
||||
tags := make([]*iamlib.IAMTag, 0, len(rec.Tags))
|
||||
for k, v := range rec.Tags {
|
||||
tags = append(tags, &iamlib.IAMTag{Key: k, Value: v})
|
||||
}
|
||||
resp.GetOpenIDConnectProviderResult.Tags = tags
|
||||
}
|
||||
return resp, nil
|
||||
}
|
||||
@@ -0,0 +1,151 @@
|
||||
package s3api
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/url"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
iamlib "github.com/seaweedfs/seaweedfs/weed/iam"
|
||||
"github.com/seaweedfs/seaweedfs/weed/iam/integration"
|
||||
"github.com/seaweedfs/seaweedfs/weed/iam/policy"
|
||||
"github.com/seaweedfs/seaweedfs/weed/iam/sts"
|
||||
)
|
||||
|
||||
// stubIntegration is the smallest IAMManagerProvider that lets the OIDC
|
||||
// dispatcher reach an IAMManager. The other IAMIntegration methods are
|
||||
// unused by these tests and panic if invoked, which is what we want — any
|
||||
// unexpected call signals a routing bug.
|
||||
type stubIntegration struct {
|
||||
IAMIntegration
|
||||
mgr *integration.IAMManager
|
||||
}
|
||||
|
||||
func (s *stubIntegration) GetIAMManager() *integration.IAMManager { return s.mgr }
|
||||
|
||||
func newOIDCTestAPI(t *testing.T) (*EmbeddedIamApiForTest, *integration.IAMManager) {
|
||||
t.Helper()
|
||||
mgr := integration.NewIAMManager()
|
||||
cfg := &integration.IAMConfig{
|
||||
STS: &sts.STSConfig{
|
||||
TokenDuration: sts.FlexibleDuration{Duration: time.Hour},
|
||||
MaxSessionLength: sts.FlexibleDuration{Duration: 12 * time.Hour},
|
||||
Issuer: "test-sts",
|
||||
SigningKey: []byte("test-signing-key-32-characters-long"),
|
||||
AccountId: "111122223333",
|
||||
Providers: []*sts.ProviderConfig{
|
||||
{
|
||||
Name: "google",
|
||||
Type: sts.ProviderTypeOIDC,
|
||||
Enabled: true,
|
||||
Config: map[string]interface{}{
|
||||
"issuer": "https://accounts.google.com",
|
||||
"clientId": "client-google",
|
||||
},
|
||||
},
|
||||
{
|
||||
Name: "github",
|
||||
Type: sts.ProviderTypeOIDC,
|
||||
Enabled: true,
|
||||
Config: map[string]interface{}{
|
||||
"issuer": "https://token.actions.githubusercontent.com",
|
||||
"clientId": "sts.amazonaws.com",
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
Policy: &policy.PolicyEngineConfig{DefaultEffect: "Deny", StoreType: "memory"},
|
||||
Roles: &integration.RoleStoreConfig{StoreType: "memory"},
|
||||
}
|
||||
if err := mgr.Initialize(cfg, func() string { return "localhost:8888" }); err != nil {
|
||||
t.Fatalf("Initialize IAM manager: %v", err)
|
||||
}
|
||||
|
||||
api := NewEmbeddedIamApiForTest()
|
||||
api.iam.iamIntegration = &stubIntegration{mgr: mgr}
|
||||
return api, mgr
|
||||
}
|
||||
|
||||
func TestListOpenIDConnectProviders(t *testing.T) {
|
||||
api, _ := newOIDCTestAPI(t)
|
||||
values := url.Values{}
|
||||
values.Set("Action", actionListOpenIDConnectProviders)
|
||||
|
||||
resp, iamErr := api.ExecuteAction(context.Background(), values, true, "test-req-1")
|
||||
if iamErr != nil {
|
||||
t.Fatalf("ExecuteAction: code=%s err=%v", iamErr.Code, iamErr.Error)
|
||||
}
|
||||
listResp, ok := resp.(*iamlib.ListOpenIDConnectProvidersResponse)
|
||||
if !ok {
|
||||
t.Fatalf("unexpected response type %T", resp)
|
||||
}
|
||||
got := listResp.ListOpenIDConnectProvidersResult.OpenIDConnectProviderList
|
||||
if len(got) != 2 {
|
||||
t.Fatalf("expected 2 providers, got %d", len(got))
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetOpenIDConnectProvider(t *testing.T) {
|
||||
api, _ := newOIDCTestAPI(t)
|
||||
arn := "arn:aws:iam::111122223333:oidc-provider/accounts.google.com"
|
||||
|
||||
values := url.Values{}
|
||||
values.Set("Action", actionGetOpenIDConnectProvider)
|
||||
values.Set("OpenIDConnectProviderArn", arn)
|
||||
|
||||
resp, iamErr := api.ExecuteAction(context.Background(), values, true, "test-req-2")
|
||||
if iamErr != nil {
|
||||
t.Fatalf("ExecuteAction: code=%s err=%v", iamErr.Code, iamErr.Error)
|
||||
}
|
||||
getResp, ok := resp.(*iamlib.GetOpenIDConnectProviderResponse)
|
||||
if !ok {
|
||||
t.Fatalf("unexpected response type %T", resp)
|
||||
}
|
||||
if getResp.GetOpenIDConnectProviderResult.Url != "https://accounts.google.com" {
|
||||
t.Fatalf("URL mismatch: %s", getResp.GetOpenIDConnectProviderResult.Url)
|
||||
}
|
||||
if len(getResp.GetOpenIDConnectProviderResult.ClientIDList) != 1 ||
|
||||
getResp.GetOpenIDConnectProviderResult.ClientIDList[0] != "client-google" {
|
||||
t.Fatalf("ClientIDList wrong: %v", getResp.GetOpenIDConnectProviderResult.ClientIDList)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetOpenIDConnectProviderMissing(t *testing.T) {
|
||||
api, _ := newOIDCTestAPI(t)
|
||||
values := url.Values{}
|
||||
values.Set("Action", actionGetOpenIDConnectProvider)
|
||||
values.Set("OpenIDConnectProviderArn", "arn:aws:iam::111122223333:oidc-provider/nope.example")
|
||||
|
||||
_, iamErr := api.ExecuteAction(context.Background(), values, true, "test-req-3")
|
||||
if iamErr == nil {
|
||||
t.Fatal("expected NoSuchEntity error")
|
||||
}
|
||||
if iamErr.Code != "NoSuchEntity" {
|
||||
t.Fatalf("expected NoSuchEntity code, got %s", iamErr.Code)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetOpenIDConnectProviderRequiresArn(t *testing.T) {
|
||||
api, _ := newOIDCTestAPI(t)
|
||||
values := url.Values{}
|
||||
values.Set("Action", actionGetOpenIDConnectProvider)
|
||||
|
||||
_, iamErr := api.ExecuteAction(context.Background(), values, true, "test-req-4")
|
||||
if iamErr == nil {
|
||||
t.Fatal("expected error for missing ARN")
|
||||
}
|
||||
if iamErr.Code != "InvalidInput" {
|
||||
t.Fatalf("expected InvalidInput code, got %s", iamErr.Code)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadOnlyAllowsOIDCList(t *testing.T) {
|
||||
api, _ := newOIDCTestAPI(t)
|
||||
api.readOnly = true
|
||||
values := url.Values{}
|
||||
values.Set("Action", actionListOpenIDConnectProviders)
|
||||
|
||||
if _, iamErr := api.ExecuteAction(context.Background(), values, true, "ro-1"); iamErr != nil {
|
||||
t.Fatalf("read-only mode should allow ListOpenIDConnectProviders: %v", iamErr.Error)
|
||||
}
|
||||
}
|
||||
+69
-17
@@ -53,6 +53,34 @@ const (
|
||||
// federationNameRegex validates the Name parameter for GetFederationToken per AWS spec
|
||||
var federationNameRegex = regexp.MustCompile(`^[\w+=,.@-]+$`)
|
||||
|
||||
// roleSessionNameRegex validates RoleSessionName per AWS spec.
|
||||
// Same character class as federation Name, but the length bounds differ
|
||||
// (RoleSessionName is 2..64).
|
||||
var roleSessionNameRegex = regexp.MustCompile(`^[\w+=,.@-]+$`)
|
||||
|
||||
const (
|
||||
minRoleSessionNameLen = 2
|
||||
maxRoleSessionNameLen = 64
|
||||
)
|
||||
|
||||
// validateRoleSessionName enforces the AWS RoleSessionName contract:
|
||||
// length 2..64, characters [\w+=,.@-]+. Returns the STS error code and a
|
||||
// descriptive error suitable for callers to surface to the caller.
|
||||
func validateRoleSessionName(name string) (STSErrorCode, error) {
|
||||
if name == "" {
|
||||
return STSErrMissingParameter, fmt.Errorf("RoleSessionName is required")
|
||||
}
|
||||
if len(name) < minRoleSessionNameLen || len(name) > maxRoleSessionNameLen {
|
||||
return STSErrInvalidParameterValue,
|
||||
fmt.Errorf("RoleSessionName must be between %d and %d characters", minRoleSessionNameLen, maxRoleSessionNameLen)
|
||||
}
|
||||
if !roleSessionNameRegex.MatchString(name) {
|
||||
return STSErrInvalidParameterValue,
|
||||
fmt.Errorf(`RoleSessionName contains invalid characters; allowed: [\w+=,.@-]`)
|
||||
}
|
||||
return "", nil
|
||||
}
|
||||
|
||||
// STS duration constants (AWS specification)
|
||||
const (
|
||||
minDurationSeconds = int64(900) // 15 minutes
|
||||
@@ -61,6 +89,27 @@ const (
|
||||
maxFederationDurationSeconds = int64(129600) // 36 hours (GetFederationToken max)
|
||||
)
|
||||
|
||||
// AWS limits inline session policies to 2048 characters for AssumeRole,
|
||||
// AssumeRoleWithWebIdentity, and AssumeRoleWithSAML. PackedPolicySize is
|
||||
// returned as a percentage of that budget so callers can detect how close
|
||||
// they are to the limit.
|
||||
const sessionPolicyBudgetBytes = 2048
|
||||
|
||||
// computePackedPolicySize returns the inline session policy size as a
|
||||
// percentage of the per-action budget, or nil when no session policy was
|
||||
// provided. Output is bounded to [0, 100] for AWS-compat reporting; the
|
||||
// actual policy size validation happens upstream in NormalizeSessionPolicy.
|
||||
func computePackedPolicySize(policyJSON string) *int64 {
|
||||
if policyJSON == "" {
|
||||
return nil
|
||||
}
|
||||
pct := int64(len(policyJSON)) * 100 / sessionPolicyBudgetBytes
|
||||
if pct > 100 {
|
||||
pct = 100
|
||||
}
|
||||
return &pct
|
||||
}
|
||||
|
||||
// parseDurationSecondsWithBounds parses and validates the DurationSeconds parameter
|
||||
// against the given min and max bounds. Returns nil if the parameter is not provided.
|
||||
func parseDurationSecondsWithBounds(r *http.Request, minSec, maxSec int64) (*int64, STSErrorCode, error) {
|
||||
@@ -170,9 +219,8 @@ func (h *STSHandlers) handleAssumeRoleWithWebIdentity(w http.ResponseWriter, r *
|
||||
return
|
||||
}
|
||||
|
||||
if roleSessionName == "" {
|
||||
h.writeSTSErrorResponse(w, r, STSErrMissingParameter,
|
||||
fmt.Errorf("RoleSessionName is required"))
|
||||
if errCode, err := validateRoleSessionName(roleSessionName); err != nil {
|
||||
h.writeSTSErrorResponse(w, r, errCode, err)
|
||||
return
|
||||
}
|
||||
|
||||
@@ -245,6 +293,7 @@ func (h *STSHandlers) handleAssumeRoleWithWebIdentity(w http.ResponseWriter, r *
|
||||
Expiration: response.Credentials.Expiration.Format(time.RFC3339),
|
||||
},
|
||||
SubjectFromWebIdentityToken: response.AssumedRoleUser.Subject,
|
||||
PackedPolicySize: computePackedPolicySize(sessionPolicyJSON),
|
||||
},
|
||||
}
|
||||
xmlResponse.ResponseMetadata.RequestId = request_id.GetFromRequest(r)
|
||||
@@ -264,9 +313,8 @@ func (h *STSHandlers) handleAssumeRole(w http.ResponseWriter, r *http.Request) {
|
||||
// Validate required parameters
|
||||
// RoleArn is optional to support S3-compatible clients that omit it
|
||||
|
||||
if roleSessionName == "" {
|
||||
h.writeSTSErrorResponse(w, r, STSErrMissingParameter,
|
||||
fmt.Errorf("RoleSessionName is required"))
|
||||
if errCode, err := validateRoleSessionName(roleSessionName); err != nil {
|
||||
h.writeSTSErrorResponse(w, r, errCode, err)
|
||||
return
|
||||
}
|
||||
|
||||
@@ -373,8 +421,9 @@ func (h *STSHandlers) handleAssumeRole(w http.ResponseWriter, r *http.Request) {
|
||||
// Build and return response
|
||||
xmlResponse := &AssumeRoleResponse{
|
||||
Result: AssumeRoleResult{
|
||||
Credentials: stsCreds,
|
||||
AssumedRoleUser: assumedUser,
|
||||
Credentials: stsCreds,
|
||||
AssumedRoleUser: assumedUser,
|
||||
PackedPolicySize: computePackedPolicySize(sessionPolicyJSON),
|
||||
},
|
||||
}
|
||||
xmlResponse.ResponseMetadata.RequestId = request_id.GetFromRequest(r)
|
||||
@@ -397,9 +446,8 @@ func (h *STSHandlers) handleAssumeRoleWithLDAPIdentity(w http.ResponseWriter, r
|
||||
return
|
||||
}
|
||||
|
||||
if roleSessionName == "" {
|
||||
h.writeSTSErrorResponse(w, r, STSErrMissingParameter,
|
||||
fmt.Errorf("RoleSessionName is required"))
|
||||
if errCode, err := validateRoleSessionName(roleSessionName); err != nil {
|
||||
h.writeSTSErrorResponse(w, r, errCode, err)
|
||||
return
|
||||
}
|
||||
|
||||
@@ -514,8 +562,9 @@ func (h *STSHandlers) handleAssumeRoleWithLDAPIdentity(w http.ResponseWriter, r
|
||||
// Build and return response
|
||||
xmlResponse := &AssumeRoleWithLDAPIdentityResponse{
|
||||
Result: LDAPIdentityResult{
|
||||
Credentials: stsCreds,
|
||||
AssumedRoleUser: assumedUser,
|
||||
Credentials: stsCreds,
|
||||
AssumedRoleUser: assumedUser,
|
||||
PackedPolicySize: computePackedPolicySize(sessionPolicyJSON),
|
||||
},
|
||||
}
|
||||
xmlResponse.ResponseMetadata.RequestId = request_id.GetFromRequest(r)
|
||||
@@ -906,6 +955,7 @@ type WebIdentityResult struct {
|
||||
Credentials STSCredentials `xml:"Credentials"`
|
||||
SubjectFromWebIdentityToken string `xml:"SubjectFromWebIdentityToken,omitempty"`
|
||||
AssumedRoleUser *AssumedRoleUser `xml:"AssumedRoleUser,omitempty"`
|
||||
PackedPolicySize *int64 `xml:"PackedPolicySize,omitempty"`
|
||||
}
|
||||
|
||||
// STSCredentials represents temporary security credentials
|
||||
@@ -933,8 +983,9 @@ type AssumeRoleResponse struct {
|
||||
|
||||
// AssumeRoleResult contains the result of AssumeRole
|
||||
type AssumeRoleResult struct {
|
||||
Credentials STSCredentials `xml:"Credentials"`
|
||||
AssumedRoleUser *AssumedRoleUser `xml:"AssumedRoleUser,omitempty"`
|
||||
Credentials STSCredentials `xml:"Credentials"`
|
||||
AssumedRoleUser *AssumedRoleUser `xml:"AssumedRoleUser,omitempty"`
|
||||
PackedPolicySize *int64 `xml:"PackedPolicySize,omitempty"`
|
||||
}
|
||||
|
||||
// AssumeRoleWithLDAPIdentityResponse is the response for AssumeRoleWithLDAPIdentity
|
||||
@@ -948,8 +999,9 @@ type AssumeRoleWithLDAPIdentityResponse struct {
|
||||
|
||||
// LDAPIdentityResult contains the result of AssumeRoleWithLDAPIdentity
|
||||
type LDAPIdentityResult struct {
|
||||
Credentials STSCredentials `xml:"Credentials"`
|
||||
AssumedRoleUser *AssumedRoleUser `xml:"AssumedRoleUser,omitempty"`
|
||||
Credentials STSCredentials `xml:"Credentials"`
|
||||
AssumedRoleUser *AssumedRoleUser `xml:"AssumedRoleUser,omitempty"`
|
||||
PackedPolicySize *int64 `xml:"PackedPolicySize,omitempty"`
|
||||
}
|
||||
|
||||
// GetCallerIdentityResponse is the response for GetCallerIdentity
|
||||
|
||||
@@ -0,0 +1,34 @@
|
||||
package s3api
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestComputePackedPolicySize(t *testing.T) {
|
||||
cases := []struct {
|
||||
name string
|
||||
policyLen int
|
||||
empty bool
|
||||
want int64
|
||||
}{
|
||||
{"empty -> nil", 0, true, 0},
|
||||
{"tiny policy -> 0%", 10, false, 0},
|
||||
{"half budget -> 50%", sessionPolicyBudgetBytes / 2, false, 50},
|
||||
{"full budget -> 100%", sessionPolicyBudgetBytes, false, 100},
|
||||
{"oversized -> capped at 100", sessionPolicyBudgetBytes * 3, false, 100},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
policy := repeat('a', tc.policyLen)
|
||||
got := computePackedPolicySize(policy)
|
||||
switch {
|
||||
case tc.empty:
|
||||
if got != nil {
|
||||
t.Fatalf("expected nil for empty input, got %d", *got)
|
||||
}
|
||||
case got == nil:
|
||||
t.Fatalf("expected non-nil result, got nil")
|
||||
case *got != tc.want:
|
||||
t.Fatalf("got=%d want=%d", *got, tc.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,47 @@
|
||||
package s3api
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestValidateRoleSessionName(t *testing.T) {
|
||||
cases := []struct {
|
||||
name string
|
||||
input string
|
||||
wantErr bool
|
||||
// wantCode is checked only when wantErr is true
|
||||
wantCode STSErrorCode
|
||||
}{
|
||||
{"empty rejected", "", true, STSErrMissingParameter},
|
||||
{"single char rejected (below min len 2)", "a", true, STSErrInvalidParameterValue},
|
||||
{"min length 2 accepted", "ab", false, ""},
|
||||
{"plain ascii accepted", "session-name_1", false, ""},
|
||||
{"all special chars allowed", "+=,.@-", false, ""},
|
||||
{"email-style accepted", "alice@example.com", false, ""},
|
||||
{"max length 64 accepted", string(make([]byte, 64)), true, STSErrInvalidParameterValue}, // zero bytes -> invalid charset
|
||||
{"max length 64 valid charset accepted", repeat('a', 64), false, ""},
|
||||
{"length 65 rejected", repeat('a', 65), true, STSErrInvalidParameterValue},
|
||||
{"space rejected", "alice bob", true, STSErrInvalidParameterValue},
|
||||
{"slash rejected", "alice/bob", true, STSErrInvalidParameterValue},
|
||||
{"colon rejected", "alice:bob", true, STSErrInvalidParameterValue},
|
||||
{"unicode rejected", "alicé", true, STSErrInvalidParameterValue},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
code, err := validateRoleSessionName(tc.input)
|
||||
gotErr := err != nil
|
||||
if gotErr != tc.wantErr {
|
||||
t.Fatalf("err mismatch: got=%v want=%v (err=%v)", gotErr, tc.wantErr, err)
|
||||
}
|
||||
if tc.wantErr && code != tc.wantCode {
|
||||
t.Fatalf("code mismatch: got=%s want=%s", code, tc.wantCode)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func repeat(b byte, n int) string {
|
||||
out := make([]byte, n)
|
||||
for i := range out {
|
||||
out[i] = b
|
||||
}
|
||||
return string(out)
|
||||
}
|
||||
Reference in New Issue
Block a user