From a249d814a1acfd312810c6ec89f6c499e50f13cf Mon Sep 17 00:00:00 2001 From: Charles Wu Date: Wed, 5 Aug 2026 15:42:29 +1000 Subject: [PATCH 1/3] feat: log Azure credential and ACR token acquisition The Azure credential path was completely silent: internal/cloudprovider/azure and internal/store/credentialprovider/azure had no logging at all, and CreateCredentialChain swallowed both credential construction errors, so there was no way to tell which identity was used or why registry auth failed. Log which credential sources are available (workload identity, managed identity) and why one is skipped, the AAD token acquisition and the ACR refresh token exchange with their durations, and the resolved credential TTL. Failures to build the chain or exchange the token are now logged at error level, and a TTL parse fallback at warn level. Signed-off-by: Charles Wu --- internal/cloudprovider/azure/tokencredential.go | 9 +++++++++ .../store/credentialprovider/azure/register.go | 14 ++++++++++++++ 2 files changed, 23 insertions(+) diff --git a/internal/cloudprovider/azure/tokencredential.go b/internal/cloudprovider/azure/tokencredential.go index b0ca899a9..9188259d6 100644 --- a/internal/cloudprovider/azure/tokencredential.go +++ b/internal/cloudprovider/azure/tokencredential.go @@ -18,6 +18,7 @@ package azure 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 @@ -34,6 +35,9 @@ func CreateCredentialChain(clientID, tenantID string) (azcore.TokenCredential, e }) if err == nil { sources = append(sources, wiCred) + logrus.Debugf("azure: workload identity credential is available (clientID=%q, tenantID=%q)", clientID, tenantID) + } else { + logrus.Debugf("azure: workload identity credential is unavailable: %v", err) } // 2. Try Managed Identity second @@ -46,8 +50,13 @@ func CreateCredentialChain(clientID, tenantID string) (azcore.TokenCredential, e miCred, err := azidentity.NewManagedIdentityCredential(miOpts) if err == nil { sources = append(sources, miCred) + logrus.Debugf("azure: managed identity credential is available (clientID=%q)", clientID) + } else { + logrus.Debugf("azure: managed identity credential is unavailable: %v", err) } + logrus.Debugf("azure: built credential chain with %d source(s)", len(sources)) + // 3. Create chained credential return azidentity.NewChainedTokenCredential(sources, nil) } diff --git a/internal/store/credentialprovider/azure/register.go b/internal/store/credentialprovider/azure/register.go index fc38e1e1c..a1aa001a1 100644 --- a/internal/store/credentialprovider/azure/register.go +++ b/internal/store/credentialprovider/azure/register.go @@ -27,9 +27,12 @@ import ( "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} + const ( // GrantTypeAccessToken is the grant type for AAD access token GrantTypeAccessToken = "access_token" @@ -93,14 +96,18 @@ func createAzureIdentityProvider(opts credentialprovider.Options) (ratify.Regist func (p *IdentityProvider) GetWithTTL(ctx context.Context, serverAddress string) (credentialprovider.CredentialWithTTL, error) { // Step 1: Create a ChainedTokenCredential in the order: workload identity, // managed identity. + log := logger.GetLogger(ctx, logOpt) + log.Debugf("resolving ACR credential for %s (clientID=%q, tenantID=%q)", serverAddress, p.clientID, p.tenantID) chain, err := azure.CreateCredentialChain(p.clientID, p.tenantID) 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) } @@ -108,8 +115,10 @@ func (p *IdentityProvider) GetWithTTL(ctx context.Context, serverAddress string) ttl, err := parseJWTTokenTTL(acrRefreshToken) if err != nil { // If JWT parsing fails, fall back to the default TTL + log.Warnf("failed to parse ACR refresh token TTL for %s, falling back to the default: %v", serverAddress, err) ttl = DefaultACRTokenTTL } + log.Debugf("resolved ACR credential for %s, expires in %s", serverAddress, ttl) return credentialprovider.CredentialWithTTL{ Credential: ratify.RegistryCredential{ @@ -122,13 +131,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 @@ -138,6 +150,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), @@ -154,6 +167,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 } From 327e47124cf8aa947a4ded40f1519dd6c39503ae Mon Sep 17 00:00:00 2001 From: Charles Wu Date: Thu, 6 Aug 2026 12:10:10 +1000 Subject: [PATCH 2/3] fix: never cache an already-expired ACR refresh token parseJWTTokenTTL now returns a sentinel error for an expired token so the new tokenCacheTTL helper can tell it apart from an unparseable one. An expired token gets a zero TTL, which stops CachedProvider from storing it and serving 401s for hours; only a genuinely unparseable expiry falls back to the default TTL. Credential-chain logging moves into an appendCredential helper so both the available and unavailable paths are covered by tests. Signed-off-by: Charles Wu --- .../cloudprovider/azure/tokencredential.go | 28 +++++----- .../azure/tokencredential_test.go | 26 +++++++++ .../credentialprovider/azure/register.go | 31 ++++++++--- .../credentialprovider/azure/register_test.go | 53 +++++++++++++++++++ 4 files changed, 120 insertions(+), 18 deletions(-) diff --git a/internal/cloudprovider/azure/tokencredential.go b/internal/cloudprovider/azure/tokencredential.go index 9188259d6..9a322f6bb 100644 --- a/internal/cloudprovider/azure/tokencredential.go +++ b/internal/cloudprovider/azure/tokencredential.go @@ -16,6 +16,8 @@ limitations under the License. package azure import ( + "fmt" + "github.com/Azure/azure-sdk-for-go/sdk/azcore" "github.com/Azure/azure-sdk-for-go/sdk/azidentity" "github.com/sirupsen/logrus" @@ -33,12 +35,7 @@ func CreateCredentialChain(clientID, tenantID string) (azcore.TokenCredential, e ClientID: clientID, TenantID: tenantID, }) - if err == nil { - sources = append(sources, wiCred) - logrus.Debugf("azure: workload identity credential is available (clientID=%q, tenantID=%q)", clientID, tenantID) - } else { - logrus.Debugf("azure: workload identity credential is unavailable: %v", err) - } + sources = appendCredential(sources, wiCred, err, "workload identity", fmt.Sprintf("clientID=%q, tenantID=%q", clientID, tenantID)) // 2. Try Managed Identity second miOpts := &azidentity.ManagedIdentityCredentialOptions{} @@ -48,15 +45,22 @@ func CreateCredentialChain(clientID, tenantID string) (azcore.TokenCredential, e miOpts.ID = azidentity.ClientID(clientID) } miCred, err := azidentity.NewManagedIdentityCredential(miOpts) - if err == nil { - sources = append(sources, miCred) - logrus.Debugf("azure: managed identity credential is available (clientID=%q)", clientID) - } else { - logrus.Debugf("azure: managed identity credential is unavailable: %v", err) - } + sources = appendCredential(sources, miCred, err, "managed identity", fmt.Sprintf("clientID=%q", clientID)) logrus.Debugf("azure: built credential chain with %d source(s)", len(sources)) // 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 a1aa001a1..4df4f0a7b 100644 --- a/internal/store/credentialprovider/azure/register.go +++ b/internal/store/credentialprovider/azure/register.go @@ -18,6 +18,7 @@ package azure import ( "context" "encoding/json" + "errors" "fmt" "time" @@ -33,6 +34,25 @@ import ( 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") + +// tokenCacheTTL reports how long an ACR refresh token may be cached, along with +// any problem found while reading its expiry. An already-expired token gets a +// zero TTL so that it is never cached and reused, while a token whose expiry +// cannot be parsed falls back to the default TTL. +func tokenCacheTTL(token string) (time.Duration, error) { + ttl, err := parseJWTTokenTTL(token) + switch { + case errors.Is(err, errTokenExpired): + return 0, err + case err != nil: + return DefaultACRTokenTTL, err + default: + return ttl, nil + } +} + const ( // GrantTypeAccessToken is the grant type for AAD access token GrantTypeAccessToken = "access_token" @@ -112,13 +132,12 @@ func (p *IdentityProvider) GetWithTTL(ctx context.Context, serverAddress string) } // Step 3: Parse the JWT token to extract the actual TTL - ttl, err := parseJWTTokenTTL(acrRefreshToken) + ttl, err := tokenCacheTTL(acrRefreshToken) if err != nil { - // If JWT parsing fails, fall back to the default TTL - log.Warnf("failed to parse ACR refresh token TTL for %s, falling back to the default: %v", serverAddress, err) - ttl = DefaultACRTokenTTL + log.Warnf("could not determine the ACR refresh token TTL for %s, caching it for %s instead: %v", serverAddress, ttl, err) + } else { + log.Debugf("resolved ACR credential for %s, expires in %s", serverAddress, ttl) } - log.Debugf("resolved ACR credential for %s, expires in %s", serverAddress, ttl) return credentialprovider.CredentialWithTTL{ Credential: ratify.RegistryCredential{ @@ -200,7 +219,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 66d2d3d35..3ab8ff434 100644 --- a/internal/store/credentialprovider/azure/register_test.go +++ b/internal/store/credentialprovider/azure/register_test.go @@ -904,3 +904,56 @@ func contains(s, substr string) bool { } return false } + +func TestTokenCacheTTL(t *testing.T) { + tests := []struct { + name string + token string + wantTTL time.Duration + wantErr bool + }{ + { + name: "valid token keeps its own TTL", + token: createTestJWTToken(map[string]interface{}{"exp": time.Now().Add(time.Hour).Unix()}), + wantTTL: 55 * time.Minute, + }, + { + name: "unparseable token falls back to the default TTL", + token: "invalid.jwt.token", + wantTTL: DefaultACRTokenTTL, + wantErr: true, + }, + { + name: "expired token is never cached", + token: createTestJWTToken(map[string]interface{}{"exp": time.Now().Add(-time.Hour).Unix()}), + wantTTL: 0, + wantErr: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + ttl, err := tokenCacheTTL(tt.token) + if (err != nil) != tt.wantErr { + t.Fatalf("tokenCacheTTL() error = %v, wantErr %v", err, tt.wantErr) + } + // 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("tokenCacheTTL() = %s, want ~%s", ttl, tt.wantTTL) + } + }) + } +} + +func TestTokenCacheTTL_ExpiredTokenIsNotCacheable(t *testing.T) { + token := createTestJWTToken(map[string]interface{}{"exp": time.Now().Add(-time.Hour).Unix()}) + + ttl, err := tokenCacheTTL(token) + if !errors.Is(err, errTokenExpired) { + t.Fatalf("tokenCacheTTL() error = %v, want errTokenExpired", err) + } + // A zero TTL keeps CachedProvider from storing an unusable token. + if ttl != 0 { + t.Errorf("tokenCacheTTL() = %s, want 0 so the token is not cached", ttl) + } +} From 1f3925c2a429b9da1db922b5977f2a5a7082bad8 Mon Sep 17 00:00:00 2001 From: Charles Wu Date: Fri, 7 Aug 2026 13:09:10 +1000 Subject: [PATCH 3/3] refactor: make the ACR token TTL decision explicit at one place The default-TTL fallback had moved out of GetWithTTL into a helper that only returned (ttl, err), so the caller logged one message for both failure modes: an expired token reported "caching it for 0s", which is not what happens - it is not cached at all. resolveTokenTTL now owns both the TTL and the explanation, naming DefaultACRTokenTTL where it is applied and logging the expired and unparseable cases distinctly. A test asserts the two messages differ. Signed-off-by: Charles Wu --- .../credentialprovider/azure/register.go | 28 ++++----- .../credentialprovider/azure/register_test.go | 58 +++++++++++++------ 2 files changed, 55 insertions(+), 31 deletions(-) diff --git a/internal/store/credentialprovider/azure/register.go b/internal/store/credentialprovider/azure/register.go index 9689e1067..f84157cf7 100644 --- a/internal/store/credentialprovider/azure/register.go +++ b/internal/store/credentialprovider/azure/register.go @@ -25,6 +25,7 @@ import ( "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" @@ -37,19 +38,23 @@ 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") -// tokenCacheTTL reports how long an ACR refresh token may be cached, along with -// any problem found while reading its expiry. An already-expired token gets a -// zero TTL so that it is never cached and reused, while a token whose expiry -// cannot be parsed falls back to the default TTL. -func tokenCacheTTL(token string) (time.Duration, error) { +// 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): - return 0, err + log.Warnf("ACR returned an already-expired refresh token for %s; it will not be cached", serverAddress) + return 0 case err != nil: - return DefaultACRTokenTTL, err + log.Warnf("failed to parse the ACR refresh token TTL for %s, falling back to %s: %v", serverAddress, DefaultACRTokenTTL, err) + return DefaultACRTokenTTL default: - return ttl, nil + log.Debugf("resolved ACR credential for %s, expires in %s", serverAddress, ttl) + return ttl } } @@ -156,12 +161,7 @@ func (p *IdentityProvider) GetWithTTL(ctx context.Context, serverAddress string) } // Step 3: Parse the JWT token to extract the actual TTL - ttl, err := tokenCacheTTL(acrRefreshToken) - if err != nil { - log.Warnf("could not determine the ACR refresh token TTL for %s, caching it for %s instead: %v", serverAddress, ttl, err) - } else { - log.Debugf("resolved ACR credential for %s, expires in %s", serverAddress, ttl) - } + ttl := resolveTokenTTL(log, serverAddress, acrRefreshToken) return credentialprovider.CredentialWithTTL{ Credential: ratify.RegistryCredential{ diff --git a/internal/store/credentialprovider/azure/register_test.go b/internal/store/credentialprovider/azure/register_test.go index 5cfda46fa..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 ( @@ -974,55 +977,76 @@ func contains(s, substr string) bool { return false } -func TestTokenCacheTTL(t *testing.T) { +func TestResolveTokenTTL(t *testing.T) { tests := []struct { name string token string wantTTL time.Duration - wantErr bool + 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, - wantErr: true, + wantLog: "falling back to", }, { name: "expired token is never cached", token: createTestJWTToken(map[string]interface{}{"exp": time.Now().Add(-time.Hour).Unix()}), wantTTL: 0, - wantErr: true, + wantLog: "already-expired refresh token", }, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - ttl, err := tokenCacheTTL(tt.token) - if (err != nil) != tt.wantErr { - t.Fatalf("tokenCacheTTL() error = %v, wantErr %v", err, tt.wantErr) - } + 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("tokenCacheTTL() = %s, want ~%s", ttl, tt.wantTTL) + 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) } }) } } -func TestTokenCacheTTL_ExpiredTokenIsNotCacheable(t *testing.T) { - token := createTestJWTToken(map[string]interface{}{"exp": time.Now().Add(-time.Hour).Unix()}) +// 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()})) - ttl, err := tokenCacheTTL(token) - if !errors.Is(err, errTokenExpired) { - t.Fatalf("tokenCacheTTL() error = %v, want errTokenExpired", err) + 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) } - // A zero TTL keeps CachedProvider from storing an unusable token. - if ttl != 0 { - t.Errorf("tokenCacheTTL() = %s, want 0 so the token is not cached", ttl) + 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 +}