@@ -35,10 +35,76 @@ 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+
3854type 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+ defer r .mu .Unlock ()
69+
70+ t , err := r .source .Token ()
71+ if err != nil {
72+ return nil , fmt .Errorf ("resettable oauth token source: %w" , err )
73+ }
74+
75+ r .rememberRefreshTokenLocked (t )
76+
77+ return t , nil
78+ }
79+
80+ func (r * resettableOAuthTokenSource ) ForceRefresh (context.Context ) (* oauth2.Token , error ) {
81+ r .mu .Lock ()
82+ defer r .mu .Unlock ()
83+
84+ refreshToken := r .refreshToken
85+ candidate := r .newSource (& oauth2.Token {RefreshToken : refreshToken })
86+
87+ t , err := candidate .Token ()
88+ if err != nil {
89+ return nil , fmt .Errorf ("resettable oauth token source refresh: %w" , err )
90+ }
91+
92+ r .source = candidate
93+ r .rememberRefreshTokenLocked (t )
94+
95+ return t , nil
96+ }
97+
98+ func (r * resettableOAuthTokenSource ) rememberRefreshTokenLocked (t * oauth2.Token ) {
99+ if t == nil {
100+ return
101+ }
102+
103+ if refreshToken := strings .TrimSpace (t .RefreshToken ); refreshToken != "" {
104+ r .refreshToken = refreshToken
105+ }
106+ }
107+
42108func newPersistingTokenSource (base oauth2.TokenSource , store secrets.Store , client string , email string , tok secrets.Token , serviceLabel string , updateEmailReferences googleauth.EmailReferenceUpdater ) oauth2.TokenSource {
43109 return & persistingTokenSource {
44110 base : base ,
@@ -52,16 +118,43 @@ func newPersistingTokenSource(base oauth2.TokenSource, store secrets.Store, clie
52118}
53119
54120func (p * persistingTokenSource ) Token () (* oauth2.Token , error ) {
121+ p .mu .Lock ()
122+ defer p .mu .Unlock ()
123+
55124 t , err := p .base .Token ()
56125 if err != nil {
57126 return nil , fmt .Errorf ("base token source: %w" , err )
58127 }
59128
60- refreshToken := strings .TrimSpace (t .RefreshToken )
129+ return p .persistTokenLocked (t )
130+ }
131+
132+ func (p * persistingTokenSource ) ForceRefresh (ctx context.Context ) error {
133+ refresher , ok := p .base .(forceRefreshTokenSource )
134+ if ! ok {
135+ return errBaseTokenSourceCannotForceRefresh
136+ }
61137
62138 p .mu .Lock ()
63139 defer p .mu .Unlock ()
64140
141+ t , err := refresher .ForceRefresh (ctx )
142+ if err != nil {
143+ return fmt .Errorf ("force token refresh: %w" , err )
144+ }
145+
146+ _ , err = p .persistTokenLocked (t )
147+
148+ return err
149+ }
150+
151+ func (p * persistingTokenSource ) persistTokenLocked (t * oauth2.Token ) (* oauth2.Token , error ) {
152+ if t == nil {
153+ return nil , errBaseTokenSourceReturnedNilToken
154+ }
155+
156+ refreshToken := strings .TrimSpace (t .RefreshToken )
157+
65158 updated := p .tok
66159 changed := false
67160 emailChangedFromIdentity := false
@@ -126,26 +219,26 @@ func (p *persistingTokenSource) Token() (*oauth2.Token, error) {
126219 }
127220
128221 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 )
222+ 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
130223 return t , nil
131224 }
132225
133226 if ! strings .EqualFold (p .email , persistEmail ) {
134227 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 )
228+ 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
136229 }
137230
138231 aliasDeleter , ok := p .store .(tokenAliasDeleter )
139232 if ! ok {
140- slog .Debug ("token store cannot delete renamed email alias" , "old_email" , p .email , "new_email" , persistEmail , "client" , p .client )
233+ 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
141234 } 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 )
235+ 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
143236 }
144237 }
145238
146239 p .tok = updated
147240 p .email = persistEmail
148- slog .Debug ("persisted refreshed token metadata" , "email" , persistEmail , "client" , p .client )
241+ slog .Debug ("persisted refreshed token metadata" , "email" , persistEmail , "client" , p .client ) //nolint:gosec // logged values are token metadata identifiers for auth diagnostics
149242
150243 return t , nil
151244}
@@ -276,7 +369,9 @@ func tokenSourceForAccountScopesWithStoredScopeCheck(
276369 // Ensure refresh-token exchanges don't hang forever.
277370 ctx = context .WithValue (ctx , oauth2 .HTTPClient , & http.Client {Timeout : tokenExchangeTimeout })
278371
279- baseSource := cfg .TokenSource (ctx , & oauth2.Token {
372+ baseSource := newResettableOAuthTokenSource (func (t * oauth2.Token ) oauth2.TokenSource {
373+ return cfg .TokenSource (ctx , t )
374+ }, & oauth2.Token {
280375 RefreshToken : tok .RefreshToken ,
281376 AccessToken : strings .TrimSpace (tok .AccessToken ),
282377 Expiry : tok .AccessTokenExpiresAt ,
0 commit comments