Improve oauth integration (#21266)
This commit is contained in:
parent
4c73cbe197
commit
11e769ea22
18 changed files with 219 additions and 208 deletions
|
|
@ -1,20 +1,19 @@
|
|||
package auth
|
||||
|
||||
// TODO
|
||||
// - configurable redirect uri
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"github.com/evcc-io/evcc/api"
|
||||
"github.com/evcc-io/evcc/server/db/settings"
|
||||
"github.com/evcc-io/evcc/server/oauth2redirect"
|
||||
"github.com/evcc-io/evcc/server/providerauth"
|
||||
"github.com/evcc-io/evcc/util"
|
||||
"github.com/evcc-io/evcc/util/oauth"
|
||||
"golang.org/x/oauth2"
|
||||
|
|
@ -31,7 +30,7 @@ type OAuth struct {
|
|||
}
|
||||
|
||||
var (
|
||||
oauthMu sync.Mutex
|
||||
// oauthMu sync.Mutex
|
||||
identities = make(map[string]*OAuth)
|
||||
)
|
||||
|
||||
|
|
@ -43,7 +42,7 @@ func addInstance(subject string, identity *OAuth) {
|
|||
identities[subject] = identity
|
||||
}
|
||||
|
||||
func init() {
|
||||
/* func init() {
|
||||
registry.AddCtx("oauth", NewOauthFromConfig)
|
||||
}
|
||||
|
||||
|
|
@ -57,23 +56,24 @@ func NewOauthFromConfig(ctx context.Context, other map[string]any) (Authorizer,
|
|||
}
|
||||
|
||||
return NewOauth(ctx, cc)
|
||||
}
|
||||
} */
|
||||
|
||||
func NewOauth(ctx context.Context, cc oauth2.Config) (*OAuth, error) {
|
||||
func NewOauth(ctx context.Context, cc oauth2.Config, instanceName string) (*OAuth, error) {
|
||||
log := util.NewLogger("oauth-generic")
|
||||
|
||||
// generate json string from oauth2 config
|
||||
bytejson, err := json.Marshal(cc)
|
||||
|
||||
if err != nil {
|
||||
log.ERROR.Printf("error converting oauth config to json: %s", err)
|
||||
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)
|
||||
sha1_hash := hex.EncodeToString(h.Sum(nil))
|
||||
fullHash := hex.EncodeToString(h.Sum(nil))
|
||||
sha256_hash := fullHash[:8]
|
||||
|
||||
subject := sha1_hash
|
||||
subject := instanceName + " (" + sha256_hash + ")"
|
||||
|
||||
// reuse instance
|
||||
if instance := getInstance(subject); instance != nil {
|
||||
|
|
@ -104,7 +104,7 @@ func NewOauth(ctx context.Context, cc oauth2.Config) (*OAuth, error) {
|
|||
addInstance(o.subject, o)
|
||||
|
||||
// register authredirect
|
||||
oauth2redirect.Register(o, subject)
|
||||
providerauth.Register(o, subject)
|
||||
|
||||
return o, nil
|
||||
}
|
||||
|
|
@ -119,9 +119,6 @@ func (o *OAuth) Transport(base http.RoundTripper) http.RoundTripper {
|
|||
|
||||
// RefreshToken implements oauth.RefreshTokenSource.
|
||||
func (o *OAuth) RefreshToken(token *oauth2.Token) (*oauth2.Token, error) {
|
||||
o.mu.Lock()
|
||||
defer o.mu.Unlock()
|
||||
|
||||
if token.RefreshToken == "" {
|
||||
return nil, api.ErrMissingToken
|
||||
}
|
||||
|
|
@ -144,19 +141,9 @@ func (o *OAuth) RefreshToken(token *oauth2.Token) (*oauth2.Token, error) {
|
|||
return token, err
|
||||
}
|
||||
|
||||
// AuthCodeURL implements api.AuthProvider.
|
||||
func (o *OAuth) AuthCodeURL(state string) string {
|
||||
o.mu.Lock()
|
||||
defer o.mu.Unlock()
|
||||
|
||||
o.cv = oauth2.GenerateVerifier()
|
||||
return o.cc.AuthCodeURL(state, oauth2.S256ChallengeOption(o.cv))
|
||||
}
|
||||
|
||||
// HandleCallback implements api.AuthProvider.
|
||||
func (o *OAuth) HandleCallback(r *http.Request) {
|
||||
q := r.URL.Query()
|
||||
code := q.Get("code")
|
||||
func (o *OAuth) HandleCallback(responseValues url.Values) error {
|
||||
code := responseValues.Get("code")
|
||||
|
||||
o.mu.Lock()
|
||||
defer o.mu.Unlock()
|
||||
|
|
@ -164,7 +151,7 @@ func (o *OAuth) HandleCallback(r *http.Request) {
|
|||
token, err := o.cc.Exchange(o.ctx, code, oauth2.VerifierOption(o.cv))
|
||||
if err != nil {
|
||||
o.log.ERROR.Printf("error during oauth exchange: %s", err)
|
||||
return
|
||||
return err
|
||||
}
|
||||
err = settings.SetJson(o.subject, token)
|
||||
if err != nil {
|
||||
|
|
@ -172,10 +159,20 @@ func (o *OAuth) HandleCallback(r *http.Request) {
|
|||
}
|
||||
|
||||
o.TokenSource = oauth.RefreshTokenSource(token, o)
|
||||
return nil
|
||||
}
|
||||
|
||||
// HandleLogout implements api.AuthProvider.
|
||||
func (o *OAuth) HandleLogout(r *http.Request) {
|
||||
// Login implements api.AuthProvider.
|
||||
func (o *OAuth) Login(state string) string {
|
||||
o.mu.Lock()
|
||||
defer o.mu.Unlock()
|
||||
|
||||
o.cv = oauth2.GenerateVerifier()
|
||||
return o.cc.AuthCodeURL(state, oauth2.S256ChallengeOption(o.cv))
|
||||
}
|
||||
|
||||
// Logout implements api.AuthProvider.
|
||||
func (o *OAuth) Logout() error {
|
||||
o.log.INFO.Printf("removing %s from database", o.subject)
|
||||
if settings.Exists(o.subject) {
|
||||
settings.Delete(o.subject)
|
||||
|
|
@ -184,4 +181,20 @@ func (o *OAuth) HandleLogout(r *http.Request) {
|
|||
o.mu.Lock()
|
||||
defer o.mu.Unlock()
|
||||
o.TokenSource = oauth.RefreshTokenSource(nil, o)
|
||||
return nil
|
||||
}
|
||||
|
||||
// DisplayName implements api.AuthProvider.
|
||||
func (o *OAuth) DisplayName() string {
|
||||
return o.subject
|
||||
}
|
||||
|
||||
// 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
|
||||
}
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue