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) {
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue