chore: preserve extra token properties during refresh
This commit is contained in:
parent
3206fd6a20
commit
b9a324ad9b
2 changed files with 27 additions and 17 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"))
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue