Refactor token refresh handling into TokenSource and TokenRefresher (#837)

This commit is contained in:
andig 2021-04-05 17:13:02 +02:00 • committed by GitHub
parent 33dc71938e
commit 8befdcf41e
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
13 changed files with 233 additions and 261 deletions

View file

@ -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
}

View 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
}

View file

@ -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
}

View file

@ -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

View file

@ -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))
}
}
}

View 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
}

View file

@ -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)
}
}

View file

@ -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")
}
}

View file

@ -1,4 +1,4 @@
package oidc
package internal
import (
"encoding/json"

View file

@ -1,4 +1,4 @@
package oidc
package internal
import (
"encoding/json"

22
util/oauth/token.go Normal file
View 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
View 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)
}

View 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)
}
}