Skip to content

Commit fb5d8f5

Browse files
fix(auth): retry stale scoped token once after 403
1 parent 213ddb6 commit fb5d8f5

4 files changed

Lines changed: 357 additions & 9 deletions

File tree

internal/googleapi/client.go

Lines changed: 10 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -148,10 +148,18 @@ func authenticatedTransportWithStoredScopeCheck(
148148
}
149149
}
150150

151-
return readOnlyTransportFromContext(ctx, NewRetryTransport(&oauth2.Transport{
151+
retryTransport := NewRetryTransport(&oauth2.Transport{
152152
Source: ts,
153153
Base: newBaseTransport(),
154-
})), nil
154+
})
155+
156+
if refresher, ok := ts.(interface {
157+
ForceRefresh(context.Context) error
158+
}); ok {
159+
retryTransport.RefreshAuth = refresher.ForceRefresh
160+
}
161+
162+
return readOnlyTransportFromContext(ctx, retryTransport), nil
155163
}
156164

157165
func optionsForAccountScopes(ctx context.Context, serviceLabel string, email string, scopes []string) ([]option.ClientOption, error) {

internal/googleapi/client_auth.go

Lines changed: 100 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -35,10 +35,78 @@ type persistingTokenSource struct {
3535
tok secrets.Token
3636
}
3737

38+
type forceRefreshTokenSource interface {
39+
ForceRefresh(context.Context) (*oauth2.Token, error)
40+
}
41+
42+
var (
43+
errBaseTokenSourceCannotForceRefresh = errors.New("base token source cannot force refresh")
44+
errBaseTokenSourceReturnedNilToken = errors.New("base token source returned nil token")
45+
)
46+
47+
type resettableOAuthTokenSource struct {
48+
mu sync.Mutex
49+
source oauth2.TokenSource
50+
newSource func(*oauth2.Token) oauth2.TokenSource
51+
refreshToken string
52+
}
53+
3854
type tokenAliasDeleter interface {
3955
DeleteTokenAlias(client string, email string) error
4056
}
4157

58+
func newResettableOAuthTokenSource(newSource func(*oauth2.Token) oauth2.TokenSource, initial *oauth2.Token) *resettableOAuthTokenSource {
59+
return &resettableOAuthTokenSource{
60+
source: newSource(initial),
61+
newSource: newSource,
62+
refreshToken: strings.TrimSpace(initial.RefreshToken),
63+
}
64+
}
65+
66+
func (r *resettableOAuthTokenSource) Token() (*oauth2.Token, error) {
67+
r.mu.Lock()
68+
source := r.source
69+
r.mu.Unlock()
70+
71+
t, err := source.Token()
72+
if err != nil {
73+
return nil, fmt.Errorf("resettable oauth token source: %w", err)
74+
}
75+
76+
r.rememberRefreshToken(t)
77+
78+
return t, nil
79+
}
80+
81+
func (r *resettableOAuthTokenSource) ForceRefresh(context.Context) (*oauth2.Token, error) {
82+
r.mu.Lock()
83+
refreshToken := r.refreshToken
84+
r.source = r.newSource(&oauth2.Token{RefreshToken: refreshToken})
85+
source := r.source
86+
r.mu.Unlock()
87+
88+
t, err := source.Token()
89+
if err != nil {
90+
return nil, fmt.Errorf("resettable oauth token source refresh: %w", err)
91+
}
92+
93+
r.rememberRefreshToken(t)
94+
95+
return t, nil
96+
}
97+
98+
func (r *resettableOAuthTokenSource) rememberRefreshToken(t *oauth2.Token) {
99+
if t == nil {
100+
return
101+
}
102+
103+
if refreshToken := strings.TrimSpace(t.RefreshToken); refreshToken != "" {
104+
r.mu.Lock()
105+
r.refreshToken = refreshToken
106+
r.mu.Unlock()
107+
}
108+
}
109+
42110
func newPersistingTokenSource(base oauth2.TokenSource, store secrets.Store, client string, email string, tok secrets.Token, serviceLabel string, updateEmailReferences googleauth.EmailReferenceUpdater) oauth2.TokenSource {
43111
return &persistingTokenSource{
44112
base: base,
@@ -57,6 +125,30 @@ func (p *persistingTokenSource) Token() (*oauth2.Token, error) {
57125
return nil, fmt.Errorf("base token source: %w", err)
58126
}
59127

128+
return p.persistToken(t)
129+
}
130+
131+
func (p *persistingTokenSource) ForceRefresh(ctx context.Context) error {
132+
refresher, ok := p.base.(forceRefreshTokenSource)
133+
if !ok {
134+
return errBaseTokenSourceCannotForceRefresh
135+
}
136+
137+
t, err := refresher.ForceRefresh(ctx)
138+
if err != nil {
139+
return fmt.Errorf("force token refresh: %w", err)
140+
}
141+
142+
_, err = p.persistToken(t)
143+
144+
return err
145+
}
146+
147+
func (p *persistingTokenSource) persistToken(t *oauth2.Token) (*oauth2.Token, error) {
148+
if t == nil {
149+
return nil, errBaseTokenSourceReturnedNilToken
150+
}
151+
60152
refreshToken := strings.TrimSpace(t.RefreshToken)
61153

62154
p.mu.Lock()
@@ -126,26 +218,26 @@ func (p *persistingTokenSource) Token() (*oauth2.Token, error) {
126218
}
127219

128220
if err := p.store.SetToken(p.client, persistEmail, updated); err != nil {
129-
slog.Warn("persist refreshed token metadata failed", "email", persistEmail, "client", p.client, "err", err)
221+
slog.Warn("persist refreshed token metadata failed", "email", persistEmail, "client", p.client, "err", err) //nolint:gosec // logged values are token metadata identifiers for auth diagnostics
130222
return t, nil
131223
}
132224

133225
if !strings.EqualFold(p.email, persistEmail) {
134226
if err := googleauth.MigrateStoredEmailReferences(p.store, p.updateEmailReferences, p.client, p.email, persistEmail); err != nil {
135-
slog.Warn("migrate renamed token email references failed", "old_email", p.email, "new_email", persistEmail, "client", p.client, "err", err)
227+
slog.Warn("migrate renamed token email references failed", "old_email", p.email, "new_email", persistEmail, "client", p.client, "err", err) //nolint:gosec // logged values are token metadata identifiers for auth diagnostics
136228
}
137229

138230
aliasDeleter, ok := p.store.(tokenAliasDeleter)
139231
if !ok {
140-
slog.Debug("token store cannot delete renamed email alias", "old_email", p.email, "new_email", persistEmail, "client", p.client)
232+
slog.Debug("token store cannot delete renamed email alias", "old_email", p.email, "new_email", persistEmail, "client", p.client) //nolint:gosec // logged values are token metadata identifiers for auth diagnostics
141233
} else if err := aliasDeleter.DeleteTokenAlias(p.client, p.email); err != nil {
142-
slog.Warn("delete renamed token alias failed", "old_email", p.email, "new_email", persistEmail, "client", p.client, "err", err)
234+
slog.Warn("delete renamed token alias failed", "old_email", p.email, "new_email", persistEmail, "client", p.client, "err", err) //nolint:gosec // logged values are token metadata identifiers for auth diagnostics
143235
}
144236
}
145237

146238
p.tok = updated
147239
p.email = persistEmail
148-
slog.Debug("persisted refreshed token metadata", "email", persistEmail, "client", p.client)
240+
slog.Debug("persisted refreshed token metadata", "email", persistEmail, "client", p.client) //nolint:gosec // logged values are token metadata identifiers for auth diagnostics
149241

150242
return t, nil
151243
}
@@ -276,7 +368,9 @@ func tokenSourceForAccountScopesWithStoredScopeCheck(
276368
// Ensure refresh-token exchanges don't hang forever.
277369
ctx = context.WithValue(ctx, oauth2.HTTPClient, &http.Client{Timeout: tokenExchangeTimeout})
278370

279-
baseSource := cfg.TokenSource(ctx, &oauth2.Token{
371+
baseSource := newResettableOAuthTokenSource(func(t *oauth2.Token) oauth2.TokenSource {
372+
return cfg.TokenSource(ctx, t)
373+
}, &oauth2.Token{
280374
RefreshToken: tok.RefreshToken,
281375
AccessToken: strings.TrimSpace(tok.AccessToken),
282376
Expiry: tok.AccessTokenExpiresAt,

internal/googleapi/transport.go

Lines changed: 58 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -11,10 +11,14 @@ import (
1111
"math/big"
1212
"net/http"
1313
"strconv"
14+
"strings"
1415
"time"
1516
)
1617

17-
const maxBufferedReplayBodyBytes = int64(16 << 20)
18+
const (
19+
maxBufferedReplayBodyBytes = int64(16 << 20)
20+
maxAuthRetryResponseBodyBytes = int64(1 << 20)
21+
)
1822

1923
var errRequestBodyTooLarge = errors.New("request body too large to buffer for retry")
2024

@@ -26,6 +30,7 @@ type RetryTransport struct {
2630
MaxRetries5xx int
2731
BaseDelay time.Duration
2832
CircuitBreaker *CircuitBreaker
33+
RefreshAuth func(context.Context) error
2934
}
3035

3136
// NewRetryTransport creates a RetryTransport with sensible defaults.
@@ -57,6 +62,7 @@ func (t *RetryTransport) RoundTrip(req *http.Request) (*http.Response, error) {
5762
var resp *http.Response
5863
retries429 := 0
5964
retries5xx := 0
65+
retriedAuth := false
6066

6167
for {
6268
// Reset body for retry
@@ -134,6 +140,28 @@ func (t *RetryTransport) RoundTrip(req *http.Request) (*http.Response, error) {
134140
continue
135141
}
136142

143+
if resp.StatusCode == http.StatusForbidden && t.RefreshAuth != nil && !retriedAuth && replayable {
144+
insufficientScopes, detectErr := responseIndicatesInsufficientScopes(resp)
145+
if detectErr != nil {
146+
slog.Debug("could not inspect auth failure response for retry", "err", detectErr)
147+
return resp, nil
148+
}
149+
150+
if insufficientScopes {
151+
slog.Debug("insufficient scopes response, refreshing auth token and retrying")
152+
153+
drainAndClose(resp.Body)
154+
155+
if err := t.RefreshAuth(req.Context()); err != nil {
156+
return nil, fmt.Errorf("refresh auth after insufficient scopes response: %w", err)
157+
}
158+
159+
retriedAuth = true
160+
161+
continue
162+
}
163+
}
164+
137165
// Other errors (4xx except 429): don't retry
138166
return resp, nil
139167
}
@@ -249,3 +277,32 @@ func drainAndClose(body io.ReadCloser) {
249277
_, _ = io.Copy(io.Discard, io.LimitReader(body, 1<<20))
250278
_ = body.Close()
251279
}
280+
281+
func responseIndicatesInsufficientScopes(resp *http.Response) (bool, error) {
282+
if resp == nil || resp.Body == nil {
283+
return false, nil
284+
}
285+
286+
if resp.ContentLength > maxAuthRetryResponseBodyBytes {
287+
return false, nil
288+
}
289+
290+
bodyBytes, err := io.ReadAll(io.LimitReader(resp.Body, maxAuthRetryResponseBodyBytes+1))
291+
_ = resp.Body.Close()
292+
resp.Body = io.NopCloser(bytes.NewReader(bodyBytes))
293+
294+
if err != nil {
295+
return false, fmt.Errorf("read auth failure response: %w", err)
296+
}
297+
298+
if int64(len(bodyBytes)) > maxAuthRetryResponseBodyBytes {
299+
return false, nil
300+
}
301+
302+
text := strings.ToLower(string(bodyBytes))
303+
normalized := strings.NewReplacer("_", " ", "-", " ").Replace(text)
304+
305+
return strings.Contains(normalized, "insufficient authentication scopes") ||
306+
strings.Contains(normalized, "insufficient scopes") ||
307+
strings.Contains(normalized, "access token scope insufficient"), nil
308+
}

0 commit comments

Comments
 (0)