diff --git a/util/oauth/tokensource.go b/util/oauth/tokensource.go index f5d91fbf2..0cdc2aaf7 100644 --- a/util/oauth/tokensource.go +++ b/util/oauth/tokensource.go @@ -46,12 +46,16 @@ func (ts *TokenSource) Token() (*oauth2.Token, error) { } if token.AccessToken == "" { - err = errors.New("token refresh failed to obtain access token") - } else { - err = ts.mergeToken(token) + return nil, errors.New("token refresh failed to obtain access token") } - return ts.token, err + if token.RefreshToken == "" { + token.RefreshToken = ts.token.RefreshToken + } + + ts.token = token + + return ts.token, nil } // mergeToken updates a token while preventing wiping the refresh token diff --git a/util/oauth/tokensource_test.go b/util/oauth/tokensource_test.go index 68d54fe09..af8d24454 100644 --- a/util/oauth/tokensource_test.go +++ b/util/oauth/tokensource_test.go @@ -2,30 +2,36 @@ package oauth import ( "testing" + "time" + "github.com/stretchr/testify/require" "golang.org/x/oauth2" ) +type tkr struct{} + +func (tkr *tkr) RefreshToken(_ *oauth2.Token) (*oauth2.Token, error) { + return (&oauth2.Token{ + AccessToken: "new", + }).WithExtra(map[string]any{ + "foo": "bar", + }), nil +} + func TestMerge(t *testing.T) { ts := &TokenSource{ token: &oauth2.Token{ AccessToken: "access", RefreshToken: "refresh", + Expiry: time.Now(), }, + refresher: new(tkr), } - r := &oauth2.Token{ - AccessToken: "new", - } + r, err := ts.Token() + require.NoError(t, err) - if err := ts.mergeToken(r); err != nil { - t.Error(err) - } - - if ts.token.AccessToken != "new" { - t.Error("unexpected access token", ts.token) - } - if ts.token.RefreshToken != "refresh" { - t.Error("unexpected refresh token", ts.token) - } + require.Equal(t, "new", r.AccessToken, "unexpected access token") + require.Equal(t, "refresh", r.RefreshToken, "unexpected refresh token") + require.Equal(t, "bar", r.Extra("foo")) }