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
88 changes: 69 additions & 19 deletions base/oauth.go
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@ import (
"io"
"net/http"
"strings"
"sync"
"time"

"github.com/keybase/go-keybase-chat-bot/kbchat"
Expand All @@ -22,17 +23,26 @@ func (e OAuthRequiredError) Error() string {
return "OAuth is required for this, permission requested."
}

// ShouldRetryAuth checks if an error indicates OAuth credentials have failed
// and should be deleted to trigger re-authentication. This consolidates the
// retry logic used across meetbot, zoombot, and gcalbot.
// ShouldRetryAuth reports whether err means the user's OAuth credentials are
// permanently unusable and should be deleted. Transient token-fetch failures
// (network, 5xx) are not treated as credential errors.
func ShouldRetryAuth(err error) bool {
if err == nil {
return false
}
errMsg := err.Error()
return strings.Contains(errMsg, "cannot fetch token") ||
strings.Contains(errMsg, "invalid_grant") ||
strings.Contains(errMsg, "token expired and refresh token is not set")
var retr *oauth2.RetrieveError
if errors.As(err, &retr) {
switch strings.ToLower(retr.ErrorCode) {
case "invalid_grant", "invalid_token":
return true
}
body := string(retr.Body)
return strings.Contains(body, "invalid_grant") ||
strings.Contains(body, "invalid_token")
}
msg := err.Error()
return strings.Contains(msg, "invalid_grant") ||
strings.Contains(msg, "token expired and refresh token is not set")
}

type OAuthStorage interface {
Expand Down Expand Up @@ -287,18 +297,58 @@ func GetOAuthClient(

return nil, OAuthRequiredError{}
}
// renew token
if token.Expiry.Before(time.Now()) {
newToken, err := config.TokenSource(ctx, token).Token()
if err != nil {
return nil, fmt.Errorf("unable to renew token: %s", err)
}
err = storage.PutToken(ctx, tokenIdentifier, newToken)
if err != nil {
return nil, fmt.Errorf("unable to update token: %s", err)
}
token = newToken

src := PersistTokenSource(ctx, token, ConfigTokenSource(ctx, config, token), func(ctx context.Context, tok *oauth2.Token) error {
return storage.PutToken(ctx, tokenIdentifier, tok)
})
if _, err := src.Token(); err != nil {
return nil, fmt.Errorf("unable to renew token: %w", err)
}
return oauth2.NewClient(ctx, src), nil
}

// ConfigTokenSource is config.TokenSource, except tokens with a zero Expiry
// and a refresh token are treated as expired. oauth2.Token.Valid treats a
// zero Expiry as never-expired, which would skip refresh forever.
func ConfigTokenSource(ctx context.Context, config *oauth2.Config, token *oauth2.Token) oauth2.TokenSource {
if token != nil && token.Expiry.IsZero() && token.RefreshToken != "" {
cp := *token
cp.Expiry = time.Now().Add(-time.Minute)
token = &cp
}
return config.TokenSource(ctx, token)
}

// PersistTokenSource wraps src and writes the token whenever AccessToken,
// RefreshToken, or Expiry changes (including refresh-token rotation).
func PersistTokenSource(ctx context.Context, token *oauth2.Token, src oauth2.TokenSource, put func(context.Context, *oauth2.Token) error) oauth2.TokenSource {
// Token() is invoked on later HTTP refreshes; the caller context may
// already be done by then, so persist independently of it.
return &persistTokenSource{ctx: context.WithoutCancel(ctx), token: token, src: src, put: put}
}

type persistTokenSource struct {
ctx context.Context
token *oauth2.Token
src oauth2.TokenSource
put func(context.Context, *oauth2.Token) error
mu sync.Mutex
}

return config.Client(ctx, token), nil
func (s *persistTokenSource) Token() (*oauth2.Token, error) {
s.mu.Lock()
defer s.mu.Unlock()
tok, err := s.src.Token()
if err != nil {
return nil, err
}
if tok.AccessToken != s.token.AccessToken ||
tok.RefreshToken != s.token.RefreshToken ||
!tok.Expiry.Equal(s.token.Expiry) {
*s.token = *tok
if err := s.put(s.ctx, tok); err != nil {
return nil, fmt.Errorf("unable to update token: %w", err)
}
}
return tok, nil
}
160 changes: 160 additions & 0 deletions base/oauth_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,160 @@
package base

import (
"context"
"errors"
"fmt"
"sync"
"testing"
"time"

"github.com/stretchr/testify/require"
"golang.org/x/oauth2"
)

func TestShouldRetryAuth(t *testing.T) {
t.Parallel()

t.Run("nil", func(t *testing.T) {
require.False(t, ShouldRetryAuth(nil))
})

t.Run("invalid_grant RetrieveError", func(t *testing.T) {
err := &oauth2.RetrieveError{ErrorCode: "invalid_grant", Body: []byte(`{"error":"invalid_grant"}`)}
require.True(t, ShouldRetryAuth(err))
require.True(t, ShouldRetryAuth(fmt.Errorf("unable to renew token: %w", err)))
})

t.Run("invalid_grant in body only", func(t *testing.T) {
err := &oauth2.RetrieveError{Body: []byte(`{"error":"invalid_grant"}`)}
require.True(t, ShouldRetryAuth(err))
})

t.Run("missing refresh token", func(t *testing.T) {
require.True(t, ShouldRetryAuth(errors.New("oauth2: token expired and refresh token is not set")))
})

t.Run("transient cannot fetch token", func(t *testing.T) {
err := &oauth2.RetrieveError{Body: []byte("connection reset"), ErrorCode: ""}
require.False(t, ShouldRetryAuth(err))
require.False(t, ShouldRetryAuth(errors.New("oauth2: cannot fetch token: 500 Internal Server Error")))
})

t.Run("unrelated", func(t *testing.T) {
require.False(t, ShouldRetryAuth(errors.New("calendar: 404 not found")))
})
}

type stubTokenSource struct {
tok *oauth2.Token
err error
}

func (s stubTokenSource) Token() (*oauth2.Token, error) {
return s.tok, s.err
}

func TestPersistTokenSource(t *testing.T) {
t.Parallel()
baseTok := oauth2.Token{
AccessToken: "old-access",
RefreshToken: "old-refresh",
Expiry: time.Now().Add(time.Hour),
}

t.Run("no write when unchanged", func(t *testing.T) {
orig := baseTok
var puts int
src := PersistTokenSource(context.Background(), &orig, stubTokenSource{tok: &orig},
func(context.Context, *oauth2.Token) error {
puts++
return nil
})
tok, err := src.Token()
require.NoError(t, err)
require.Equal(t, "old-access", tok.AccessToken)
require.Equal(t, 0, puts)
})

t.Run("writes on refresh token rotation", func(t *testing.T) {
stored := baseTok
rotated := &oauth2.Token{
AccessToken: "new-access",
RefreshToken: "new-refresh",
Expiry: time.Now().Add(2 * time.Hour),
}
var got *oauth2.Token
src := PersistTokenSource(context.Background(), &stored, stubTokenSource{tok: rotated},
func(_ context.Context, tok *oauth2.Token) error {
got = tok
return nil
})
tok, err := src.Token()
require.NoError(t, err)
require.Equal(t, "new-refresh", tok.RefreshToken)
require.Equal(t, "new-refresh", stored.RefreshToken)
require.Equal(t, "new-refresh", got.RefreshToken)
})

t.Run("put uses uncancelled context", func(t *testing.T) {
orig := baseTok
ctx, cancel := context.WithCancel(context.Background())
cancel()
rotated := &oauth2.Token{
AccessToken: "new-access",
RefreshToken: "new-refresh",
Expiry: time.Now().Add(time.Hour),
}
src := PersistTokenSource(ctx, &orig, stubTokenSource{tok: rotated},
func(putCtx context.Context, _ *oauth2.Token) error {
require.NoError(t, putCtx.Err())
return nil
})
_, err := src.Token()
require.NoError(t, err)
})

t.Run("put error", func(t *testing.T) {
orig := baseTok
rotated := &oauth2.Token{
AccessToken: "new-access",
RefreshToken: "old-refresh",
Expiry: time.Now().Add(time.Hour),
}
src := PersistTokenSource(context.Background(), &orig, stubTokenSource{tok: rotated},
func(context.Context, *oauth2.Token) error {
return errors.New("db down")
})
_, err := src.Token()
require.ErrorContains(t, err, "unable to update token")
require.ErrorContains(t, err, "db down")
})
}

func TestConfigTokenSourceZeroExpiry(t *testing.T) {
t.Parallel()
token := &oauth2.Token{AccessToken: "a", RefreshToken: "r"}
require.True(t, token.Valid(), "zero expiry is Valid() in oauth2")

//nolint:gosec // G101: False positive - TokenURL is a dummy loopback address, not credentials
cfg := &oauth2.Config{Endpoint: oauth2.Endpoint{TokenURL: "http://127.0.0.1:1"}}
src := ConfigTokenSource(context.Background(), cfg, token)
_, err := src.Token()
require.Error(t, err, "zero-expiry token with refresh token must be refreshed, not reused")
}

func TestPersistTokenSourceConcurrent(t *testing.T) {
t.Parallel()
token := &oauth2.Token{AccessToken: "a", RefreshToken: "r", Expiry: time.Now().Add(time.Hour)}
rotated := &oauth2.Token{AccessToken: "b", RefreshToken: "r2", Expiry: time.Now().Add(2 * time.Hour)}
src := PersistTokenSource(context.Background(), token, stubTokenSource{tok: rotated},
func(context.Context, *oauth2.Token) error { return nil })
var wg sync.WaitGroup
for range 8 {
wg.Go(func() {
_, err := src.Token()
require.NoError(t, err)
})
}
wg.Wait()
}
Loading
Loading