diff --git a/internal/cloudprovider/azure/tokencredential.go b/internal/cloudprovider/azure/tokencredential.go index e4324f0f6..323185e3a 100644 --- a/internal/cloudprovider/azure/tokencredential.go +++ b/internal/cloudprovider/azure/tokencredential.go @@ -20,6 +20,7 @@ import ( "github.com/Azure/azure-sdk-for-go/sdk/azcore" "github.com/Azure/azure-sdk-for-go/sdk/azidentity" + "github.com/sirupsen/logrus" ) // CreateCredentialChain creates a ChainedTokenCredential with the specified @@ -36,8 +37,10 @@ func CreateCredentialChainWithIdentityBinding(clientID, tenantID string, ibConfi if ibConfig != nil && ibConfig.SNIName != "" { ibCred, err := newIdentityBindingCredential(clientID, *ibConfig) if err != nil { + logrus.Debugf("azure: identity binding credential is unavailable (clientID=%q, sniName=%q): %v", clientID, ibConfig.SNIName, err) return nil, fmt.Errorf("failed to create identity binding credential: %w", err) } + logrus.Debugf("azure: using the identity binding credential exclusively (clientID=%q, sniName=%q)", clientID, ibConfig.SNIName) return ibCred, nil } @@ -49,9 +52,7 @@ func CreateCredentialChainWithIdentityBinding(clientID, tenantID string, ibConfi ClientID: clientID, TenantID: tenantID, }) - if err == nil { - sources = append(sources, wiCred) - } + sources = appendCredential(sources, wiCred, err, "workload identity", fmt.Sprintf("clientID=%q, tenantID=%q", clientID, tenantID)) // 2. Try Managed Identity second miOpts := &azidentity.ManagedIdentityCredentialOptions{} @@ -61,9 +62,9 @@ func CreateCredentialChainWithIdentityBinding(clientID, tenantID string, ibConfi miOpts.ID = azidentity.ClientID(clientID) } miCred, err := azidentity.NewManagedIdentityCredential(miOpts) - if err == nil { - sources = append(sources, miCred) - } + sources = appendCredential(sources, miCred, err, "managed identity", fmt.Sprintf("clientID=%q", clientID)) + + logrus.Debugf("azure: built credential chain with %d source(s)", len(sources)) // Fail clearly when no credential source could be constructed rather than // deferring to a less obvious error from the chained credential. @@ -74,3 +75,15 @@ func CreateCredentialChainWithIdentityBinding(clientID, tenantID string, ibConfi // 3. Create chained credential return azidentity.NewChainedTokenCredential(sources, nil) } + +// appendCredential adds cred to sources when it was created successfully. Both +// outcomes are logged so that operators can tell which identity sources the +// pod actually has available. +func appendCredential(sources []azcore.TokenCredential, cred azcore.TokenCredential, err error, name, details string) []azcore.TokenCredential { + if err != nil { + logrus.Debugf("azure: %s credential is unavailable (%s): %v", name, details, err) + return sources + } + logrus.Debugf("azure: %s credential is available (%s)", name, details) + return append(sources, cred) +} diff --git a/internal/cloudprovider/azure/tokencredential_test.go b/internal/cloudprovider/azure/tokencredential_test.go index a3b3674aa..7dadd1b92 100644 --- a/internal/cloudprovider/azure/tokencredential_test.go +++ b/internal/cloudprovider/azure/tokencredential_test.go @@ -16,7 +16,12 @@ limitations under the License. package azure import ( + "context" + "errors" "testing" + + "github.com/Azure/azure-sdk-for-go/sdk/azcore" + "github.com/Azure/azure-sdk-for-go/sdk/azcore/policy" ) func TestCreateCredentialChain(t *testing.T) { @@ -285,3 +290,24 @@ func TestCreateCredentialChain_AllPaths(t *testing.T) { }) } } + +// stubCredential is a placeholder credential used to verify chain assembly. +type stubCredential struct{} + +func (stubCredential) GetToken(_ context.Context, _ policy.TokenRequestOptions) (azcore.AccessToken, error) { + return azcore.AccessToken{}, nil +} + +func TestAppendCredential(t *testing.T) { + cred := stubCredential{} + + got := appendCredential(nil, cred, nil, "workload identity", `clientID="id"`) + if len(got) != 1 { + t.Fatalf("expected a successfully created credential to be appended, got %d source(s)", len(got)) + } + + got = appendCredential(got, nil, errors.New("unavailable"), "managed identity", `clientID=""`) + if len(got) != 1 { + t.Fatalf("expected a failed credential to be skipped, got %d source(s)", len(got)) + } +} diff --git a/internal/store/credentialprovider/azure/register.go b/internal/store/credentialprovider/azure/register.go index 59502e9a5..f84157cf7 100644 --- a/internal/store/credentialprovider/azure/register.go +++ b/internal/store/credentialprovider/azure/register.go @@ -18,18 +18,46 @@ package azure import ( "context" "encoding/json" + "errors" "fmt" "time" "github.com/Azure/azure-sdk-for-go/sdk/azcore" "github.com/Azure/azure-sdk-for-go/sdk/azcore/policy" azcontainerregistry "github.com/Azure/azure-sdk-for-go/sdk/containers/azcontainerregistry" + dcontext "github.com/docker/distribution/context" "github.com/golang-jwt/jwt/v5" "github.com/notaryproject/ratify-go" "github.com/notaryproject/ratify/v2/internal/cloudprovider/azure" + "github.com/notaryproject/ratify/v2/internal/logger" "github.com/notaryproject/ratify/v2/internal/store/credentialprovider" ) +var logOpt = logger.Option{ComponentType: logger.AuthProvider} + +// errTokenExpired is returned when the ACR refresh token is already expired. +var errTokenExpired = errors.New("JWT token has already expired") + +// resolveTokenTTL reports how long an ACR refresh token may be cached and +// records why. A token whose expiry cannot be parsed falls back to +// DefaultACRTokenTTL, but an already-expired one gets a zero TTL: CachedProvider +// only caches a positive TTL, so the token is discarded instead of being replayed +// for hours and returning 401 on every request. +func resolveTokenTTL(log dcontext.Logger, serverAddress, token string) time.Duration { + ttl, err := parseJWTTokenTTL(token) + switch { + case errors.Is(err, errTokenExpired): + log.Warnf("ACR returned an already-expired refresh token for %s; it will not be cached", serverAddress) + return 0 + case err != nil: + log.Warnf("failed to parse the ACR refresh token TTL for %s, falling back to %s: %v", serverAddress, DefaultACRTokenTTL, err) + return DefaultACRTokenTTL + default: + log.Debugf("resolved ACR credential for %s, expires in %s", serverAddress, ttl) + return ttl + } +} + const ( // GrantTypeAccessToken is the grant type for AAD access token GrantTypeAccessToken = "access_token" @@ -117,23 +145,23 @@ func (p *IdentityProvider) GetWithTTL(ctx context.Context, serverAddress string) // Step 1: Create the Azure token credential. When identity binding is // configured it is used exclusively; otherwise the chain is workload // identity followed by managed identity. + log := logger.GetLogger(ctx, logOpt) + log.Debugf("resolving ACR credential for %s (clientID=%q, tenantID=%q, identityBinding=%t)", serverAddress, p.clientID, p.tenantID, p.ibConfig != nil) chain, err := azure.CreateCredentialChainWithIdentityBinding(p.clientID, p.tenantID, p.ibConfig) if err != nil { + log.Errorf("failed to create Azure credential chain for %s: %v", serverAddress, err) return credentialprovider.CredentialWithTTL{}, fmt.Errorf("failed to create credential chain: %w", err) } // Step 2: Exchange an AAD token for an ACR refresh token using ExchangeAADAccessTokenForACRRefreshToken acrRefreshToken, err := p.exchangeAADTokenForACRToken(ctx, chain, serverAddress) if err != nil { + log.Errorf("failed to exchange AAD token for ACR refresh token for %s: %v", serverAddress, err) return credentialprovider.CredentialWithTTL{}, fmt.Errorf("failed to exchange AAD token for ACR refresh token: %w", err) } // Step 3: Parse the JWT token to extract the actual TTL - ttl, err := parseJWTTokenTTL(acrRefreshToken) - if err != nil { - // If JWT parsing fails, fall back to the default TTL - ttl = DefaultACRTokenTTL - } + ttl := resolveTokenTTL(log, serverAddress, acrRefreshToken) return credentialprovider.CredentialWithTTL{ Credential: ratify.RegistryCredential{ @@ -146,13 +174,16 @@ func (p *IdentityProvider) GetWithTTL(ctx context.Context, serverAddress string) // exchangeAADTokenForACRToken exchanges an AAD access token for an ACR refresh // token. func (p *IdentityProvider) exchangeAADTokenForACRToken(ctx context.Context, credential azcore.TokenCredential, serverAddress string) (string, error) { + log := logger.GetLogger(ctx, logOpt) // Get an AAD access token + start := time.Now() token, err := credential.GetToken(ctx, policy.TokenRequestOptions{ Scopes: []string{AADResource}, }) if err != nil { return "", fmt.Errorf("failed to get AAD access token: %w", err) } + log.Debugf("acquired AAD access token for scope %s in %dms", AADResource, time.Since(start).Milliseconds()) // Create ACR authentication client serverURL := "https://" + serverAddress @@ -162,6 +193,7 @@ func (p *IdentityProvider) exchangeAADTokenForACRToken(ctx context.Context, cred } // Exchange AAD token for ACR refresh token + exchangeStart := time.Now() response, err := client.ExchangeAADAccessTokenForACRRefreshToken( ctx, azcontainerregistry.PostContentSchemaGrantType(GrantTypeAccessToken), @@ -178,6 +210,7 @@ func (p *IdentityProvider) exchangeAADTokenForACRToken(ctx context.Context, cred if response.RefreshToken == nil { return "", fmt.Errorf("received nil refresh token from ACR") } + log.Debugf("exchanged AAD token for an ACR refresh token at %s in %dms", serverAddress, time.Since(exchangeStart).Milliseconds()) return *response.RefreshToken, nil } @@ -210,7 +243,7 @@ func parseJWTTokenTTL(token string) (time.Duration, error) { // If token is already expired, return 0 TTL if expTime.Before(now) { - return 0, fmt.Errorf("JWT token has already expired") + return 0, errTokenExpired } // Calculate TTL with a small buffer (subtract 1 minute for safety) diff --git a/internal/store/credentialprovider/azure/register_test.go b/internal/store/credentialprovider/azure/register_test.go index e1734c5e2..aca47c15b 100644 --- a/internal/store/credentialprovider/azure/register_test.go +++ b/internal/store/credentialprovider/azure/register_test.go @@ -25,8 +25,11 @@ import ( "github.com/Azure/azure-sdk-for-go/sdk/azcore" "github.com/Azure/azure-sdk-for-go/sdk/azcore/policy" + dcontext "github.com/docker/distribution/context" "github.com/notaryproject/ratify/v2/internal/cloudprovider/azure" "github.com/notaryproject/ratify/v2/internal/store/credentialprovider" + "github.com/sirupsen/logrus" + "github.com/sirupsen/logrus/hooks/test" ) const ( @@ -973,3 +976,77 @@ func contains(s, substr string) bool { } return false } + +func TestResolveTokenTTL(t *testing.T) { + tests := []struct { + name string + token string + wantTTL time.Duration + wantLog string + }{ + { + name: "valid token keeps its own TTL", + token: createTestJWTToken(map[string]interface{}{"exp": time.Now().Add(time.Hour).Unix()}), + wantTTL: 55 * time.Minute, + wantLog: "resolved ACR credential", + }, + { + name: "unparseable token falls back to the default TTL", + token: "invalid.jwt.token", + wantTTL: DefaultACRTokenTTL, + wantLog: "falling back to", + }, + { + name: "expired token is never cached", + token: createTestJWTToken(map[string]interface{}{"exp": time.Now().Add(-time.Hour).Unix()}), + wantTTL: 0, + wantLog: "already-expired refresh token", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + log, hook := newTestLogger() + ttl := resolveTokenTTL(log, testRegistry, tt.token) + + // Allow a small delta because the TTL is computed from time.Now(). + if delta := ttl - tt.wantTTL; delta > time.Minute || delta < -time.Minute { + t.Errorf("resolveTokenTTL() = %s, want ~%s", ttl, tt.wantTTL) + } + entry := hook.LastEntry() + if entry == nil { + t.Fatalf("resolveTokenTTL() logged nothing, want a message containing %q", tt.wantLog) + } + if !contains(entry.Message, tt.wantLog) { + t.Errorf("resolveTokenTTL() logged %q, want it to contain %q", entry.Message, tt.wantLog) + } + }) + } +} + +// An expired token and an unreadable one must not be reported identically: only +// the expired one is dropped from the cache. +func TestResolveTokenTTL_ExpiredAndUnparseableDiffer(t *testing.T) { + expiredLog, expiredHook := newTestLogger() + expiredTTL := resolveTokenTTL(expiredLog, testRegistry, createTestJWTToken(map[string]interface{}{"exp": time.Now().Add(-time.Hour).Unix()})) + + unparseableLog, unparseableHook := newTestLogger() + unparseableTTL := resolveTokenTTL(unparseableLog, testRegistry, "invalid.jwt.token") + + if expiredTTL != 0 { + t.Errorf("expired token TTL = %s, want 0 so it is not cached", expiredTTL) + } + if unparseableTTL != DefaultACRTokenTTL { + t.Errorf("unparseable token TTL = %s, want %s", unparseableTTL, DefaultACRTokenTTL) + } + if expiredHook.LastEntry().Message == unparseableHook.LastEntry().Message { + t.Errorf("expired and unparseable tokens logged the same message: %q", expiredHook.LastEntry().Message) + } +} + +func newTestLogger() (dcontext.Logger, *test.Hook) { + base := logrus.New() + base.SetLevel(logrus.DebugLevel) + hook := test.NewLocal(base) + return logrus.NewEntry(base), hook +}