OAuth2: refactor authorization framework (BC) (#23978)
This commit is contained in:
parent
c642732588
commit
2780278752
11 changed files with 241 additions and 311 deletions
|
|
@ -1,9 +0,0 @@
|
|||
package auth
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
)
|
||||
|
||||
type Authorizer interface {
|
||||
Transport(base http.RoundTripper) http.RoundTripper
|
||||
}
|
||||
|
|
@ -6,12 +6,13 @@ import (
|
|||
"strings"
|
||||
|
||||
reg "github.com/evcc-io/evcc/util/registry"
|
||||
"golang.org/x/oauth2"
|
||||
)
|
||||
|
||||
var registry = reg.New[Authorizer]("auth")
|
||||
var registry = reg.New[oauth2.TokenSource]("auth")
|
||||
|
||||
// NewFromConfig creates auth from configuration
|
||||
func NewFromConfig(ctx context.Context, typ string, other map[string]any) (Authorizer, error) {
|
||||
func NewFromConfig(ctx context.Context, typ string, other map[string]any) (oauth2.TokenSource, error) {
|
||||
factory, err := registry.Get(strings.ToLower(typ))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
|
|
|
|||
|
|
@ -1,27 +0,0 @@
|
|||
package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
|
||||
"github.com/evcc-io/evcc/util"
|
||||
)
|
||||
|
||||
type nop struct{}
|
||||
|
||||
func init() {
|
||||
registry.AddCtx("nop", NewNopFromConfig)
|
||||
}
|
||||
|
||||
func NewNopFromConfig(ctx context.Context, other map[string]any) (Authorizer, error) {
|
||||
var cc struct{}
|
||||
if err := util.DecodeOther(other, &cc); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return new(nop), nil
|
||||
}
|
||||
|
||||
func (p *nop) Transport(base http.RoundTripper) http.RoundTripper {
|
||||
return base
|
||||
}
|
||||
|
|
@ -4,9 +4,8 @@ import (
|
|||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"net/http"
|
||||
"fmt"
|
||||
"net/url"
|
||||
"strings"
|
||||
"sync"
|
||||
|
|
@ -22,58 +21,58 @@ import (
|
|||
type OAuth struct {
|
||||
oauth2.TokenSource
|
||||
mu sync.Mutex
|
||||
cc oauth2.Config
|
||||
log *util.Logger
|
||||
oc *oauth2.Config
|
||||
subject string
|
||||
cv string
|
||||
log *util.Logger
|
||||
ctx context.Context
|
||||
}
|
||||
|
||||
var (
|
||||
// oauthMu sync.Mutex
|
||||
oauthMu sync.Mutex
|
||||
identities = make(map[string]*OAuth)
|
||||
)
|
||||
|
||||
func getInstance(subject string) *OAuth {
|
||||
oauthMu.Lock()
|
||||
defer oauthMu.Unlock()
|
||||
return identities[subject]
|
||||
}
|
||||
|
||||
func addInstance(subject string, identity *OAuth) {
|
||||
oauthMu.Lock()
|
||||
defer oauthMu.Unlock()
|
||||
identities[subject] = identity
|
||||
}
|
||||
|
||||
/* func init() {
|
||||
func init() {
|
||||
registry.AddCtx("oauth", NewOauthFromConfig)
|
||||
}
|
||||
|
||||
func NewOauthFromConfig(ctx context.Context, other map[string]any) (Authorizer, error) {
|
||||
oauthMu.Lock()
|
||||
defer oauthMu.Unlock()
|
||||
// parse oauth config from yaml
|
||||
var cc oauth2.Config
|
||||
func NewOauthFromConfig(ctx context.Context, other map[string]any) (oauth2.TokenSource, error) {
|
||||
var cc struct {
|
||||
Name string
|
||||
oauth2.Config `mapstructure:",squash"`
|
||||
}
|
||||
|
||||
if err := util.DecodeOther(other, &cc); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return NewOauth(ctx, cc)
|
||||
} */
|
||||
return NewOauth(ctx, &cc.Config, cc.Name)
|
||||
}
|
||||
|
||||
func NewOauth(ctx context.Context, cc oauth2.Config, instanceName string) (*OAuth, error) {
|
||||
func NewOauth(ctx context.Context, oc *oauth2.Config, instanceName string) (oauth2.TokenSource, error) {
|
||||
log := util.NewLogger("oauth-generic")
|
||||
|
||||
if instanceName == "" {
|
||||
return nil, errors.New("instance name must not be empty")
|
||||
}
|
||||
|
||||
// generate json string from oauth2 config
|
||||
bytejson, _ := json.Marshal(cc)
|
||||
|
||||
h := sha256.New()
|
||||
h.Write(bytejson)
|
||||
fullHash := hex.EncodeToString(h.Sum(nil))
|
||||
sha256_hash := fullHash[:8]
|
||||
|
||||
subject := instanceName + " (" + sha256_hash + ")"
|
||||
// hash oauth2 config
|
||||
h := sha256.Sum256(fmt.Append(nil, oc))
|
||||
hash := hex.EncodeToString(h[:])[:8]
|
||||
subject := instanceName + " (" + hash + ")"
|
||||
|
||||
// reuse instance
|
||||
if instance := getInstance(subject); instance != nil {
|
||||
|
|
@ -83,7 +82,7 @@ func NewOauth(ctx context.Context, cc oauth2.Config, instanceName string) (*OAut
|
|||
// create new instance
|
||||
o := &OAuth{
|
||||
subject: subject,
|
||||
cc: cc,
|
||||
oc: oc,
|
||||
log: log,
|
||||
ctx: ctx,
|
||||
}
|
||||
|
|
@ -103,20 +102,12 @@ func NewOauth(ctx context.Context, cc oauth2.Config, instanceName string) (*OAut
|
|||
// add instance
|
||||
addInstance(o.subject, o)
|
||||
|
||||
// register authredirect
|
||||
providerauth.Register(o, subject)
|
||||
// register auth redirect
|
||||
providerauth.Register(subject, o)
|
||||
|
||||
return o, nil
|
||||
}
|
||||
|
||||
func (o *OAuth) Transport(base http.RoundTripper) http.RoundTripper {
|
||||
transport := oauth2.Transport{
|
||||
Base: base,
|
||||
Source: o,
|
||||
}
|
||||
return &transport
|
||||
}
|
||||
|
||||
// RefreshToken implements oauth.RefreshTokenSource.
|
||||
func (o *OAuth) RefreshToken(token *oauth2.Token) (*oauth2.Token, error) {
|
||||
if token.RefreshToken == "" {
|
||||
|
|
@ -127,7 +118,7 @@ func (o *OAuth) RefreshToken(token *oauth2.Token) (*oauth2.Token, error) {
|
|||
o.log.DEBUG.Printf("refreshing token for %s", o.subject)
|
||||
|
||||
// refresh token source
|
||||
token, err := o.cc.TokenSource(o.ctx, token).Token()
|
||||
token, err := o.oc.TokenSource(o.ctx, token).Token()
|
||||
if err != nil {
|
||||
if strings.Contains(err.Error(), "invalid_grant") {
|
||||
if settings.Exists(o.subject) {
|
||||
|
|
@ -148,13 +139,13 @@ func (o *OAuth) HandleCallback(responseValues url.Values) error {
|
|||
o.mu.Lock()
|
||||
defer o.mu.Unlock()
|
||||
|
||||
token, err := o.cc.Exchange(o.ctx, code, oauth2.VerifierOption(o.cv))
|
||||
token, err := o.oc.Exchange(o.ctx, code, oauth2.VerifierOption(o.cv))
|
||||
if err != nil {
|
||||
o.log.ERROR.Printf("error during oauth exchange: %s", err)
|
||||
return err
|
||||
}
|
||||
err = settings.SetJson(o.subject, token)
|
||||
if err != nil {
|
||||
|
||||
if err := settings.SetJson(o.subject, token); err != nil {
|
||||
o.log.ERROR.Printf("error saving token: %s", err)
|
||||
}
|
||||
|
||||
|
|
@ -168,7 +159,7 @@ func (o *OAuth) Login(state string) string {
|
|||
defer o.mu.Unlock()
|
||||
|
||||
o.cv = oauth2.GenerateVerifier()
|
||||
return o.cc.AuthCodeURL(state, oauth2.S256ChallengeOption(o.cv))
|
||||
return o.oc.AuthCodeURL(state, oauth2.S256ChallengeOption(o.cv))
|
||||
}
|
||||
|
||||
// Logout implements api.AuthProvider.
|
||||
|
|
@ -191,10 +182,6 @@ func (o *OAuth) DisplayName() string {
|
|||
|
||||
// Authenticated implements api.AuthProvider.
|
||||
func (o *OAuth) Authenticated() bool {
|
||||
// check if token is valid
|
||||
if token, err := o.TokenSource.Token(); err == nil {
|
||||
return token.Valid()
|
||||
} else {
|
||||
return false
|
||||
}
|
||||
token, err := o.TokenSource.Token()
|
||||
return err == nil && token.Valid()
|
||||
}
|
||||
|
|
|
|||
|
|
@ -39,10 +39,9 @@ func init() {
|
|||
|
||||
type Viessmann struct {
|
||||
client *http.Client
|
||||
ts oauth2.TokenSource
|
||||
}
|
||||
|
||||
func NewViessmannFromConfig(ctx context.Context, other map[string]any) (Authorizer, error) {
|
||||
func NewViessmannFromConfig(ctx context.Context, other map[string]any) (oauth2.TokenSource, error) {
|
||||
var cc struct {
|
||||
ClientID string
|
||||
User, Password string
|
||||
|
|
@ -57,22 +56,13 @@ func NewViessmannFromConfig(ctx context.Context, other map[string]any) (Authoriz
|
|||
|
||||
oc := OAuth2Config(cc.ClientID)
|
||||
|
||||
ctx = context.WithValue(context.Background(), oauth2.HTTPClient, v.client)
|
||||
ctx = context.WithValue(ctx, oauth2.HTTPClient, v.client)
|
||||
token, err := v.login(ctx, oc, cc.User, cc.Password)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
v.ts = oc.TokenSource(ctx, token)
|
||||
|
||||
return v, nil
|
||||
}
|
||||
|
||||
func (v *Viessmann) Transport(base http.RoundTripper) http.RoundTripper {
|
||||
return &oauth2.Transport{
|
||||
Base: base,
|
||||
Source: v.ts,
|
||||
}
|
||||
return oc.TokenSource(ctx, token), nil
|
||||
}
|
||||
|
||||
func (v *Viessmann) login(ctx context.Context, oc *oauth2.Config, user, password string) (*oauth2.Token, error) {
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ import (
|
|||
"github.com/evcc-io/evcc/util"
|
||||
"github.com/evcc-io/evcc/util/transport"
|
||||
"github.com/jpfielding/go-http-digest/pkg/digest"
|
||||
"golang.org/x/oauth2"
|
||||
)
|
||||
|
||||
// Auth is the authorization config
|
||||
|
|
@ -47,11 +48,14 @@ func (p *Auth) Transport(ctx context.Context, log *util.Logger, base http.RoundT
|
|||
p.Other["password"] = p.Password
|
||||
}
|
||||
|
||||
authorizer, err := auth.NewFromConfig(ctx, p.Source, p.Other)
|
||||
ts, err := auth.NewFromConfig(ctx, p.Source, p.Other)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return authorizer.Transport(base), nil
|
||||
return &oauth2.Transport{
|
||||
Source: ts,
|
||||
Base: base,
|
||||
}, nil
|
||||
}
|
||||
}
|
||||
|
|
|
|||
171
server/providerauth/handler.go
Normal file
171
server/providerauth/handler.go
Normal file
|
|
@ -0,0 +1,171 @@
|
|||
package providerauth
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/evcc-io/evcc/api"
|
||||
"github.com/evcc-io/evcc/core/keys"
|
||||
"github.com/evcc-io/evcc/util"
|
||||
)
|
||||
|
||||
// Handler manages a dynamic map of routes for handling the redirect during
|
||||
// OAuth authentication. When a route is registered a token OAuth state is returned.
|
||||
// On GET request the generic handler identifies route and target handler
|
||||
// by request state obtained from the request and delegates to the registered handler.
|
||||
type Handler struct {
|
||||
mu sync.Mutex
|
||||
secret []byte
|
||||
providers map[string]api.AuthProvider
|
||||
states map[string]string
|
||||
log *util.Logger
|
||||
}
|
||||
|
||||
func (a *Handler) Publish(paramC chan<- util.Param) {
|
||||
a.mu.Lock()
|
||||
defer a.mu.Unlock()
|
||||
|
||||
apMap := make(map[string]*AuthProvider)
|
||||
|
||||
for id, provider := range a.providers {
|
||||
ap := &AuthProvider{
|
||||
ID: url.QueryEscape(id),
|
||||
Authenticated: provider.Authenticated(),
|
||||
}
|
||||
apMap[provider.DisplayName()] = ap
|
||||
}
|
||||
|
||||
a.log.TRACE.Printf("publishing %d auth providers", len(apMap))
|
||||
|
||||
// publish the updated auth providers
|
||||
paramC <- util.Param{Key: keys.AuthProviders, Val: apMap}
|
||||
}
|
||||
|
||||
func (a *Handler) register(name string, handler api.AuthProvider) error {
|
||||
a.mu.Lock()
|
||||
defer a.mu.Unlock()
|
||||
|
||||
if a.providers[name] != nil {
|
||||
return fmt.Errorf("provider already registered: %s", name)
|
||||
}
|
||||
|
||||
a.log.DEBUG.Printf("registering provider: %s", name)
|
||||
a.providers[name] = handler
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (a *Handler) handleLogin(w http.ResponseWriter, r *http.Request) {
|
||||
id := r.URL.Query().Get("id")
|
||||
a.log.DEBUG.Printf("login request for: %s", id)
|
||||
|
||||
a.mu.Lock()
|
||||
defer a.mu.Unlock()
|
||||
|
||||
provider, ok := a.providers[id]
|
||||
if !ok {
|
||||
w.WriteHeader(http.StatusBadRequest)
|
||||
fmt.Fprintf(w, "invalid id")
|
||||
return
|
||||
}
|
||||
|
||||
// Generate a new state and store the provider
|
||||
state := NewState()
|
||||
encryptedState := state.Encrypt(a.secret)
|
||||
a.states[encryptedState] = id
|
||||
|
||||
// Schedule cleanup for stale state entries after state becomes invalid
|
||||
time.AfterFunc(stateValidity, func() {
|
||||
a.mu.Lock()
|
||||
defer a.mu.Unlock()
|
||||
delete(a.states, encryptedState)
|
||||
})
|
||||
|
||||
// return authorization URL
|
||||
res := struct {
|
||||
LoginUri string `json:"loginUri"`
|
||||
}{
|
||||
LoginUri: provider.Login(encryptedState),
|
||||
}
|
||||
|
||||
if err := json.NewEncoder(w).Encode(res); err != nil {
|
||||
a.log.ERROR.Printf("failed to encode login URI response: %v", err)
|
||||
}
|
||||
|
||||
w.WriteHeader(http.StatusFound)
|
||||
}
|
||||
|
||||
func (a *Handler) handleLogout(w http.ResponseWriter, r *http.Request) {
|
||||
id := r.URL.Query().Get("id")
|
||||
a.log.DEBUG.Printf("logout request for: %s", id)
|
||||
|
||||
a.mu.Lock()
|
||||
defer a.mu.Unlock()
|
||||
|
||||
provider, ok := a.providers[id]
|
||||
if !ok {
|
||||
w.WriteHeader(http.StatusBadRequest)
|
||||
fmt.Fprintf(w, "invalid id")
|
||||
return
|
||||
}
|
||||
|
||||
// Handle logout
|
||||
if err := provider.Logout(); err != nil {
|
||||
a.log.ERROR.Printf("logout for provider %s failed: %v", id, err)
|
||||
}
|
||||
|
||||
http.Redirect(w, r, "/", http.StatusFound)
|
||||
}
|
||||
|
||||
func (a *Handler) handleCallback(w http.ResponseWriter, r *http.Request) {
|
||||
q := r.URL.Query()
|
||||
|
||||
if q.Has("error") {
|
||||
w.WriteHeader(http.StatusBadRequest)
|
||||
fmt.Fprintf(w, "error: %s: %s\n", q.Get("error"), q.Get("error_description"))
|
||||
return
|
||||
}
|
||||
|
||||
encryptedState := q.Get("state")
|
||||
state, err := DecryptState(encryptedState, a.secret)
|
||||
if err != nil || !state.Valid() {
|
||||
w.WriteHeader(http.StatusBadRequest)
|
||||
fmt.Fprintf(w, "invalid state")
|
||||
return
|
||||
}
|
||||
|
||||
a.mu.Lock()
|
||||
defer a.mu.Unlock()
|
||||
|
||||
// Find the corresponding provider
|
||||
id, ok := a.states[encryptedState]
|
||||
if !ok {
|
||||
w.WriteHeader(http.StatusBadRequest)
|
||||
fmt.Fprintf(w, "no provider found for state")
|
||||
return
|
||||
}
|
||||
|
||||
provider, ok := a.providers[id]
|
||||
if !ok {
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
fmt.Fprintf(w, "internal provider state unexpected")
|
||||
return
|
||||
}
|
||||
|
||||
// Remove the state from the map
|
||||
delete(a.states, encryptedState)
|
||||
|
||||
// Handle the callback
|
||||
if err := provider.HandleCallback(r.URL.Query()); err != nil {
|
||||
a.log.ERROR.Printf("callback handling for provider %s failed: %v", id, err)
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
fmt.Fprintf(w, "callback handling failed")
|
||||
return
|
||||
}
|
||||
|
||||
http.Redirect(w, r, "/", http.StatusFound)
|
||||
}
|
||||
|
|
@ -2,35 +2,18 @@ package providerauth
|
|||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/evcc-io/evcc/api"
|
||||
"github.com/evcc-io/evcc/core/keys"
|
||||
"github.com/evcc-io/evcc/util"
|
||||
"github.com/gorilla/mux"
|
||||
)
|
||||
|
||||
var instance *Handler
|
||||
|
||||
// Handler manages a dynamic map of routes for handling the redirect during
|
||||
// OAuth authentication. When a route is registered a token OAuth state is returned.
|
||||
// On GET request the generic handler identifies route and target handler
|
||||
// by request state obtained from the request and delegates to the registered handler.
|
||||
type Handler struct {
|
||||
mu sync.Mutex
|
||||
secret []byte
|
||||
providers map[string]api.AuthProvider
|
||||
states map[string]string
|
||||
log *util.Logger
|
||||
}
|
||||
|
||||
type AuthProvider struct {
|
||||
ID string `json:"id"`
|
||||
Authenticated bool `json:"authenticated"`
|
||||
|
|
@ -38,9 +21,7 @@ type AuthProvider struct {
|
|||
|
||||
func init() {
|
||||
var secret [16]byte
|
||||
_, err := io.ReadFull(rand.Reader, secret[:])
|
||||
|
||||
if err != nil {
|
||||
if _, err := io.ReadFull(rand.Reader, secret[:]); err != nil {
|
||||
panic(err)
|
||||
}
|
||||
|
||||
|
|
@ -71,172 +52,7 @@ func Setup(router *mux.Router, paramC chan<- util.Param) {
|
|||
}()
|
||||
}
|
||||
|
||||
func (a *Handler) Publish(paramC chan<- util.Param) {
|
||||
a.mu.Lock()
|
||||
defer a.mu.Unlock()
|
||||
|
||||
apMap := make(map[string]*AuthProvider)
|
||||
|
||||
for id, provider := range a.providers {
|
||||
ap := &AuthProvider{
|
||||
ID: url.QueryEscape(id),
|
||||
Authenticated: provider.Authenticated(),
|
||||
}
|
||||
apMap[provider.DisplayName()] = ap
|
||||
}
|
||||
|
||||
a.log.TRACE.Printf("publishing %d auth providers", len(apMap))
|
||||
|
||||
// publish the updated auth providers
|
||||
paramC <- util.Param{Key: keys.AuthProviders, Val: apMap}
|
||||
}
|
||||
|
||||
// Register registers a specific AuthProvider. Returns login path as string.
|
||||
func Register(handler api.AuthProvider, name string) error {
|
||||
return instance.register(handler, name)
|
||||
}
|
||||
|
||||
func (a *Handler) register(handler api.AuthProvider, name string) error {
|
||||
a.mu.Lock()
|
||||
defer a.mu.Unlock()
|
||||
|
||||
if a.providers[name] != nil {
|
||||
a.log.ERROR.Printf("provider with name %s already registered", name)
|
||||
return errors.New("provider already registered")
|
||||
}
|
||||
a.log.INFO.Printf("registering oauth provider: %s", name)
|
||||
a.providers[name] = handler
|
||||
return nil
|
||||
}
|
||||
|
||||
func (a *Handler) handleLogin(w http.ResponseWriter, r *http.Request) {
|
||||
// Find corresponding provider
|
||||
q := r.URL.Query()
|
||||
id := q.Get("id")
|
||||
if id == "" {
|
||||
w.WriteHeader(http.StatusBadRequest)
|
||||
fmt.Fprintf(w, "missing id")
|
||||
return
|
||||
}
|
||||
|
||||
a.log.DEBUG.Printf("login request for provider: %s", id)
|
||||
|
||||
a.mu.Lock()
|
||||
defer a.mu.Unlock()
|
||||
|
||||
provider, ok := a.providers[id]
|
||||
if !ok {
|
||||
w.WriteHeader(http.StatusBadRequest)
|
||||
fmt.Fprintf(w, "invalid id")
|
||||
return
|
||||
}
|
||||
|
||||
// Generate a new state and store the provider
|
||||
state := util.NewState()
|
||||
encryptedState := state.Encrypt(a.secret)
|
||||
a.states[encryptedState] = id
|
||||
|
||||
// Schedule cleanup for stale state entries after state becomes invalid
|
||||
go func(state string) {
|
||||
time.Sleep(util.StateValidity)
|
||||
a.mu.Lock()
|
||||
defer a.mu.Unlock()
|
||||
delete(a.states, state)
|
||||
}(encryptedState)
|
||||
|
||||
// Build authorization URL
|
||||
loginUri := provider.Login(encryptedState)
|
||||
|
||||
responseVal := struct {
|
||||
LoginUri string `json:"loginUri"`
|
||||
}{
|
||||
LoginUri: loginUri,
|
||||
}
|
||||
if err := json.NewEncoder(w).Encode(responseVal); err != nil {
|
||||
a.log.ERROR.Printf("failed to encode login URI response: %v", err)
|
||||
}
|
||||
w.WriteHeader(http.StatusFound)
|
||||
}
|
||||
|
||||
func (a *Handler) handleLogout(w http.ResponseWriter, r *http.Request) {
|
||||
// Find corresponding provider
|
||||
q := r.URL.Query()
|
||||
id := q.Get("id")
|
||||
if id == "" {
|
||||
w.WriteHeader(http.StatusBadRequest)
|
||||
fmt.Fprintf(w, "missing id")
|
||||
return
|
||||
}
|
||||
|
||||
a.mu.Lock()
|
||||
defer a.mu.Unlock()
|
||||
|
||||
provider, ok := a.providers[id]
|
||||
if !ok {
|
||||
w.WriteHeader(http.StatusBadRequest)
|
||||
fmt.Fprintf(w, "invalid id")
|
||||
return
|
||||
}
|
||||
|
||||
// Handle logout
|
||||
if err := provider.Logout(); err != nil {
|
||||
a.log.ERROR.Printf("logout for provider %s failed: %v", id, err)
|
||||
}
|
||||
|
||||
http.Redirect(w, r, "/", http.StatusFound)
|
||||
}
|
||||
|
||||
func (a *Handler) handleCallback(w http.ResponseWriter, r *http.Request) {
|
||||
q := r.URL.Query()
|
||||
|
||||
if q.Has("error") {
|
||||
w.WriteHeader(http.StatusBadRequest)
|
||||
fmt.Fprintf(w, "error: %s: %s\n", q.Get("error"), q.Get("error_description"))
|
||||
return
|
||||
}
|
||||
|
||||
encryptedState := q.Get("state")
|
||||
state, err := util.DecryptState(encryptedState, a.secret)
|
||||
if err != nil {
|
||||
w.WriteHeader(http.StatusBadRequest)
|
||||
fmt.Fprintf(w, "failed to decrypt state")
|
||||
return
|
||||
}
|
||||
|
||||
if err := state.Validate(); err != nil {
|
||||
w.WriteHeader(http.StatusBadRequest)
|
||||
fmt.Fprintf(w, "invalid state")
|
||||
return
|
||||
}
|
||||
|
||||
a.mu.Lock()
|
||||
defer a.mu.Unlock()
|
||||
|
||||
// Find the corresponding provider
|
||||
id, ok := a.states[encryptedState]
|
||||
if !ok {
|
||||
w.WriteHeader(http.StatusBadRequest)
|
||||
fmt.Fprintf(w, "no provider found for state")
|
||||
return
|
||||
}
|
||||
|
||||
provider, ok := a.providers[id]
|
||||
if !ok {
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
fmt.Fprintf(w, "internal provider state unexpected")
|
||||
return
|
||||
}
|
||||
|
||||
// Remove the state from the map
|
||||
delete(a.states, encryptedState)
|
||||
|
||||
// Handle the callback
|
||||
if err := provider.HandleCallback(r.URL.Query()); err != nil {
|
||||
a.log.ERROR.Printf("callback handling for provider %s failed: %v", id, err)
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
fmt.Fprintf(w, "callback handling failed")
|
||||
return
|
||||
}
|
||||
|
||||
http.Redirect(w, r, "/", http.StatusFound)
|
||||
func Register(name string, handler api.AuthProvider) error {
|
||||
return instance.register(name, handler)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
package util
|
||||
package providerauth
|
||||
|
||||
import (
|
||||
"crypto/aes"
|
||||
|
|
@ -12,17 +12,15 @@ import (
|
|||
"time"
|
||||
)
|
||||
|
||||
var ErrStateExpired = fmt.Errorf("state expired")
|
||||
|
||||
const StateValidity = 2 * time.Minute
|
||||
const stateValidity = 2 * time.Minute
|
||||
|
||||
type State struct {
|
||||
Time time.Time
|
||||
Created time.Time
|
||||
}
|
||||
|
||||
func NewState() State {
|
||||
return State{
|
||||
Time: time.Now(),
|
||||
Created: time.Now(),
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -84,10 +82,6 @@ func (c *State) Encrypt(key []byte) string {
|
|||
return base64.URLEncoding.EncodeToString(ciphertext)
|
||||
}
|
||||
|
||||
func (c *State) Validate() error {
|
||||
if time.Since(c.Time) <= StateValidity {
|
||||
return nil
|
||||
}
|
||||
|
||||
return ErrStateExpired
|
||||
func (c *State) Valid() bool {
|
||||
return time.Since(c.Created) <= stateValidity
|
||||
}
|
||||
|
|
@ -19,11 +19,11 @@ type VolvoConnected struct {
|
|||
}
|
||||
|
||||
func init() {
|
||||
registry.Add("volvo-connected", NewVolvoConnectedFromConfig)
|
||||
registry.AddCtx("volvo-connected", NewVolvoConnectedFromConfig)
|
||||
}
|
||||
|
||||
// NewVolvoConnectedFromConfig creates a new VolvoConnected vehicle
|
||||
func NewVolvoConnectedFromConfig(other map[string]interface{}) (api.Vehicle, error) {
|
||||
func NewVolvoConnectedFromConfig(ctx context.Context, other map[string]interface{}) (api.Vehicle, error) {
|
||||
cc := struct {
|
||||
embed `mapstructure:",squash"`
|
||||
VIN string
|
||||
|
|
@ -42,14 +42,15 @@ func NewVolvoConnectedFromConfig(other map[string]interface{}) (api.Vehicle, err
|
|||
log := util.NewLogger("volvo-connected").Redact(cc.VIN, cc.VccApiKey)
|
||||
|
||||
// create oauth2 config
|
||||
config := connected.Oauth2Config(cc.Credentials.ID, cc.Credentials.Secret, cc.RedirectUri)
|
||||
ctx := context.WithValue(context.Background(), oauth2.HTTPClient, request.NewClient(log))
|
||||
authorizer, err := auth.NewOauth(ctx, *config, cc.embed.GetTitle())
|
||||
oc := connected.Oauth2Config(cc.Credentials.ID, cc.Credentials.Secret, cc.RedirectUri)
|
||||
ctx = context.WithValue(ctx, oauth2.HTTPClient, request.NewClient(log))
|
||||
|
||||
ts, err := auth.NewOauth(ctx, oc, cc.embed.GetTitle())
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
api := connected.NewAPI(log, cc.VccApiKey, authorizer)
|
||||
api := connected.NewAPI(log, cc.VccApiKey, ts)
|
||||
|
||||
cc.VIN, err = ensureVehicle(cc.VIN, api.Vehicles)
|
||||
|
||||
|
|
|
|||
|
|
@ -3,11 +3,11 @@ package connected
|
|||
import (
|
||||
"fmt"
|
||||
|
||||
"github.com/evcc-io/evcc/plugin/auth"
|
||||
"github.com/evcc-io/evcc/util"
|
||||
"github.com/evcc-io/evcc/util/request"
|
||||
"github.com/evcc-io/evcc/util/transport"
|
||||
"github.com/samber/lo"
|
||||
"golang.org/x/oauth2"
|
||||
)
|
||||
|
||||
// api constants
|
||||
|
|
@ -21,18 +21,20 @@ type API struct {
|
|||
}
|
||||
|
||||
// NewAPI creates a new api client
|
||||
func NewAPI(log *util.Logger, vccapikey string, authorizer auth.Authorizer) *API {
|
||||
func NewAPI(log *util.Logger, vccapikey string, ts oauth2.TokenSource) *API {
|
||||
v := &API{
|
||||
Helper: request.NewHelper(log),
|
||||
}
|
||||
|
||||
decoratedTransport := &transport.Decorator{
|
||||
Base: v.Client.Transport,
|
||||
Decorator: transport.DecorateHeaders(map[string]string{
|
||||
"vcc-api-key": vccapikey,
|
||||
}),
|
||||
v.Client.Transport = &oauth2.Transport{
|
||||
Source: ts,
|
||||
Base: &transport.Decorator{
|
||||
Decorator: transport.DecorateHeaders(map[string]string{
|
||||
"vcc-api-key": vccapikey,
|
||||
}),
|
||||
Base: v.Client.Transport,
|
||||
},
|
||||
}
|
||||
v.Client.Transport = authorizer.Transport(decoratedTransport)
|
||||
|
||||
return v
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue