@@ -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+
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+ 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+
42110func 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 ,
0 commit comments