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
3 changes: 3 additions & 0 deletions sdk/azidentity/CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,9 @@
### Other Changes
* `AzureCLICredential` imposes its default timeout only when the `Context`
passed to `GetToken()` has no deadline
* Added `NewCredentialUnavailableError()`. This function constructs an error indicating
a credential can't authenticate and an encompassing `ChainedTokenCredential` should
try its next credential, if any.

## 1.3.0-beta.1 (2022-12-13)

Expand Down
17 changes: 11 additions & 6 deletions sdk/azidentity/chained_token_credential.go
Original file line number Diff line number Diff line change
Expand Up @@ -81,10 +81,13 @@ func (c *ChainedTokenCredential) GetToken(ctx context.Context, opts policy.Token
}
}

var err error
var errs []error
var token azcore.AccessToken
var successfulCredential azcore.TokenCredential
var (
err error
errs []error
successfulCredential azcore.TokenCredential
token azcore.AccessToken
unavailableErr *credentialUnavailableError
)
for _, cred := range c.sources {
token, err = cred.GetToken(ctx, opts)
if err == nil {
Expand All @@ -93,12 +96,14 @@ func (c *ChainedTokenCredential) GetToken(ctx context.Context, opts policy.Token
break
}
errs = append(errs, err)
if _, ok := err.(*credentialUnavailableError); !ok {
// continue to the next source iff this one returned credentialUnavailableError
if !errors.As(err, &unavailableErr) {
break
}
}
if c.iterating {
c.cond.L.Lock()
// this is nil when all credentials returned an error
c.successfulCredential = successfulCredential
c.iterating = false
c.cond.L.Unlock()
Expand All @@ -108,7 +113,7 @@ func (c *ChainedTokenCredential) GetToken(ctx context.Context, opts policy.Token
if err != nil {
// return credentialUnavailableError iff all sources did so; return AuthenticationFailedError otherwise
msg := createChainedErrorMessage(errs)
if _, ok := err.(*credentialUnavailableError); ok {
if errors.As(err, &unavailableErr) {
err = newCredentialUnavailableError(c.name, msg)
} else {
res := getResponseFromError(err)
Expand Down
7 changes: 5 additions & 2 deletions sdk/azidentity/chained_token_credential_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -82,11 +82,14 @@ func TestChainedTokenCredential_NilSource(t *testing.T) {
}

func TestChainedTokenCredential_GetTokenSuccess(t *testing.T) {
// ChainedTokenCredential should continue iterating when a source returns credentialUnavailableError, wrapped or not
c1 := NewFakeCredential()
c1.SetResponse(azcore.AccessToken{}, newCredentialUnavailableError("test", "something went wrong"))
c2 := NewFakeCredential()
c2.SetResponse(azcore.AccessToken{Token: tokenValue, ExpiresOn: time.Now().Add(time.Hour)}, nil)
cred, err := NewChainedTokenCredential([]azcore.TokenCredential{c1, c2}, nil)
c2.SetResponse(azcore.AccessToken{}, fmt.Errorf("%w", newCredentialUnavailableError("...", "...")))
c3 := NewFakeCredential()
c3.SetResponse(azcore.AccessToken{Token: tokenValue, ExpiresOn: time.Now().Add(time.Hour)}, nil)
cred, err := NewChainedTokenCredential([]azcore.TokenCredential{c1, c2, c3}, nil)
if err != nil {
t.Fatal(err)
}
Expand Down
27 changes: 17 additions & 10 deletions sdk/azidentity/errors.go
Original file line number Diff line number Diff line change
Expand Up @@ -101,24 +101,31 @@ func (*AuthenticationFailedError) NonRetriable() {

var _ errorinfo.NonRetriable = (*AuthenticationFailedError)(nil)

// credentialUnavailableError indicates a credential can't attempt
// authentication because it lacks required data or state.
// credentialUnavailableError indicates a credential can't attempt authentication because it lacks required
// data or state
type credentialUnavailableError struct {
credType string
message string
message string
}

// newCredentialUnavailableError is an internal helper that ensures consistent error message formatting
func newCredentialUnavailableError(credType, message string) error {
return &credentialUnavailableError{credType: credType, message: message}
msg := fmt.Sprintf("%s: %s", credType, message)
return &credentialUnavailableError{msg}
}

func (e *credentialUnavailableError) Error() string {
return e.credType + ": " + e.message
// NewCredentialUnavailableError constructs an error indicating a credential can't attempt authentication
// because it lacks required data or state. When [ChainedTokenCredential] receives this error it will try
// its next credential, if any.
func NewCredentialUnavailableError(message string) error {
Comment thread
jhendrixMSFT marked this conversation as resolved.
return &credentialUnavailableError{message}
}

// NonRetriable indicates that this error should not be retried.
func (e *credentialUnavailableError) NonRetriable() {
// marker method
// Error implements the error interface. Note that the message contents are not contractual and can change over time.
func (e *credentialUnavailableError) Error() string {
return e.message
}

// NonRetriable is a marker method indicating this error should not be retried. It has no implementation.
func (e *credentialUnavailableError) NonRetriable() {}

var _ errorinfo.NonRetriable = (*credentialUnavailableError)(nil)