Refactor token refresh handling into TokenSource and TokenRefresher (#837)
This commit is contained in:
parent
33dc71938e
commit
8befdcf41e
13 changed files with 233 additions and 261 deletions
|
|
@ -8,9 +8,9 @@ import (
|
|||
"time"
|
||||
|
||||
"github.com/andig/evcc/api"
|
||||
"github.com/andig/evcc/internal/vehicle/oidc"
|
||||
"github.com/andig/evcc/provider"
|
||||
"github.com/andig/evcc/util"
|
||||
"github.com/andig/evcc/util/oauth"
|
||||
"github.com/andig/evcc/util/request"
|
||||
)
|
||||
|
||||
|
|
@ -24,7 +24,7 @@ type Ford struct {
|
|||
*embed
|
||||
*request.Helper
|
||||
user, password, vin string
|
||||
tokens oidc.Token
|
||||
tokens oauth.Token
|
||||
chargeStateG func() (float64, error)
|
||||
}
|
||||
|
||||
|
|
@ -84,7 +84,7 @@ func (v *Ford) login(user, password string) error {
|
|||
return err
|
||||
}
|
||||
|
||||
var tokens oidc.Token
|
||||
var tokens oauth.Token
|
||||
if err = v.DoJSON(req, &tokens); err == nil {
|
||||
v.tokens = tokens
|
||||
}
|
||||
|
|
|
|||
37
internal/vehicle/id/refresh.go
Normal file
37
internal/vehicle/id/refresh.go
Normal file
|
|
@ -0,0 +1,37 @@
|
|||
package id
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/andig/evcc/util"
|
||||
"github.com/andig/evcc/util/oauth"
|
||||
"github.com/andig/evcc/util/request"
|
||||
"golang.org/x/oauth2"
|
||||
)
|
||||
|
||||
type tokenRefresher struct {
|
||||
*request.Helper
|
||||
}
|
||||
|
||||
func Refresher(log *util.Logger) oauth.TokenRefresher {
|
||||
return &tokenRefresher{
|
||||
Helper: request.NewHelper(log),
|
||||
}
|
||||
}
|
||||
|
||||
// Refresh is the oauth.TokenRefresher
|
||||
func (tr *tokenRefresher) Refresh(token *oauth2.Token) (*oauth2.Token, error) {
|
||||
uri := "https://login.apps.emea.vwapps.io/refresh/v1"
|
||||
|
||||
req, err := request.New(http.MethodGet, uri, nil, map[string]string{
|
||||
"Accept": "application/json",
|
||||
"Authorization": "Bearer " + token.RefreshToken,
|
||||
})
|
||||
|
||||
var res Token
|
||||
if err == nil {
|
||||
err = tr.DoJSON(req, &token)
|
||||
}
|
||||
|
||||
return (*oauth2.Token)(&res), err
|
||||
}
|
||||
|
|
@ -2,11 +2,8 @@ package id
|
|||
|
||||
import (
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
"github.com/andig/evcc/util"
|
||||
"github.com/andig/evcc/util/request"
|
||||
"golang.org/x/oauth2"
|
||||
)
|
||||
|
||||
|
|
@ -29,42 +26,3 @@ func (t *Token) UnmarshalJSON(data []byte) error {
|
|||
|
||||
return err
|
||||
}
|
||||
|
||||
func (t *Token) TokenSource(log *util.Logger) oauth2.TokenSource {
|
||||
return &TokenSource{
|
||||
Helper: request.NewHelper(log),
|
||||
token: t,
|
||||
}
|
||||
}
|
||||
|
||||
type TokenSource struct {
|
||||
*request.Helper
|
||||
token *Token
|
||||
}
|
||||
|
||||
func (ts *TokenSource) Token() (*oauth2.Token, error) {
|
||||
var err error
|
||||
if time.Until(ts.token.Expiry) < time.Minute {
|
||||
err = ts.refreshToken()
|
||||
}
|
||||
|
||||
return (*oauth2.Token)(ts.token), err
|
||||
}
|
||||
|
||||
func (ts *TokenSource) refreshToken() error {
|
||||
uri := "https://login.apps.emea.vwapps.io/refresh/v1"
|
||||
|
||||
req, err := request.New(http.MethodGet, uri, nil, map[string]string{
|
||||
"Accept": "application/json",
|
||||
"Authorization": "Bearer " + ts.token.RefreshToken,
|
||||
})
|
||||
|
||||
if err == nil {
|
||||
var token Token
|
||||
if err = ts.DoJSON(req, &token); err == nil {
|
||||
ts.token = &token
|
||||
}
|
||||
}
|
||||
|
||||
return err
|
||||
}
|
||||
|
|
|
|||
|
|
@ -3,7 +3,6 @@ package vehicle
|
|||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/url"
|
||||
|
|
@ -12,10 +11,11 @@ import (
|
|||
|
||||
"github.com/andig/evcc/api"
|
||||
"github.com/andig/evcc/internal/vehicle/kamereon"
|
||||
"github.com/andig/evcc/internal/vehicle/oidc"
|
||||
"github.com/andig/evcc/provider"
|
||||
"github.com/andig/evcc/util"
|
||||
"github.com/andig/evcc/util/oauth"
|
||||
"github.com/andig/evcc/util/request"
|
||||
"golang.org/x/oauth2"
|
||||
)
|
||||
|
||||
// Credits to
|
||||
|
|
@ -46,8 +46,6 @@ type Nissan struct {
|
|||
*request.Helper
|
||||
log *util.Logger
|
||||
user, password, vin string
|
||||
userID string
|
||||
tokens oidc.Token
|
||||
*kamereon.API
|
||||
}
|
||||
|
||||
|
|
@ -82,10 +80,17 @@ func NewNissanFromConfig(other map[string]interface{}) (api.Vehicle, error) {
|
|||
vin: strings.ToUpper(cc.VIN),
|
||||
}
|
||||
|
||||
err := v.authFlow()
|
||||
token, err := v.authFlow()
|
||||
if err == nil {
|
||||
// replace transport client with authenticated client
|
||||
v.Helper.Client.Transport = &oauth2.Transport{
|
||||
Source: oauth.RefreshTokenSource((*oauth2.Token)(&token), v),
|
||||
Base: v.Helper.Client.Transport,
|
||||
}
|
||||
}
|
||||
|
||||
if err == nil && cc.VIN == "" {
|
||||
v.vin, err = findVehicle(v.vehicles(v.userID))
|
||||
v.vin, err = findVehicle(v.vehicles())
|
||||
if err == nil {
|
||||
log.DEBUG.Printf("found vehicle: %v", v.vin)
|
||||
}
|
||||
|
|
@ -121,9 +126,10 @@ type nissanToken struct {
|
|||
Realm string `json:"realm"`
|
||||
}
|
||||
|
||||
func (v *Nissan) authFlow() error {
|
||||
uri := fmt.Sprintf("%s/json/realms/root/realms/%s/authenticate", nissanAuthBaseURL, nissanRealm)
|
||||
func (v *Nissan) authFlow() (oauth.Token, error) {
|
||||
client := request.NewHelper(v.log) // no underlying oauth transport
|
||||
|
||||
uri := fmt.Sprintf("%s/json/realms/root/realms/%s/authenticate", nissanAuthBaseURL, nissanRealm)
|
||||
req, err := request.New(http.MethodPost, uri, nil, map[string]string{
|
||||
"Accept-Api-Version": nissanAPIVersion,
|
||||
"X-Username": "anonymous",
|
||||
|
|
@ -131,14 +137,15 @@ func (v *Nissan) authFlow() error {
|
|||
"Accept": "application/json",
|
||||
})
|
||||
|
||||
var oauth nissanToken
|
||||
var nToken nissanToken
|
||||
var realm string
|
||||
var resp *http.Response
|
||||
var code string
|
||||
|
||||
if err == nil {
|
||||
var res nissanAuth
|
||||
if err = v.DoJSON(req, &res); err != nil {
|
||||
return err
|
||||
if err = client.DoJSON(req, &res); err != nil {
|
||||
return oauth.Token{}, err
|
||||
}
|
||||
|
||||
for id, cb := range res.Callbacks {
|
||||
|
|
@ -164,13 +171,12 @@ func (v *Nissan) authFlow() error {
|
|||
}
|
||||
|
||||
if err == nil {
|
||||
err = v.DoJSON(req, &oauth)
|
||||
err = client.DoJSON(req, &nToken)
|
||||
realm = strings.Trim(nToken.Realm, "/")
|
||||
}
|
||||
}
|
||||
|
||||
if err == nil {
|
||||
uri := fmt.Sprintf("%s/oauth2/%s/authorize", nissanAuthBaseURL, strings.Trim(oauth.Realm, "/"))
|
||||
|
||||
data := url.Values{
|
||||
"client_id": []string{nissanClientID},
|
||||
"redirect_uri": []string{nissanRedirectURI},
|
||||
|
|
@ -179,15 +185,15 @@ func (v *Nissan) authFlow() error {
|
|||
"nonce": []string{"sdfdsfez"},
|
||||
}
|
||||
|
||||
uri += "?" + data.Encode()
|
||||
uri := fmt.Sprintf("%s/oauth2/%s/authorize?%s", nissanAuthBaseURL, realm, data.Encode())
|
||||
req, err = request.New(http.MethodGet, uri, nil, map[string]string{
|
||||
"Cookie": "i18next=en-UK; amlbcookie=05; kauthSession=" + oauth.TokenID,
|
||||
"Cookie": "i18next=en-UK; amlbcookie=05; kauthSession=" + nToken.TokenID,
|
||||
})
|
||||
|
||||
if err == nil {
|
||||
v.Client.CheckRedirect = func(req *http.Request, via []*http.Request) error { return http.ErrUseLastResponse }
|
||||
resp, err = v.Do(req)
|
||||
v.Client.CheckRedirect = nil
|
||||
client.CheckRedirect = func(req *http.Request, via []*http.Request) error { return http.ErrUseLastResponse }
|
||||
resp, err = client.Do(req)
|
||||
client.CheckRedirect = nil
|
||||
|
||||
if err == nil {
|
||||
resp.Body.Close()
|
||||
|
|
@ -202,9 +208,8 @@ func (v *Nissan) authFlow() error {
|
|||
}
|
||||
}
|
||||
|
||||
var res oauth.Token
|
||||
if err == nil {
|
||||
uri = fmt.Sprintf("%s/oauth2/%s/access_token", nissanAuthBaseURL, strings.Trim(oauth.Realm, "/"))
|
||||
|
||||
data := url.Values{
|
||||
"code": []string{code},
|
||||
"client_id": []string{nissanClientID},
|
||||
|
|
@ -213,71 +218,38 @@ func (v *Nissan) authFlow() error {
|
|||
"grant_type": []string{"authorization_code"},
|
||||
}
|
||||
|
||||
uri += "?" + data.Encode()
|
||||
uri = fmt.Sprintf("%s/oauth2/%s/access_token?%s", nissanAuthBaseURL, realm, data.Encode())
|
||||
req, err = request.New(http.MethodPost, uri, nil, request.URLEncoding)
|
||||
if err == nil {
|
||||
if err = v.DoJSON(req, &v.tokens); err == nil && v.tokens.AccessToken == "" {
|
||||
err = errors.New("missing access token")
|
||||
}
|
||||
err = client.DoJSON(req, &res)
|
||||
}
|
||||
}
|
||||
|
||||
if err == nil {
|
||||
uri = fmt.Sprintf("%s/v1/users/current", nissanUserAdapterBaseURL)
|
||||
|
||||
var user struct{ UserID string }
|
||||
if req, err = request.New(http.MethodGet, uri, nil, nil); err == nil {
|
||||
if err = v.request(req, &user); err == nil {
|
||||
v.userID = user.UserID
|
||||
}
|
||||
}
|
||||
|
||||
if v.userID == "" {
|
||||
err = errors.New("missing user id")
|
||||
}
|
||||
}
|
||||
|
||||
return err
|
||||
return res, err
|
||||
}
|
||||
|
||||
func (v *Nissan) refreshToken() error {
|
||||
uri := fmt.Sprintf("%s/oauth2/%s/access_token", nissanAuthBaseURL, nissanRealm)
|
||||
|
||||
func (v *Nissan) Refresh(token *oauth2.Token) (*oauth2.Token, error) {
|
||||
data := url.Values{
|
||||
"client_id": []string{nissanClientID},
|
||||
"client_secret": []string{nissanClientSecret},
|
||||
"grant_type": {"refresh_token"},
|
||||
"refresh_token": {v.tokens.RefreshToken},
|
||||
"refresh_token": {token.RefreshToken},
|
||||
}
|
||||
|
||||
uri += "?" + data.Encode()
|
||||
uri := fmt.Sprintf("%s/oauth2/%s/access_token?%s", nissanAuthBaseURL, nissanRealm, data.Encode())
|
||||
req, err := request.New(http.MethodPost, uri, nil, request.URLEncoding)
|
||||
|
||||
var res oauth.Token
|
||||
if err == nil {
|
||||
if err = v.DoJSON(req, &v.tokens); err == nil && v.tokens.AccessToken == "" {
|
||||
err = errors.New("missing access token")
|
||||
}
|
||||
client := request.NewHelper(v.log)
|
||||
err = client.DoJSON(req, &res)
|
||||
}
|
||||
|
||||
return err
|
||||
}
|
||||
|
||||
// request executes given request and handles token refresh
|
||||
func (v *Nissan) request(req *http.Request, res interface{}) error {
|
||||
req.Header.Set("Authorization", "Bearer "+v.tokens.AccessToken)
|
||||
err := v.DoJSON(req, &res)
|
||||
|
||||
// repeat auth if error
|
||||
if err != nil {
|
||||
if err = v.refreshToken(); err != nil {
|
||||
err = v.authFlow()
|
||||
}
|
||||
if err == nil {
|
||||
req.Header.Set("Authorization", "Bearer "+v.tokens.AccessToken)
|
||||
err = v.DoJSON(req, &res)
|
||||
}
|
||||
res, err = v.authFlow()
|
||||
}
|
||||
|
||||
return err
|
||||
return (*oauth2.Token)(&res), err
|
||||
}
|
||||
|
||||
type nissanVehicles struct {
|
||||
|
|
@ -290,13 +262,15 @@ type nissanVehicle struct {
|
|||
PictureURL string
|
||||
}
|
||||
|
||||
func (v *Nissan) vehicles(userID string) ([]string, error) {
|
||||
uri := fmt.Sprintf("%s/v2/users/%s/cars", nissanUserBaseURL, userID)
|
||||
func (v *Nissan) vehicles() ([]string, error) {
|
||||
var user struct{ UserID string }
|
||||
uri := fmt.Sprintf("%s/v1/users/current", nissanUserAdapterBaseURL)
|
||||
err := v.GetJSON(uri, &user)
|
||||
|
||||
var res nissanVehicles
|
||||
req, err := request.New(http.MethodGet, uri, nil, nil)
|
||||
if err == nil {
|
||||
err = v.request(req, &res)
|
||||
uri := fmt.Sprintf("%s/v2/users/%s/cars", nissanUserBaseURL, user.UserID)
|
||||
err = v.GetJSON(uri, &res)
|
||||
}
|
||||
|
||||
var vehicles []string
|
||||
|
|
@ -321,16 +295,13 @@ func (v *Nissan) batteryAPI() (interface{}, error) {
|
|||
|
||||
var res kamereon.Response
|
||||
if err == nil {
|
||||
err = v.request(req, &res)
|
||||
err = v.DoJSON(req, &res)
|
||||
}
|
||||
|
||||
// request battery status
|
||||
if err == nil {
|
||||
uri = fmt.Sprintf("%s/v1/cars/%s/battery-status", nissanCarAdapterBaseURL, v.vin)
|
||||
|
||||
if req, err = request.New(http.MethodGet, uri, nil, nil); err == nil {
|
||||
err = v.request(req, &res)
|
||||
}
|
||||
err = v.GetJSON(uri, &res)
|
||||
}
|
||||
|
||||
return res, err
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ import (
|
|||
|
||||
"github.com/andig/evcc/internal/vehicle/id"
|
||||
"github.com/andig/evcc/util"
|
||||
"github.com/andig/evcc/util/oauth"
|
||||
"github.com/andig/evcc/util/request"
|
||||
"golang.org/x/net/publicsuffix"
|
||||
"golang.org/x/oauth2"
|
||||
|
|
@ -148,9 +149,9 @@ func (v *Identity) Login(query url.Values, user, password string) error {
|
|||
})
|
||||
|
||||
if err == nil {
|
||||
var token Token
|
||||
var token oauth.Token
|
||||
if err = v.DoJSON(req, &token); err == nil {
|
||||
v.TokenSource = token.TokenSource(v.log, v.clientID)
|
||||
v.TokenSource = oauth.RefreshTokenSource((*oauth2.Token)(&token), refresher(v.log, v.clientID))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -174,7 +175,7 @@ func (v *Identity) Login(query url.Values, user, password string) error {
|
|||
if err == nil {
|
||||
var token id.Token
|
||||
if err = v.DoJSON(req, &token); err == nil {
|
||||
v.TokenSource = token.TokenSource(v.log)
|
||||
v.TokenSource = oauth.RefreshTokenSource((*oauth2.Token)(&token), id.Refresher(v.log))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
45
internal/vehicle/vw/refresh.go
Normal file
45
internal/vehicle/vw/refresh.go
Normal file
|
|
@ -0,0 +1,45 @@
|
|||
package vw
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
|
||||
"github.com/andig/evcc/util"
|
||||
"github.com/andig/evcc/util/oauth"
|
||||
"github.com/andig/evcc/util/request"
|
||||
"golang.org/x/oauth2"
|
||||
)
|
||||
|
||||
type tokenRefresher struct {
|
||||
*request.Helper
|
||||
clientID string
|
||||
}
|
||||
|
||||
func refresher(log *util.Logger, clientID string) oauth.TokenRefresher {
|
||||
return &tokenRefresher{
|
||||
Helper: request.NewHelper(log),
|
||||
clientID: clientID,
|
||||
}
|
||||
}
|
||||
|
||||
// Refresh is the oauth.TokenRefresher
|
||||
func (tr *tokenRefresher) Refresh(token *oauth2.Token) (*oauth2.Token, error) {
|
||||
data := url.Values(map[string][]string{
|
||||
"grant_type": {"refresh_token"},
|
||||
"refresh_token": {token.RefreshToken},
|
||||
"scope": {"sc2:fal"},
|
||||
})
|
||||
|
||||
req, err := request.New(http.MethodPost, OauthTokenURI, strings.NewReader(data.Encode()), map[string]string{
|
||||
"Content-Type": "application/x-www-form-urlencoded",
|
||||
"X-Client-Id": tr.clientID,
|
||||
})
|
||||
|
||||
var res oauth.Token
|
||||
if err == nil {
|
||||
err = tr.DoJSON(req, &res)
|
||||
}
|
||||
|
||||
return (*oauth2.Token)(&res), err
|
||||
}
|
||||
|
|
@ -1,80 +0,0 @@
|
|||
package vw
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/andig/evcc/internal/vehicle/oidc"
|
||||
"github.com/andig/evcc/util"
|
||||
"github.com/andig/evcc/util/request"
|
||||
"github.com/imdario/mergo"
|
||||
"golang.org/x/oauth2"
|
||||
)
|
||||
|
||||
// Token is the VW token
|
||||
type Token oauth2.Token
|
||||
|
||||
func (t *Token) UnmarshalJSON(data []byte) error {
|
||||
var o oidc.Token
|
||||
|
||||
err := json.Unmarshal(data, &o)
|
||||
if err == nil {
|
||||
*t = (Token)(o.Token)
|
||||
}
|
||||
|
||||
return err
|
||||
}
|
||||
|
||||
func (t *Token) TokenSource(log *util.Logger, clientID string) oauth2.TokenSource {
|
||||
return &TokenSource{
|
||||
Helper: request.NewHelper(log),
|
||||
clientID: clientID,
|
||||
token: t,
|
||||
}
|
||||
}
|
||||
|
||||
type TokenSource struct {
|
||||
*request.Helper
|
||||
clientID string
|
||||
token *Token
|
||||
}
|
||||
|
||||
func (ts *TokenSource) Token() (*oauth2.Token, error) {
|
||||
var err error
|
||||
if time.Until(ts.token.Expiry) < time.Minute {
|
||||
err = ts.refreshToken()
|
||||
}
|
||||
|
||||
return (*oauth2.Token)(ts.token), err
|
||||
}
|
||||
|
||||
func (ts *TokenSource) refreshToken() error {
|
||||
data := url.Values(map[string][]string{
|
||||
"grant_type": {"refresh_token"},
|
||||
"refresh_token": {ts.token.RefreshToken},
|
||||
"scope": {"sc2:fal"},
|
||||
})
|
||||
|
||||
req, err := request.New(http.MethodPost, OauthTokenURI, strings.NewReader(data.Encode()), map[string]string{
|
||||
"Content-Type": "application/x-www-form-urlencoded",
|
||||
"X-Client-Id": ts.clientID,
|
||||
})
|
||||
|
||||
if err == nil {
|
||||
var token Token
|
||||
if err = ts.DoJSON(req, &token); err == nil {
|
||||
ts.mergeToken(token)
|
||||
}
|
||||
}
|
||||
|
||||
return err
|
||||
}
|
||||
|
||||
func (ts *TokenSource) mergeToken(t Token) {
|
||||
if err := mergo.Merge(ts.token, &t, mergo.WithOverride); err != nil {
|
||||
panic(err)
|
||||
}
|
||||
}
|
||||
|
|
@ -1,56 +0,0 @@
|
|||
package vw
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestUnmarshalJSON(t *testing.T) {
|
||||
var tok Token
|
||||
str := `{"access_token":"access","refresh_token":"refresh","token_type":"bearer","expires_in":3600}`
|
||||
|
||||
if err := json.Unmarshal([]byte(str), &tok); err != nil {
|
||||
t.Error(err)
|
||||
}
|
||||
|
||||
if tok.AccessToken != "access" {
|
||||
t.Error("AccessToken")
|
||||
}
|
||||
|
||||
if tok.RefreshToken != "refresh" {
|
||||
t.Error("RefreshToken")
|
||||
}
|
||||
|
||||
if tok.TokenType != "bearer" {
|
||||
t.Error("TokenType")
|
||||
}
|
||||
|
||||
if tok.Expiry.IsZero() {
|
||||
t.Error("Expiry")
|
||||
}
|
||||
}
|
||||
|
||||
func TestMerge(t *testing.T) {
|
||||
ts := &TokenSource{
|
||||
token: &Token{
|
||||
AccessToken: "access1",
|
||||
RefreshToken: "refresh1",
|
||||
},
|
||||
}
|
||||
|
||||
new := Token{
|
||||
AccessToken: "access2",
|
||||
}
|
||||
|
||||
ts.mergeToken(new)
|
||||
|
||||
tok := ts.token
|
||||
|
||||
if tok.AccessToken != "access2" {
|
||||
t.Error("AccessToken")
|
||||
}
|
||||
|
||||
if tok.RefreshToken != "refresh1" {
|
||||
t.Error("RefreshToken")
|
||||
}
|
||||
}
|
||||
|
|
@ -1,4 +1,4 @@
|
|||
package oidc
|
||||
package internal
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
|
|
@ -1,4 +1,4 @@
|
|||
package oidc
|
||||
package internal
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
22
util/oauth/token.go
Normal file
22
util/oauth/token.go
Normal file
|
|
@ -0,0 +1,22 @@
|
|||
package oauth
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
|
||||
"github.com/andig/evcc/util/oauth/internal"
|
||||
"golang.org/x/oauth2"
|
||||
)
|
||||
|
||||
// Token is an OAuth token that supports the expires_in attribute
|
||||
type Token oauth2.Token
|
||||
|
||||
func (t *Token) UnmarshalJSON(data []byte) error {
|
||||
var o internal.Token
|
||||
|
||||
err := json.Unmarshal(data, &o)
|
||||
if err == nil {
|
||||
*t = (Token)(o.Token)
|
||||
}
|
||||
|
||||
return err
|
||||
}
|
||||
43
util/oauth/tokensource.go
Normal file
43
util/oauth/tokensource.go
Normal file
|
|
@ -0,0 +1,43 @@
|
|||
package oauth
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"time"
|
||||
|
||||
"github.com/imdario/mergo"
|
||||
"golang.org/x/oauth2"
|
||||
)
|
||||
|
||||
type TokenRefresher interface {
|
||||
Refresh(token *oauth2.Token) (*oauth2.Token, error)
|
||||
}
|
||||
|
||||
type TokenSource struct {
|
||||
token *oauth2.Token
|
||||
refresher TokenRefresher
|
||||
}
|
||||
|
||||
func RefreshTokenSource(token *oauth2.Token, refresher TokenRefresher) oauth2.TokenSource {
|
||||
return &TokenSource{token, refresher}
|
||||
}
|
||||
|
||||
func (ts *TokenSource) Token() (*oauth2.Token, error) {
|
||||
var err error
|
||||
if time.Until(ts.token.Expiry) < time.Minute {
|
||||
var token *oauth2.Token
|
||||
if token, err = ts.refresher.Refresh(ts.token); err == nil {
|
||||
if token.AccessToken == "" {
|
||||
err = errors.New("token refresh failed to obtain access token")
|
||||
} else {
|
||||
err = ts.mergeToken(token)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return ts.token, err
|
||||
}
|
||||
|
||||
// mergeToken updates a token while preventing wiping the refresh token
|
||||
func (ts *TokenSource) mergeToken(t *oauth2.Token) error {
|
||||
return mergo.Merge(ts.token, t, mergo.WithOverride)
|
||||
}
|
||||
31
util/oauth/tokensource_test.go
Normal file
31
util/oauth/tokensource_test.go
Normal file
|
|
@ -0,0 +1,31 @@
|
|||
package oauth
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"golang.org/x/oauth2"
|
||||
)
|
||||
|
||||
func TestMerge(t *testing.T) {
|
||||
ts := &TokenSource{
|
||||
token: &oauth2.Token{
|
||||
AccessToken: "access",
|
||||
RefreshToken: "refresh",
|
||||
},
|
||||
}
|
||||
|
||||
r := &oauth2.Token{
|
||||
AccessToken: "new",
|
||||
}
|
||||
|
||||
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)
|
||||
}
|
||||
}
|
||||
Loading…
Add table
Add a link
Reference in a new issue