OAuth2: refactor authorization framework (BC) (#23978)

This commit is contained in:
andig 2025-09-30 18:35:10 +02:00 • committed by GitHub
parent c642732588
commit 2780278752
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
11 changed files with 241 additions and 311 deletions

View file

@ -1,9 +0,0 @@
package auth
import (
"net/http"
)
type Authorizer interface {
Transport(base http.RoundTripper) http.RoundTripper
}

View file

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

View file

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

View file

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

View file

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

View file

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

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

View file

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

View file

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

View file

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

View file

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