Implement oauth login proxy (#2425)
This commit is contained in:
parent
c18eb34df6
commit
cb44888dba
6 changed files with 182 additions and 139 deletions
|
|
@ -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
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
}
|
||||
|
|
|
|||
85
server/auth/auth.go
Normal file
85
server/auth/auth.go
Normal file
|
|
@ -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)
|
||||
}
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
|
@ -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)
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue