Provider auth: require authenticated session for OAuth login/logout (#33115)

Co-authored-by: Michael Geers <michael@geers.tv>
This commit is contained in:
Tilo Alexander 2026-08-24 15:47:06 +02:00 • committed by GitHub
parent 226425e22e
commit 4adbf141d6
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
7 changed files with 65 additions and 17 deletions

View file

@ -320,14 +320,14 @@ func (s *HTTPd) RegisterSystemHandler(site *core.Site, pub publisher, cache *uti
}
// API key endpoints require an authenticated session.
ensureAuth := ensureAuthHandler(auth)
ensureAuth := EnsureAuthHandler(auth)
api.Methods("GET").Path("/apikey").Handler(ensureAuth(apiKeyStatusHandler(auth)))
api.Methods("POST").Path("/apikey").Handler(ensureAuth(regenerateApiKeyHandler(auth)))
}
{ // api/config
api := api.PathPrefix("/config").Subrouter()
api.Use(ensureAuthHandler(auth))
api.Use(EnsureAuthHandler(auth))
routes := map[string]route{
"auth": {"POST", "/auth", authHandler},
@ -423,7 +423,7 @@ func (s *HTTPd) RegisterSystemHandler(site *core.Site, pub publisher, cache *uti
{ // api/system
api := api.PathPrefix("/system").Subrouter()
api.Use(ensureAuthHandler(auth))
api.Use(EnsureAuthHandler(auth))
routes := map[string]route{
"log": {"GET", "/log", logHandler},

View file

@ -186,7 +186,8 @@ func logoutHandler(w http.ResponseWriter, r *http.Request) {
})
}
func ensureAuthHandler(authObject auth.Auth) mux.MiddlewareFunc {
// EnsureAuthHandler returns middleware that rejects unauthenticated requests
func EnsureAuthHandler(authObject auth.Auth) mux.MiddlewareFunc {
return func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if next == nil {

View file

@ -75,3 +75,15 @@ func TestRequireCriticalConfig(t *testing.T) {
})
}
}
func TestEnsureAuthHandler(t *testing.T) {
get := func(a auth.Auth) int {
next := http.HandlerFunc(func(http.ResponseWriter, *http.Request) {})
rec := httptest.NewRecorder()
EnsureAuthHandler(a)(next).ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/login", nil))
return rec.Code
}
assert.Equal(t, http.StatusUnauthorized, get(fakeAuth{mode: auth.Enabled}), "enabled must reject without credentials")
assert.Equal(t, http.StatusOK, get(fakeAuth{mode: auth.Disabled}), "disabled must pass through")
}

View file

@ -32,14 +32,14 @@ func init() {
}
}
// Setup connects the redirect handler to the router and registers the callback channel
func Setup(router *mux.Router, paramC chan<- util.Param) {
// callback?code=...&state=...
// Setup connects the redirect handler to the router and registers the callback channel.
// Callback stays open: the cross-site IdP redirect carries no session cookie, the state token gates it.
func Setup(router *mux.Router, paramC chan<- util.Param, authMiddleware mux.MiddlewareFunc) {
gate := func(h http.HandlerFunc) http.Handler { return authMiddleware(h) }
router.Methods(http.MethodGet).Path("/callback").HandlerFunc(instance.handleCallback)
// login?id=...
router.Methods(http.MethodGet).Path("/login").HandlerFunc(instance.handleLogin)
// logout?id=...
router.Methods(http.MethodGet).Path("/logout").HandlerFunc(instance.handleLogout)
router.Methods(http.MethodGet).Path("/login").Handler(gate(instance.handleLogin))
router.Methods(http.MethodGet).Path("/logout").Handler(gate(instance.handleLogout))
go instance.run(paramC)
}

View file

@ -0,0 +1,34 @@
package providerauth
import (
"net/http"
"net/http/httptest"
"testing"
"github.com/evcc-io/evcc/util"
"github.com/gorilla/mux"
"github.com/stretchr/testify/assert"
)
// stand-in auth middleware rejecting every request
func blockAll(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
http.Error(w, "Unauthorized", http.StatusUnauthorized)
})
}
// login/logout are gated, callback stays open (302 to error page, not middleware 401)
func TestSetupGating(t *testing.T) {
router := mux.NewRouter()
Setup(router, make(chan util.Param, 1), blockAll)
get := func(path string) int {
rec := httptest.NewRecorder()
router.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, path, nil))
return rec.Code
}
assert.Equal(t, http.StatusUnauthorized, get("/login?id=x"), "login must be gated")
assert.Equal(t, http.StatusUnauthorized, get("/logout?id=x"), "logout must be gated")
assert.Equal(t, http.StatusFound, get("/callback"), "callback must stay open (gated by state token, not session)")
}