diff --git a/internal/auth/auth.go b/internal/auth/auth.go index 624b6247..0e9d4af4 100644 --- a/internal/auth/auth.go +++ b/internal/auth/auth.go @@ -24,6 +24,7 @@ import ( "regexp" "sort" "strings" + "sync" "time" "github.com/opentracing/opentracing-go" @@ -54,6 +55,11 @@ type Client struct { config *config.Config io iostreams.IOStreamer fs afero.Fs + + // rotationAttempted tracks refresh tokens that were already used for token + // rotation during this process so each is attempted at most once + rotationAttempted map[string]bool + rotationAttemptedMu sync.Mutex } type AuthInterface interface { @@ -99,11 +105,12 @@ type AuthInterface interface { // NewClient returns a new, empty instance of the Client func NewClient(apiClient api.APIInterface, appClient *app.Client, config *config.Config, io iostreams.IOStreamer, fs afero.Fs) *Client { var client = Client{ - api: apiClient, - appClient: appClient, - config: config, - io: io, - fs: fs, + api: apiClient, + appClient: appClient, + config: config, + io: io, + fs: fs, + rotationAttempted: map[string]bool{}, } return &client @@ -307,7 +314,14 @@ func (c *Client) rotateTokenAll(ctx context.Context, auths types.AuthByTeamDomai // We also do not want to stop the entire process: so we will not return here. // The user should go ahead with the bad token and the api will handle the // return of the appropriate error to the user. - // We only want to warn the user about what we tried to do. + if slackerror.ToSlackError(err).Code == slackerror.ErrInvalidRefreshToken { + // The refresh token can never succeed again, so remove it to stop + // retrying token rotation on every command + auth.RefreshToken = "" + updatedAuths[authKey] = auth + updated = true + c.io.PrintWarning(ctx, "Your credentials for '%s' have expired and can no longer be refreshed. Run %s to authorize again.", auth.TeamDomain, style.Commandf("login", false)) + } c.io.PrintDebug(ctx, "Your auth token for '%s' is outdated. Tried refreshing the credentials but encountered the following error:\n%s", auth.TeamDomain, err.Error()) } else if tokenIsUpdated { updated = true @@ -325,10 +339,21 @@ func (c *Client) rotateToken(ctx context.Context, auth types.SlackAuth) (types.S return auth, false /* tokenIsUpdated */, nil } + // Attempt rotation with a refresh token at most once per process to avoid + // repeated requests when rotation fails + c.rotationAttemptedMu.Lock() + if c.rotationAttempted[auth.RefreshToken] { + c.rotationAttemptedMu.Unlock() + return auth, false /* tokenIsUpdated */, nil + } + c.rotationAttempted[auth.RefreshToken] = true + c.rotationAttemptedMu.Unlock() + // Store the current apiHost before rotation // We need this because we need to restore // the apiHost to what it was before rotating each of the user's auths activeAPIHostBeforeRotation := c.api.Host() + defer c.api.SetHost(activeAPIHostBeforeRotation) if auth.APIHost != nil { c.api.SetHost(*auth.APIHost) @@ -339,7 +364,6 @@ func (c *Client) rotateToken(ctx context.Context, auth types.SlackAuth) (types.S var result, err = c.api.RotateToken(ctx, auth) if err != nil { - // handle token rotation failure by sending meaningful messages to the users and remove already expired auth return auth, false /* tokenIsUpdated */, err } @@ -348,9 +372,6 @@ func (c *Client) rotateToken(ctx context.Context, auth types.SlackAuth) (types.S auth.RefreshToken = result.RefreshToken auth.LastUpdated = time.Now() - // now restore the previous default apiHost - c.api.SetHost(activeAPIHostBeforeRotation) - return auth, true /* tokenIsUpdated */, nil } diff --git a/internal/auth/auth_test.go b/internal/auth/auth_test.go index e0a4c70c..16aaa18e 100644 --- a/internal/auth/auth_test.go +++ b/internal/auth/auth_test.go @@ -32,6 +32,7 @@ import ( "github.com/slackapi/slack-cli/internal/slackerror" "github.com/spf13/afero" "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/mock" "github.com/stretchr/testify/require" ) @@ -426,10 +427,13 @@ func Test_AuthsRotation(t *testing.T) { require.Fail(t, "should not return error when saving auths") } + apiHostBefore := authClient.api.Host() // track the api host before we call User Auths updatedAuths, err := authClient.auths(ctx) // call the function + apiHostAfter := authClient.api.Host() // track the api host after the function runs. // Assertions require.Equal(t, workspaceAuthA, updatedAuths[authATeamID], "should return the same auth") + require.Equal(t, apiHostBefore, apiHostAfter, "api host before and after should be the same") // because the token rotation results in an error we expect the old stuff back require.Equal(t, workspaceAuthB, updatedAuths[authBTeamID], "should return the same auth") @@ -437,6 +441,67 @@ func Test_AuthsRotation(t *testing.T) { require.Equal(t, len(auths), len(updatedAuths), "we expect the same number of auths even if token rotation failed for one") require.NoError(t, err, "Should not return an error when the contents of credentials are valid") }) + + t.Run("token rotation with an invalid refresh token removes the refresh token", func(t *testing.T) { + ctx, authClient := setup(t) + fiveMinutesAgo := int(time.Now().Unix()) - 60*5 + apiMock := &api.APIMock{} + apiMock.AddDefaultMocks() + apiMock.On("SetHost", mock.Anything) + apiMock.On("RotateToken", mock.Anything, mock.Anything). + Return(api.RotateTokenResult{}, slackerror.NewAPIError(slackerror.ErrInvalidRefreshToken, "", nil, "tooling.tokens.rotate")) + authClient.api = apiMock + expiredAuth := types.SlackAuth{ + Token: "expiredToken", + RefreshToken: "dead-refresh-token", + ExpiresAt: fiveMinutesAgo, + TeamDomain: "workspace-a", + TeamID: "T123456789A", + } + _, err := authClient.setAuths(ctx, types.AuthByTeamDomain{expiredAuth.TeamID: expiredAuth}) + require.NoError(t, err) + + updatedAuths, err := authClient.auths(ctx) + require.NoError(t, err) + assert.Empty(t, updatedAuths[expiredAuth.TeamID].RefreshToken) + assert.Equal(t, expiredAuth.Token, updatedAuths[expiredAuth.TeamID].Token) + authClient.io.(*iostreams.IOStreamsMock).AssertCalled(t, "PrintWarning", mock.Anything, mock.Anything, mock.Anything) + + // A new process reads the saved credentials without attempting rotation + authClient.rotationAttempted = map[string]bool{} + updatedAuths, err = authClient.auths(ctx) + require.NoError(t, err) + assert.Empty(t, updatedAuths[expiredAuth.TeamID].RefreshToken) + apiMock.AssertNumberOfCalls(t, "RotateToken", 1) + }) + + t.Run("token rotation with a transient error is attempted once per process", func(t *testing.T) { + ctx, authClient := setup(t) + fiveMinutesAgo := int(time.Now().Unix()) - 60*5 + apiMock := &api.APIMock{} + apiMock.AddDefaultMocks() + apiMock.On("SetHost", mock.Anything) + apiMock.On("RotateToken", mock.Anything, mock.Anything). + Return(api.RotateTokenResult{}, slackerror.NewAPIError(slackerror.ErrInternal, "", nil, "tooling.tokens.rotate")) + authClient.api = apiMock + expiredAuth := types.SlackAuth{ + Token: "expiredToken", + RefreshToken: "valid-refresh-token", + ExpiresAt: fiveMinutesAgo, + TeamDomain: "workspace-a", + TeamID: "T123456789A", + } + _, err := authClient.setAuths(ctx, types.AuthByTeamDomain{expiredAuth.TeamID: expiredAuth}) + require.NoError(t, err) + + for range 3 { + updatedAuths, err := authClient.auths(ctx) + require.NoError(t, err) + assert.Equal(t, expiredAuth, updatedAuths[expiredAuth.TeamID]) + } + apiMock.AssertNumberOfCalls(t, "RotateToken", 1) + apiMock.AssertCalled(t, "SetHost", "https://slack.com") + }) } func Test_Auths(t *testing.T) {