diff --git a/api/api.go b/api/api.go index 385b11404..a4d89be88 100644 --- a/api/api.go +++ b/api/api.go @@ -203,8 +203,7 @@ type WebController interface { // ProviderLogin is the ability to provide OAuth authentication through the ui type ProviderLogin interface { - SetCallbackParams(uri string, authenticated chan<- bool) + SetCallbackParams(baseURL, redirectURL string, authenticated chan<- bool) LoginHandler() http.HandlerFunc LogoutHandler() http.HandlerFunc - CallbackHandler(baseURI string) http.HandlerFunc } diff --git a/cmd/config.go b/cmd/config.go index 12079c5ac..fee86303f 100644 --- a/cmd/config.go +++ b/cmd/config.go @@ -3,8 +3,6 @@ package cmd import ( "fmt" "net/http" - "net/url" - "strings" "time" "github.com/dustin/go-humanize" @@ -14,6 +12,7 @@ import ( "github.com/evcc-io/evcc/provider/mqtt" "github.com/evcc-io/evcc/push" "github.com/evcc-io/evcc/server" + autoauth "github.com/evcc-io/evcc/server/auth" "github.com/evcc-io/evcc/util" "github.com/evcc-io/evcc/vehicle" "github.com/evcc-io/evcc/vehicle/wrapper" @@ -203,51 +202,39 @@ func (cp *ConfigProvider) configureVehicles(conf config) error { return nil } -func canonicalName(s string) string { - return strings.ToLower(strings.ReplaceAll(s, " ", "_")) -} - // webControl handles routing for devices. For now only api.ProviderLogin related routes func (cp *ConfigProvider) webControl(httpd *server.HTTPd, paramC chan<- util.Param) { router := httpd.Router() - auth := router.PathPrefix("/auth").Subrouter() + auth := router.PathPrefix("/oauth").Subrouter() auth.Use(handlers.CompressHandler) auth.Use(handlers.CORS( handlers.AllowedHeaders([]string{"Content-Type"}), )) + // wire the handler + autoauth.Setup(auth) + // initialize cp.auth = util.NewAuthCollection(paramC) + // TODO make evccURI configurable, add warnings for any network/ localhost + evccURI := fmt.Sprintf("http://%s", httpd.Addr) + authURI := fmt.Sprintf("%s/oauth", evccURI) + + var id int for _, v := range cp.vehicles { if provider, ok := v.(api.ProviderLogin); ok { - title := url.QueryEscape(canonicalName(v.Title())) - basePath := fmt.Sprintf("vehicles/%s", title) + id += 1 - // TODO make evccURI configurable, add warnings for any network/ localhost - evccURI := fmt.Sprintf("http://%s", httpd.Addr) + basePath := fmt.Sprintf("vehicles/%d", id) baseURI := fmt.Sprintf("%s/auth/%s", evccURI, basePath) // register vehicle - ap := cp.auth.Register(v.Title(), baseURI) + ap := cp.auth.Register(baseURI, v.Title()) - redirectURI := fmt.Sprintf("%s/callback", baseURI) - provider.SetCallbackParams(redirectURI, ap.Handler()) - log.INFO.Printf("ensure the oauth client redirect/callback is configured for %s: %s", v.Title(), redirectURI) + provider.SetCallbackParams(evccURI, authURI, ap.Handler()) - // TODO how to handle multiple vehicles of the same type - // - // problems, thoughts and ideas: - // conflicting callbacks! - // - some unique part has to be added. - // - or a general callback handler and the specific vehicle is transported in the state? - // - callback handler needs an option to set the token at the right vehicle and use the right code exchange - - auth. - Methods(http.MethodGet). - Path(fmt.Sprintf("/%s/callback", basePath)). - HandlerFunc(provider.CallbackHandler(evccURI)) auth. Methods(http.MethodPost). Path(fmt.Sprintf("/%s/login", basePath)). @@ -259,5 +246,9 @@ func (cp *ConfigProvider) webControl(httpd *server.HTTPd, paramC chan<- util.Par } } + if id > 0 { + log.INFO.Printf("ensure the oauth client redirect/callback is configured for: %s", authURI) + } + cp.auth.Publish() } diff --git a/server/auth/auth.go b/server/auth/auth.go new file mode 100644 index 000000000..dcfcd7f08 --- /dev/null +++ b/server/auth/auth.go @@ -0,0 +1,85 @@ +package auth + +import ( + "crypto/rand" + "fmt" + "io" + "net/http" + "sync" + + "github.com/evcc-io/evcc/util" + "github.com/gorilla/mux" +) + +var instance *Auth + +type Auth struct { + mu sync.Mutex + secret []byte + routes map[string]http.HandlerFunc +} + +func generateSecret() ([]byte, error) { + var b [16]byte + _, err := io.ReadFull(rand.Reader, b[:]) + return b[:], err +} + +func init() { + secret, err := generateSecret() + if err != nil { + panic(err) + } + + instance = &Auth{ + secret: secret, + routes: make(map[string]http.HandlerFunc), + } +} + +func Setup(router *mux.Router) { + router.Methods(http.MethodGet).HandlerFunc(instance.handle) +} + +func Register(handler http.HandlerFunc) string { + return instance.register(handler) +} + +func (a *Auth) register(handler http.HandlerFunc) string { + a.mu.Lock() + defer a.mu.Unlock() + + state := util.NewState() + key := state.Encrypt(a.secret) + + a.routes[key] = handler + + return key +} + +func (a *Auth) handle(w http.ResponseWriter, r *http.Request) { + vars := mux.Vars(r) + + if error, ok := vars["error"]; ok { + w.WriteHeader(http.StatusBadRequest) + fmt.Fprintf(w, "error: %s: %s\n", error, vars["error_description"]) + return + } + + state, err := util.DecryptState(vars["state"], a.secret) + if err == nil { + err = state.Validate() + } + + a.mu.Lock() + handler := a.routes[vars["state"]] + a.mu.Unlock() + + if err != nil || handler == nil { + w.WriteHeader(http.StatusBadRequest) + fmt.Fprintf(w, "invalid state") + return + } + + handler(w, r) +} diff --git a/util/providerauth.go b/util/providerauth.go index 72204f1e4..248c0d7c5 100644 --- a/util/providerauth.go +++ b/util/providerauth.go @@ -15,7 +15,7 @@ func NewAuthCollection(paramC chan<- Param) *AuthCollection { } } -func (ac *AuthCollection) Register(title, baseURI string) *AuthProvider { +func (ac *AuthCollection) Register(baseURI, title string) *AuthProvider { ap := &AuthProvider{ ac: ac, Uri: baseURI, diff --git a/vehicle/mercedes/state.go b/util/state.go similarity index 62% rename from vehicle/mercedes/state.go rename to util/state.go index fdeeb9ae3..1cef42f23 100644 --- a/vehicle/mercedes/state.go +++ b/util/state.go @@ -1,4 +1,4 @@ -package mercedes +package util import ( "crypto/aes" @@ -6,35 +6,64 @@ import ( "crypto/rand" "encoding/base64" "encoding/json" + "errors" "fmt" "io" "time" ) -var ErrExpiredState = fmt.Errorf("state expired") +var ErrStateExpired = fmt.Errorf("state expired") const stateValidity = 2 * time.Minute type State struct { - key []byte Time time.Time } -// TODO Move to another more general place in the repo -func NewState(key []byte) State { +func NewState() State { return State{ - key: key, Time: time.Now(), } } -func (c *State) Encrypt() string { +func DecryptState(enc string, key []byte) (*State, error) { + ciphertext, err := base64.URLEncoding.DecodeString(enc) + if err != nil { + return nil, fmt.Errorf("failed to base64 decode encrypted state: %w", err) + } + + block, err := aes.NewCipher(key) + if err != nil { + return nil, err + } + + // The IV needs to be unique, but not secure. Therefore it's common to + // include it at the beginning of the ciphertext. + if len(ciphertext) < aes.BlockSize { + return nil, errors.New("ciphertext too short") + } + + iv := ciphertext[:aes.BlockSize] + ciphertext = ciphertext[aes.BlockSize:] + + stream := cipher.NewCFBDecrypter(block, iv) + + // XORKeyStream can work in-place if the two arguments are the same. + stream.XORKeyStream(ciphertext, ciphertext) + + var state State + err = json.Unmarshal(ciphertext, &state) + + return &state, err +} + +func (c *State) Encrypt(key []byte) string { plain, err := json.Marshal(c) if err != nil { panic(err) } - block, err := aes.NewCipher(c.key) + block, err := aes.NewCipher(key) if err != nil { panic(err) } @@ -55,48 +84,10 @@ func (c *State) Encrypt() string { return base64.URLEncoding.EncodeToString(ciphertext) } -func Decrypt(enc string, key []byte) (State, error) { - ciphertext, err := base64.URLEncoding.DecodeString(enc) - if err != nil { - return State{}, fmt.Errorf("failed to base64 decode encrypted state: %w", err) - } - - block, err := aes.NewCipher(key) - if err != nil { - panic(err) - } - - // The IV needs to be unique, but not secure. Therefore it's common to - // include it at the beginning of the ciphertext. - if len(ciphertext) < aes.BlockSize { - panic("ciphertext too short") - } - - iv := ciphertext[:aes.BlockSize] - ciphertext = ciphertext[aes.BlockSize:] - - stream := cipher.NewCFBDecrypter(block, iv) - - // XORKeyStream can work in-place if the two arguments are the same. - stream.XORKeyStream(ciphertext, ciphertext) - - var state State - if err := json.Unmarshal(ciphertext, &state); err != nil { - return State{}, fmt.Errorf("failed to unmarshal encrypted state: %w", err) - } - - return state, nil -} - -func Validate(rawState string, encryptionKey []byte) error { - state, err := Decrypt(rawState, encryptionKey) - if err != nil { - return fmt.Errorf("failed to validate state: %w", err) - } - - if state.Time.Add(stateValidity).After(time.Now()) { +func (c *State) Validate() error { + if c.Time.Add(stateValidity).After(time.Now()) { return nil } - return ErrExpiredState + return ErrStateExpired } diff --git a/vehicle/mercedes/identity.go b/vehicle/mercedes/identity.go index 4079166b5..2609e78b1 100644 --- a/vehicle/mercedes/identity.go +++ b/vehicle/mercedes/identity.go @@ -2,16 +2,15 @@ package mercedes import ( "context" - "crypto/rand" "encoding/json" "fmt" - "io" "net/http" "net/url" "github.com/coreos/go-oidc" "github.com/evcc-io/evcc/api" "github.com/evcc-io/evcc/provider" + "github.com/evcc-io/evcc/server/auth" "github.com/evcc-io/evcc/util" "golang.org/x/oauth2" ) @@ -29,15 +28,9 @@ func WithToken(t *oauth2.Token) IdentityOptions { type Identity struct { log *util.Logger *ReuseTokenSource - sessionSecret []byte - oc *oauth2.Config - authC chan<- bool -} - -func generateSecret() ([]byte, error) { - var b [16]byte - _, err := io.ReadFull(rand.Reader, b[:]) - return b[:], err + oc *oauth2.Config + baseURL string + authC chan<- bool } // TODO SessionSecret from config/persistence @@ -66,7 +59,6 @@ func NewIdentity(log *util.Logger, id, secret string, options ...IdentityOptions ts.Apply(nil) v.ReuseTokenSource = ts - v.sessionSecret, err = generateSecret() for _, o := range options { if err == nil { @@ -86,19 +78,20 @@ func (v *Identity) invalidToken() { var _ api.ProviderLogin = (*Identity)(nil) -func (v *Identity) SetCallbackParams(uri string, authC chan<- bool) { - v.oc.RedirectURL = uri +func (v *Identity) SetCallbackParams(baseURL, redirectURL string, authC chan<- bool) { + v.baseURL = baseURL + v.oc.RedirectURL = redirectURL v.authC = authC } func (v *Identity) LoginHandler() http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { - state := NewState(v.sessionSecret) + state := auth.Register(v.callbackHandler) b, _ := json.Marshal(struct { LoginUri string `json:"loginUri"` }{ - LoginUri: v.oc.AuthCodeURL(state.Encrypt(), oauth2.AccessTypeOffline, + LoginUri: v.oc.AuthCodeURL(state, oauth2.AccessTypeOffline, oauth2.SetAuthURLParam("prompt", "login consent"), ), }) @@ -118,50 +111,34 @@ func (v *Identity) LogoutHandler() http.HandlerFunc { } } -func (v *Identity) CallbackHandler(baseURI string) http.HandlerFunc { - return func(w http.ResponseWriter, r *http.Request) { - v.log.TRACE.Println("callback request retrieved") +func (v *Identity) callbackHandler(w http.ResponseWriter, r *http.Request) { + v.log.TRACE.Println("callback request retrieved") - data, err := url.ParseQuery(r.URL.RawQuery) - if err != nil { - fmt.Fprintln(w, "invalid response:", data) - return - } - - if error, ok := data["error"]; ok { - fmt.Fprintf(w, "error: %s: %s\n", error, data["error_description"]) - return - } - - states, ok := data["state"] - if !ok || len(states) != 1 { - fmt.Fprintln(w, "invalid state response:", data) - return - } else if err := Validate(states[0], v.sessionSecret); err != nil { - fmt.Fprintf(w, "failed state validation: %s", err) - return - } - - codes, ok := data["code"] - if !ok || len(codes) != 1 { - fmt.Fprintln(w, "invalid response:", data) - return - } - - token, err := v.oc.Exchange(context.Background(), codes[0]) - if err != nil { - fmt.Fprintln(w, "token error:", err) - return - } - - if token.Valid() { - v.log.TRACE.Println("sending login update...") - v.ReuseTokenSource.Apply(token) - v.authC <- true - - provider.ResetCached() - } - - http.Redirect(w, r, baseURI, http.StatusFound) + data, err := url.ParseQuery(r.URL.RawQuery) + if err != nil { + fmt.Fprintln(w, "invalid response:", data) + return } + + codes, ok := data["code"] + if !ok || len(codes) != 1 { + fmt.Fprintln(w, "invalid response:", data) + return + } + + token, err := v.oc.Exchange(context.Background(), codes[0]) + if err != nil { + fmt.Fprintln(w, "token error:", err) + return + } + + if token.Valid() { + v.log.TRACE.Println("sending login update...") + v.ReuseTokenSource.Apply(token) + v.authC <- true + + provider.ResetCached() + } + + http.Redirect(w, r, v.baseURL, http.StatusFound) }