Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
25 changes: 19 additions & 6 deletions internal/cloudprovider/azure/tokencredential.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
}

Expand All @@ -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{}
Expand All @@ -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.
Expand All @@ -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)
}
26 changes: 26 additions & 0 deletions internal/cloudprovider/azure/tokencredential_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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) {
Expand Down Expand Up @@ -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))
}
}
45 changes: 39 additions & 6 deletions internal/store/credentialprovider/azure/register.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down Expand Up @@ -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{
Expand All @@ -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
Expand All @@ -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),
Expand All @@ -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
}
Expand Down Expand Up @@ -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)
Expand Down
77 changes: 77 additions & 0 deletions internal/store/credentialprovider/azure/register_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand Down Expand Up @@ -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
}
Loading